-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathCap_block_fit.m
More file actions
187 lines (150 loc) · 6.5 KB
/
Copy pathCap_block_fit.m
File metadata and controls
187 lines (150 loc) · 6.5 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
%% Function for fitting the gut inference model to subject data (after pre-processing)
% Samuel Taylor and Ryan Smith
% 1/5/2021
% file: full path of csv file containing subject data
% options: model fitting options in a table
% subdat: contains preprocessed subject data (obtained from Cap_fit)
function results = Cap_block_fit(file, options, subdat)
TpB = size(subdat,1); %trials per block
% Store actions and observations in a structure array
sub.u = ones(1, 2, TpB);
sub.o = ones(1, 2, TpB);
action = ones(1, 1, TpB);
vib = ones(1, 1, TpB);
block = zeros(1, 1, TpB);
for i = 1:TpB
vib (1, 1, i) = subdat.Vib(i);
sub.o (1,2,i) = subdat.Vib(i) + 2;
action(1, 1,i) = subdat.Press(i);
sub.u (1, :,i) = [1 1];
block(1, 1, i) = subdat.Block(i);
end
o_all = sub.o;
u_all = sub.u;
% Count correct and incorrect choices
for i = 1:size(o_all,3)
if vib(:,:,i) == (action(:,:,i))
correct(i) = 1;
else
correct(i) = 0;
end
end
% Store accuracy of subject
accuracy = mean(correct)*100;
% Store true/false positive/negatives
TN = zeros(1,size(o_all,3));
TP = zeros(1,size(o_all,3));
FN = zeros(1,size(o_all,3));
FP = zeros(1,size(o_all,3));
for i = 1:size(o_all,3)
if vib(:,:,i) == action(:,:,i)
if vib(:,:,i) == 0
TN(i) = 1;
elseif vib(:,:,i) == 1
TP(i) = 1;
end
elseif vib(:,:,i) == 0
FP(i) = 1;
elseif vib(:,:,i) == 1
FN(i) = 1;
end
end
%storing true/false negatives/positives
TP_FP_FN_TN = array2table([TP; FP; FN; TN]', 'VariableNames', {'TP', 'FP', 'FN', 'TN'});
%% Params
%--------------------------------------------------------------------------
IP = 0.95; % precision of tone (0-1)
eta = 0.5; % learning rate (between 0-1; default = 1)
etaV = 0.5; % learning rate (between 0-1; default = 1)
etaNV = 0.5; % learning rate (between 0-1; default = 1)
pV = 0.5; % prior bias (0-1; higher = prior favoring detecting vibration)
IPdiff = 0.25;
%% Invert model and try to recover original parameters:
%==========================================================================
params = struct( ...
'IP', IP, ... # Represents IP in paper (Interoceptive Precision)
'pV', pV, ... # Represents pV in paper (prior)
'etaA', eta, ... # Single learning rate for A matrix
'etaAV', etaV, ... # Learning rate for vibrations in A matrix
'etaANV', etaNV, ... # Learning rate for no vibrations in A matrix
'etaB', eta, ... # Single learning rate for B matrix
'etaBV', etaV, ... # Learning rate for vibrations in B matrix
'etaBNV', etaNV, ... # Learning rate for no vibrations in B matrix
'IPdiff', IPdiff ... # Difference in IP between normal and enhanced blocks
);
% Generate model structure from model options
MDP = Cap_gen_mdp(params, options);
MDP.TpB = TpB; % trials per block
MDP.action = action + 2;
MDP.vib = vib;
MDP.block = block;
DCM.MDP = MDP; % MDP model
% Specify set of params to fit, based on selected
% options for the model.
DCM.field = {'IP' 'pV'};
if options.b_mode == 2
DCM.field = [DCM.field {'etaBV' 'etaBNV'}];
elseif options.b_mode == 1
DCM.field = [DCM.field {'etaB'}];
end
if options.a_mode == 2
DCM.field = [DCM.field {'etaAV' 'etaANV'}];
elseif options.a_mode == 1
DCM.field = [DCM.field {'etaA'}];
end
if options.IPdiff_on && options.a_mode == 0
DCM.field = [DCM.field {'IPdiff'}];
end
DCM.U = {o_all}; % trial specification (stimuli)
DCM.Y = {u_all}; % responses (action)
% Perform model inversion
DCM = Cap_inversion(DCM);
%--------------------------------------------------------------------------
% re-transform values and compare prior with posterior estimates
%--------------------------------------------------------------------------
field = fieldnames(DCM.M.pE);
prior = zeros(1, size(field, 1));
posterior = zeros(1, size(field, 1));
% Extract the fitted parameters, returning the parameters
% to the appropriate space
for i = 1:length(field)
disp(field{i});
if strcmp(field{i},'etaB')
prior(i) = 1/(1+exp(-DCM.M.pE.(field{i})));
posterior(i) = 1/(1+exp(-DCM.Ep.(field{i})));
elseif strcmp(field{i},'etaA')
prior(i) = 1/(1+exp(-DCM.M.pE.(field{i})));
posterior(i) = 1/(1+exp(-DCM.Ep.(field{i})));
elseif strcmp(field{i},'etaBV')
prior(i) = 1/(1+exp(-DCM.M.pE.(field{i})));
posterior(i) = 1/(1+exp(-DCM.Ep.(field{i})));
elseif strcmp(field{i},'etaBNV')
prior(i) = 1/(1+exp(-DCM.M.pE.(field{i})));
posterior(i) = 1/(1+exp(-DCM.Ep.(field{i})));
elseif strcmp(field{i},'etaAV')
prior(i) = 1/(1+exp(-DCM.M.pE.(field{i})));
posterior(i) = 1/(1+exp(-DCM.Ep.(field{i})));
elseif strcmp(field{i},'etaANV')
prior(i) = 1/(1+exp(-DCM.M.pE.(field{i})));
posterior(i) = 1/(1+exp(-DCM.Ep.(field{i})));
elseif strcmp(field{i},'IP')
prior(i) = 1/(1+exp(-DCM.M.pE.(field{i})));
posterior(i) = 1/(1+exp(-DCM.Ep.(field{i})));
elseif strcmp(field{i},'pV')
prior(i) = 1/(1+exp(-DCM.M.pE.(field{i})));
posterior(i) = 1/(1+exp(-DCM.Ep.(field{i})));
elseif strcmp(field{i},'IPdiff')
prior(i) = 1/(1+exp(-DCM.M.pE.(field{i})));
posterior(i) = 1/(1+exp(-DCM.Ep.(field{i})));
else
prior(i) = exp(DCM.M.pE.(field{i}));
posterior(i) = exp(DCM.Ep.(field{i}));
end
end
% Save all relevant model information and return the results
prior = array2table(prior, 'VariableNames', field);
posterior = array2table(posterior, 'VariableNames', field);
[model_acc, P_avg, button_accuracy, BP_avg, nobutton_accuracy, NBP_avg] = Cap_acc(posterior, options, file);
avg_delay = mean(subdat.Delay, 'omitnan');
results = {{file} prior posterior DCM accuracy TP_FP_FN_TN avg_delay model_acc P_avg button_accuracy BP_avg nobutton_accuracy NBP_avg};
end