-
Notifications
You must be signed in to change notification settings - Fork 2
Expand file tree
/
Copy pathmodel.py
More file actions
375 lines (316 loc) · 16 KB
/
Copy pathmodel.py
File metadata and controls
375 lines (316 loc) · 16 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
import numpy as np
import pandas as pd
from pgmpy.models import BayesianNetwork
from pgmpy.factors.discrete import TabularCPD
from pgmpy.inference import BeliefPropagation
from scipy.stats import entropy
import copy
import torch
import torch.distributions as dist
import logging
logging.getLogger("pgmpy").setLevel(logging.ERROR)
from sampler import Sampler
from dirichlet import Dirichlet
class CIM():
def __init__(
self,
# Hyperpriors
alpha_prior = 1., # habit learning hyperprior
beta_prior = 1., # task set learning hyperprior
theta_prior = 1., # context (context cue and volatility) learning hyperprior
# Option to hard code effects
p_a__c_s = None, # S-R association
p_o__c_s_a = None, # Task set
p_c__c_s = None, # Context associations
p_c_prev = None, # Previous context
p_i__c = None, # Instruction
# Other parameters
n_c = 2,
sampling_threshold = 100,
):
# Define hyperpriors
if p_a__c_s is None:
p_a__c_s = np.ones((2, n_c, 2, 2)) * .5
self.alpha = Dirichlet(shape=(2, 2, 2, 2), params=p_a__c_s * alpha_prior * 2)
if p_o__c_s_a is None:
p_o__c_s_a = np.ones((2, n_c, 2, 2, 2)) * .5
self.beta = Dirichlet(shape=(2, n_c, 2, 2, 2), params=p_o__c_s_a * beta_prior * 2)
self.beta_prior = beta_prior
if p_c__c_s is None:
p_c__c_s = np.ones((n_c, n_c, 2)) * (1/n_c)
self.theta = Dirichlet(shape=(n_c, n_c, 2), params=p_c__c_s * theta_prior * 2)
self.theta_prior = theta_prior
if p_i__c is None:
p_i__c = np.ones((n_c, n_c)) * (1/n_c)
self.p_c_prev = p_c_prev
if p_c_prev is None:
self.p_c_prev = np.ones(n_c) / n_c
# Define model structure (without previous context)
self.model = BayesianNetwork([
('s0', 'c'),
('s0', 'a'),
('s1', 'a'),
('s0', 'o'),
('s1', 'o'),
('a', 'o'),
('c_prev', 'c'),
('c', 'a'),
('c', 'o'),
('c', 'i'),
])
# Define CPDs
cpd_s0 = TabularCPD('s0', 2, [[.5], [.5]])
cpd_s1 = TabularCPD('s1', 2, [[.5], [.5]])
cpd_c_prev = TabularCPD('c_prev', n_c, self.p_c_prev.reshape(n_c, -1))
cpd_c = TabularCPD('c', n_c, evidence=['c_prev', 's0'], evidence_card=[n_c, 2],
values=self.theta.get_MAP_cpd().reshape(n_c, -1))
cpd_a = TabularCPD('a', 2, evidence=['c', 's0', 's1'], evidence_card=[n_c, 2, 2],
values=self.alpha.get_MAP_cpd().reshape(2, -1))
cpd_o = TabularCPD('o', 2, evidence=['c', 's0', 's1', 'a'], evidence_card=[n_c, 2, 2, 2],
values=self.beta.get_MAP_cpd().reshape(2, -1))
cpd_i = TabularCPD('i', n_c, evidence=['c'], evidence_card=[n_c],
values=p_i__c.reshape(n_c, -1))
self.model.add_cpds(cpd_s0, cpd_s1, cpd_c_prev, cpd_c, cpd_a, cpd_o, cpd_i)
# Init data saving
self.data_dict = [{}]
# Other parameters
self.n_c = n_c
self.sampling_threshold = sampling_threshold
self.first_trial = True
self.additional_counts = 0
self.prev_task = None
def goal_inference(self, stim, task, goal=0, sample=False, instruction=False):
self.stim = stim
self.task = task
self.goal = goal
self.sample = sample
# Inference
inference = BeliefPropagation(self.model)
self.goal_prior = inference.query(['a', 'c'],
evidence={'s0': self.stim[0],
's1': self.stim[1]}).values
if not sample:
if not instruction:
self.goal_posterior = inference.query(['a', 'c'],
evidence={'s0': self.stim[0],
's1': self.stim[1],
'o': self.goal}).values
else:
self.goal_posterior = inference.query(['a', 'c'],
evidence={'s0': self.stim[0],
's1': self.stim[1],
'i': self.task,
'o': self.goal}).values
n_samples = None
trace = None
elif sample:
inference = Sampler(self.model, threshold=self.sampling_threshold)
self.goal_posterior, n_samples, trace, self.sample_cost, self.sample_error, self.a_rate, self.n_reject, self.n_accept = inference.query(stim=stim, task=task, outcome=goal, return_it=True, include_instruction=instruction)
self.goal_posterior_a = self.goal_posterior.sum(axis=1)
self.goal_posterior_c = self.goal_posterior.sum(axis=0)
# IT measures
goal_cost = entropy(self.goal_posterior.flatten(), self.goal_prior.flatten())
goal_meta_cost = entropy(self.goal_posterior_c, self.goal_prior.sum(axis=0))
goal_control_cost = goal_cost - goal_meta_cost
likelihood_o = self.model.get_cpds('o').values[self.goal, :, self.stim[0], self.stim[1], :].T # transpose to [a, c]
goal_error = -1 * (self.goal_posterior.flatten() * np.log(likelihood_o.flatten())).sum()
self.goal_surprise = goal_cost + goal_error
if instruction:
likelihood_i = self.model.get_cpds('i').values[task, :]
instr_error = -1 * (self.goal_posterior_c * np.log(likelihood_i)).sum()
self.goal_surprise = goal_cost + goal_error + instr_error
# Store results
self.data_dict[-1].update({
'stim': self.stim,
'congruency': self.stim[0] == self.stim[1],
'repeat': self.prev_task == self.task,
'task': self.task,
'goal': self.goal,
'prior_a': self.goal_prior.sum(axis=1)[0],
'post_a': self.goal_posterior_a[0],
'prior_c': self.goal_prior.sum(axis=0)[0],
'post_c0': self.goal_posterior_c[0],
'control_cost': goal_control_cost,
'meta_cost': goal_meta_cost,
'error': goal_error,
'surprise_goal_inference': self.goal_surprise.copy(),
})
if sample:
self.data_dict[-1].update({
'goal_n_samples': n_samples,
'accept_rate': self.a_rate,
'n_c': self.n_c,
})
for i in range(1, self.n_c):
self.data_dict[-1].update({
f'post_c{i}': self.goal_posterior_c[i],
})
return trace
def outcome_inference(self):
# Soft evidence for p(a) and p(c)
self.model.remove_edges_from([('s0', 'a'), ('s1', 'a'), ('c', 'a'), ('s0', 'c'), ('c_prev', 'c')])
self.model.remove_node('c_prev')
cpd_c = TabularCPD('c', self.n_c, self.goal_posterior_c.reshape(self.n_c, -1))
cpd_a = TabularCPD('a', 2, self.goal_posterior_a.reshape(2, -1))
self.model.add_cpds(cpd_c, cpd_a)
# Inference
inference = BeliefPropagation(self.model)
outcome_prior = inference.query(
['c'],
evidence={'s0': self.stim[0],
's1': self.stim[1]}).values # TODO check
self.outcome_posterior_c = inference.query(
['c'],
evidence={'s0': self.stim[0],
's1': self.stim[1],
'a': self.action,
'o': self.outcome}).values
# IT measures
outcome_cost = entropy(self.outcome_posterior_c, outcome_prior)
likelihood_o = self.goal_posterior_a @ self.model.get_cpds('o').values[self.outcome, :, self.stim[0], self.stim[1], :].T
outcome_error = -1 * (self.outcome_posterior_c * np.log(likelihood_o)).sum()
self.outcome_surprise = outcome_cost + outcome_error
# Reinstantiate model
self.model.add_node('c_prev')
self.model.add_edges_from([('s0', 'a'), ('s1', 'a'), ('c', 'a'), ('s0', 'c'), ('c_prev', 'c')])
cpd_c_prev = TabularCPD('c_prev', self.n_c, (np.ones(self.n_c) / self.n_c).reshape(self.n_c, -1))
cpd_c = TabularCPD('c', self.n_c, evidence=['c_prev', 's0'], evidence_card=[self.n_c, 2],
values=self.theta.get_MAP_cpd().reshape(self.n_c, -1))
cpd_a = TabularCPD('a', 2, evidence=['c', 's0', 's1'], evidence_card=[self.n_c, 2, 2],
values=self.alpha.get_MAP_cpd().reshape(2, -1))
self.model.add_cpds(cpd_c_prev, cpd_c, cpd_a)
# Store results
self.data_dict[-1].update({
'posterior_c_outcome': self.outcome_posterior_c[0],
'outcome_cost': outcome_cost,
'outcome_error': outcome_error,
'outcome_surprise': self.outcome_surprise,
})
def interact(self, ground_truth, outcome_uncertainty=0, map_estimate=True):
self.action = np.argmax(self.goal_posterior_a) if map_estimate else np.random.choice([0, 1], p=self.goal_posterior_a)
self.outcome = ground_truth[self.task, self.stim[0], self.stim[1], self.action]
if outcome_uncertainty > 0:
self.outcome = np.random.choice([self.outcome, (self.outcome - 1) * -1], p=[1-outcome_uncertainty, outcome_uncertainty])
true_accuracy = self.goal_posterior_a[np.where(ground_truth[self.task, self.stim[0], self.stim[1], :] == self.goal)][0]
self.data_dict[-1].update({
'true_accuracy': true_accuracy,
'true_error': 1 - true_accuracy,
'chosen_action': int(self.action),
'outcome': int(self.outcome),
})
def update_context(self):
if hasattr(self, 'outcome_posterior_c') and self.outcome_posterior_c is not None:
self.p_c_prev = self.outcome_posterior_c
else:
self.p_c_prev = self.goal_posterior_c
cpd_c_prev = TabularCPD('c_prev', self.n_c, values=self.p_c_prev.reshape(self.n_c, -1))
self.model.add_cpds(cpd_c_prev)
self.prev_task = self.task
def learn_habit(self, dim=None):
obs = np.zeros((2, self.n_c, 2, 2))
if dim is None:
if self.first_trial: # TODO generalize for case where p_c_prev is given
obs[:, 0, self.stim[0], self.stim[1]] = self.goal_posterior_a
else:
obs[:, :, self.stim[0], self.stim[1]] = self.goal_posterior
elif dim == 0:
if self.first_trial:
obs[:, 0, self.stim[0], 0] = self.goal_posterior_a
obs[:, 0, self.stim[0], 1] = self.goal_posterior_a
else:
obs[:, :, self.stim[0], 0] = self.goal_posterior
obs[:, :, self.stim[0], 1] = self.goal_posterior
elif dim == 1:
if self.first_trial:
obs[:, 0, 0, self.stim[1]] = self.goal_posterior_a
obs[:, 0, 1, self.stim[1]] = self.goal_posterior_a
else:
obs[:, :, 0, self.stim[1]] = self.goal_posterior
obs[:, :, 1, self.stim[1]] = self.goal_posterior
# Infer parameters based on pseudo-counts
dir_tm1 = dist.Dirichlet(torch.tensor(self.alpha.values.copy()))
self.alpha.infer(obs)
cpd_a = TabularCPD('a', 2, evidence=['c', 's0', 's1'], evidence_card=[self.n_c, 2, 2],
values=self.alpha.get_MAP_cpd().reshape(2, -1))
self.model.add_cpds(cpd_a)
# Recalculate goal inference to get updated surprise after learning
inference = BeliefPropagation(self.model)
goal_prior = inference.query(['a', 'c'], evidence={'s0': self.stim[0], 's1': self.stim[1]}).values
goal_posterior = inference.query(['a', 'c'], evidence={'s0': self.stim[0], 's1': self.stim[1], 'o': self.goal}).values
# IT measures
habit_learn_cost = dist.kl.kl_divergence(dist.Dirichlet(torch.tensor(self.alpha.values.copy())), dir_tm1)
habit_learn_cost = float(habit_learn_cost.sum())
goal_cost = entropy(goal_posterior.flatten(), goal_prior.flatten())
likelihood_o = self.model.get_cpds('o').values[self.goal, :, self.stim[0], self.stim[1], :].T # transpose to [a, c]
goal_error = -1 * (goal_posterior.flatten() * np.log(likelihood_o.flatten())).sum()
overall_goal_surprise = goal_cost + goal_error + habit_learn_cost
self.data_dict[-1].update({
'habit_learn_cost': habit_learn_cost,
'overall_goal_surprise': overall_goal_surprise,
})
def learn_task_set(self):
# Pad outcome_posterior_c if a new context was created after inference (shape mismatch)
if hasattr(self, 'outcome_posterior_c') and self.outcome_posterior_c is not None:
posterior_c = self.outcome_posterior_c
else:
posterior_c = self.goal_posterior_c
# Get pseudo-counts
obs = np.zeros((2, self.n_c, 2, 2, 2))
if self.first_trial:
obs[self.outcome, 0, self.stim[0], self.stim[1], self.action] = 1
else:
obs[self.outcome, :, self.stim[0], self.stim[1], self.action] = posterior_c
# Infer parameters based on pseudo-counts
dir_tm1 = dist.Dirichlet(torch.tensor(self.beta.values.copy()))
self.beta.infer(obs)
cpd_o = TabularCPD('o', 2, evidence=['c', 's0', 's1', 'a'], evidence_card=[self.n_c, 2, 2, 2],
values=self.beta.get_MAP_cpd().reshape(2, -1))
self.model.add_cpds(cpd_o)
# Calculate costs
task_set_learn_cost = dist.kl.kl_divergence(dist.Dirichlet(torch.tensor(self.beta.values.copy())), dir_tm1)
task_set_learn_cost = float(task_set_learn_cost.sum())
self.data_dict[-1].update({
'task_set_learn_cost': task_set_learn_cost,
})
def learn_context(self, fr=0, learn_item=True, learn_temporal=True):
if hasattr(self, 'outcome_posterior_c') and self.outcome_posterior_c is not None:
posterior_c = self.outcome_posterior_c
else:
posterior_c = self.goal_posterior_c
# Get pseudo-counts
obs = np.zeros((self.n_c, self.n_c, 2))
if learn_item and learn_temporal:
obs[:, :, self.stim[0]] = np.outer(posterior_c, self.p_c_prev) # add information from c_t-1
elif not learn_item:
if self.first_trial:
obs[0, 0, :] = 1
else:
obs[:, 0, 0] = self.goal_posterior_c * self.p_c_prev[0]
obs[:, 0, 1] = self.goal_posterior_c * self.p_c_prev[0]
obs[:, 1, 0] = self.goal_posterior_c * self.p_c_prev[1]
obs[:, 1, 1] = self.goal_posterior_c * self.p_c_prev[1]
elif not learn_temporal:
if self.first_trial:
obs[0, :, self.stim[0]] = 1
else:
obs[:, 0, self.stim[0]] = self.goal_posterior_c
obs[:, 1, self.stim[0]] = self.goal_posterior_c
# Apply forgetting rate
self.theta.forget(fr=fr)
dir_tm1 = dist.Dirichlet(torch.tensor(self.theta.values.copy()))
# Infer parameters based on pseudo-counts
self.theta.infer(obs)
cpd_c = TabularCPD('c', self.n_c, evidence=['c_prev', 's0'], evidence_card=[self.n_c, 2],
values=self.theta.get_MAP_cpd().reshape(self.n_c, -1))
self.model.add_cpds(cpd_c)
# Calculate IT metrics
context_learn_cost = dist.kl.kl_divergence(dist.Dirichlet(torch.tensor(self.theta.values.copy())), dir_tm1)
context_learn_cost = float(context_learn_cost.sum())
self.data_dict[-1].update({
'context_learn_cost': context_learn_cost,
})
def save_data(self):
self.data = pd.DataFrame(self.data_dict)
self.data_dict.append({})
if self.first_trial:
self.first_trial = False