diff --git a/src/pspm_cfg/pspm_cfg_first_level_ps.m b/src/pspm_cfg/pspm_cfg_first_level_ps.m index bc27948f1..c11709f84 100644 --- a/src/pspm_cfg/pspm_cfg_first_level_ps.m +++ b/src/pspm_cfg/pspm_cfg_first_level_ps.m @@ -8,6 +8,6 @@ cfg = cfg_repeat; cfg.name = 'Pupil size'; cfg.tag = 'ps'; -cfg.values = {pspm_cfg_glm_ps_fc}; % Values in a cfg_repeat can be any cfg_item objects +cfg.values = {pspm_cfg_glm_ps_fc,pspm_cfg_tam}; % Values in a cfg_repeat can be any cfg_item objects cfg.forcestruct = true; cfg.help = {''}; diff --git a/src/pspm_cfg/pspm_cfg_glm_hp_rew.m b/src/pspm_cfg/pspm_cfg_glm_hp_rew.m index f5415e2d0..c63a42027 100644 --- a/src/pspm_cfg/pspm_cfg_glm_hp_rew.m +++ b/src/pspm_cfg/pspm_cfg_glm_hp_rew.m @@ -1,5 +1,5 @@ function [glm_hp_rew] = pspm_cfg_glm_hp_rew -% GLM HP FC +% GLM HP reward conditioning % Initialise global settings diff --git a/src/pspm_cfg/pspm_cfg_run_tam.m b/src/pspm_cfg/pspm_cfg_run_tam.m new file mode 100644 index 000000000..32c6816a9 --- /dev/null +++ b/src/pspm_cfg/pspm_cfg_run_tam.m @@ -0,0 +1,73 @@ +function out = pspm_cfg_run_tam(job) + +model = struct(); +options = struct(); + +%% Data & design +[newmodel, newoptions] = pspm_cfg_selector_data_design('run', job); + +fields = fieldnames(newmodel); +for i = 1:numel(fields) + model.(fields{i}) = newmodel.(fields{i}); +end + +%% Output +model.modelfile = pspm_cfg_selector_outputfile('run', job); + +%% TAM settings +model.modality = 'pupil'; +model.modelspec = job.modelspec; +model.window = job.window; +model.norm = job.norm; +model.baseline = job.baseline; +model.norm_max = job.norm_max; + +%% Channel +model.channel = pspm_cfg_selector_channel('run', job.chan); + +%% Filter +model.filter = pspm_cfg_selector_filter('run', job.filter); + +if ischar(model.filter) && strcmpi(model.filter, 'none') + model = rmfield(model, 'filter'); +else + if isnumeric(model.filter.lpfreq) && isnan(model.filter.lpfreq) + model.filter.lpfreq = 'none'; + end + + if isnumeric(model.filter.hpfreq) && isnan(model.filter.hpfreq) + model.filter.hpfreq = 'none'; + end + + if isnumeric(model.filter.down) && isnan(model.filter.down) + model.filter.down = 'none'; + end +end + +%% Standard experimental condition +if isfield(job.std_exp_cond, 'none') + model.std_exp_cond = 'none'; +elseif isfield(job.std_exp_cond, 'name') + model.std_exp_cond = job.std_exp_cond.name; +elseif isfield(job.std_exp_cond, 'index') + model.std_exp_cond = job.std_exp_cond.index; +end + +%% Marker channel +if isfield(newoptions, 'marker_chan_num') + options.marker_chan = newoptions.marker_chan_num; +end + +%% Options +options = pspm_update_struct(options, job.output, {'overwrite'}); + +%% Run +[sts, ~] = pspm_tam(model, options); + +if sts < 1 + error('TAM estimation failed.'); +end + +out = {model.modelfile}; + +end \ No newline at end of file diff --git a/src/pspm_cfg/pspm_cfg_selector_data_design.m b/src/pspm_cfg/pspm_cfg_selector_data_design.m index 585b792fd..72ec19e63 100644 --- a/src/pspm_cfg/pspm_cfg_selector_data_design.m +++ b/src/pspm_cfg/pspm_cfg_selector_data_design.m @@ -17,7 +17,10 @@ % datafile model.datafile{iSession,1} = job.session(iSession).datafile{1}; % missing epochs - model.missing{1,iSession} = pspm_cfg_selector_missing_epochs('run', job.session(iSession)); + % model.missing{1,iSession} = pspm_cfg_selector_missing_epochs('run', job.session(iSession)); + if isfield(job.session(iSession), 'missing') + model.missing{1,iSession} = pspm_cfg_selector_missing_epochs('run', job.session(iSession)); + end % data & design if isfield(job.session(iSession).data_design,'no_condition') model.timing = {}; @@ -120,7 +123,7 @@ else condfile.help = vertcat(helptext{[1, 3]}); end - case 'extract' + case {'extract', 'tam'} condfile.help = helptext{1}; end @@ -217,7 +220,7 @@ else condition.val = {condname, onsets, pmod_rep}; end - case 'extract' + case {'extract', 'tam'} condition.val = {condname, onsets}; end condition.help = {''}; @@ -289,6 +292,8 @@ timing.values = {condfile, condition_rep, marker_cond ,no_condition}; case 'extract' timing.values = {condfile, condition_rep, marker_cond}; + case 'tam' + timing.values = {condfile, condition_rep}; end timing.help = {['Specify the timing of the events within the design matrix. Timing can '... @@ -319,6 +324,8 @@ session.val = {datafile, missing, timing, nuisancefile}; case 'extract' session.val = {datafile, missing, timing}; + case 'tam' + session.val = {datafile, timing}; end session.help = {''}; diff --git a/src/pspm_cfg/pspm_cfg_selector_filter.m b/src/pspm_cfg/pspm_cfg_selector_filter.m index c64ac4936..3211428f2 100644 --- a/src/pspm_cfg/pspm_cfg_selector_filter.m +++ b/src/pspm_cfg/pspm_cfg_selector_filter.m @@ -55,7 +55,7 @@ lpfreq.name = 'Cutoff frequency'; lpfreq.tag = 'freq'; lpfreq.strtype = 'r'; -if isfield(default,'lpfreq') +if isfield(default,'lpfreq') && isnumeric(default.lpfreq) lpfreq.val = {default.lpfreq}; end lpfreq.num = [1 1]; @@ -80,8 +80,13 @@ lowpass = cfg_choice; lowpass.name = 'Low-pass filter'; lowpass.tag = 'lowpass'; -lowpass.val = {enable_lp}; +% lowpass.val = {enable_lp}; lowpass.values = {enable_lp, disable}; +if isfield(default,'lpfreq') && ischar(default.lpfreq) && strcmpi(default.lpfreq, 'none') + lowpass.val = {disable}; +else + lowpass.val = {enable_lp}; +end lowpass.help = {''}; % High pass diff --git a/src/pspm_cfg/pspm_cfg_tam.m b/src/pspm_cfg/pspm_cfg_tam.m new file mode 100644 index 000000000..ba2cb2c11 --- /dev/null +++ b/src/pspm_cfg/pspm_cfg_tam.m @@ -0,0 +1,95 @@ +function tam = pspm_cfg_tam + +global settings +if isempty(settings), pspm_init; end + +%% Standard items +output = pspm_cfg_selector_outputfile('Model'); +session_rep = pspm_cfg_selector_data_design('tam'); +timeunits = pspm_cfg_selector_timeunits; +chan = pspm_cfg_selector_channel('pupil'); +normalise = pspm_cfg_selector_norm; + +filter_default = settings.tam(1).filter; + +if ischar(filter_default.lpfreq) && strcmpi(filter_default.lpfreq, 'none') + filter_default.lpfreq = NaN; +end + +if ischar(filter_default.hpfreq) && strcmpi(filter_default.hpfreq, 'none') + filter_default.hpfreq = NaN; +end + +if ischar(filter_default.down) && strcmpi(filter_default.down, 'none') + filter_default.down = NaN; +end + +filter = pspm_cfg_selector_filter(filter_default); + +%% Model specification +modelspec = cfg_menu; +modelspec.name = 'Model'; +modelspec.tag = 'modelspec'; +modelspec.labels = {'Dilation', 'Constriction'}; +modelspec.values = {'dilation', 'constriction'}; +modelspec.val = {'dilation'}; +modelspec.help = pspm_cfg_help_format('pspm_tam', 'model.modelspec'); + +%% Window +window = cfg_entry; +window.name = 'Window'; +window.tag = 'window'; +window.strtype = 'r'; +window.num = [1 1]; +window.help = pspm_cfg_help_format('pspm_tam', 'model.window'); + +%% Baseline +baseline = cfg_entry; +baseline.name = 'Baseline'; +baseline.tag = 'baseline'; +baseline.strtype = 'r'; +baseline.num = [1 1]; +baseline.val = {0}; +baseline.help = pspm_cfg_help_format('pspm_tam', 'model.baseline'); + +%% Normalize maximum +norm_max = cfg_menu; +norm_max.name = 'Normalize maximum'; +norm_max.tag = 'norm_max'; +norm_max.labels = {'No', 'Yes'}; +norm_max.values = {0, 1}; +norm_max.val = {0}; +norm_max.help = pspm_cfg_help_format('pspm_tam', 'model.norm_max'); + +%% Standard experimental condition +std_none = cfg_const; +std_none.name = 'None'; +std_none.tag = 'none'; +std_none.val = {'none'}; + +std_name = cfg_entry; +std_name.name = 'Condition name'; +std_name.tag = 'name'; +std_name.strtype = 's'; + +std_index = cfg_entry; +std_index.name = 'Condition index'; +std_index.tag = 'index'; +std_index.strtype = 'i'; +std_index.num = [1 1]; + +std_exp_cond = cfg_choice; +std_exp_cond.name = 'Standard experimental condition'; +std_exp_cond.tag = 'std_exp_cond'; +std_exp_cond.val = {std_none}; +std_exp_cond.values = {std_none, std_name, std_index}; +std_exp_cond.help = pspm_cfg_help_format('pspm_tam', 'model.std_exp_cond'); + +%% Executable branch +tam = cfg_exbranch; +tam.name = 'Trial Average Model'; +tam.tag = 'tam'; +tam.val = {output, chan, timeunits, session_rep, modelspec, window, normalise, filter, baseline, norm_max, std_exp_cond}; +tam.prog = @pspm_cfg_run_tam; +tam.vout = @pspm_cfg_vout_outfile; +tam.help = pspm_cfg_help_format('pspm_tam'); \ No newline at end of file diff --git a/src/pspm_check_model.m b/src/pspm_check_model.m index 320d4c2bc..6e165418b 100644 --- a/src/pspm_check_model.m +++ b/src/pspm_check_model.m @@ -430,7 +430,7 @@ model.baseline = 0; elseif ~isnumeric(model.baseline) warning('ID:invalid_input','''model.baseline'' has to be a numeric.'); return; - elseif model.baseline > model.window || model.baseline < 0 + elseif model.baseline >= model.window || model.baseline < 0 warning('ID:invalid_input',['''model.baseline'' has to be positive ',... 'and smaller than ''model.window''.']); return; end @@ -441,22 +441,20 @@ ' index corresponding to it.']; if ~isfield(model,'std_exp_cond') model.std_exp_cond = 'none'; - elseif ~ischar(model.std_exp_cond) && ~isnumeric(model.std_exp_cond) - warning('ID:invalid_input',std_cond_war_msg); return; elseif ischar(model.std_exp_cond) - tmp_ind = cellfun(@(x) strcmpi(model.std_exp_cond,x), model.timing{1}.names); - if ~any(tmp_ind) - warning('ID:invalid_input',std_cond_war_msg); return; - end - std_exp_cond.name = model.std_exp_cond; - std_exp_cond.ind = find(tmp_ind); + % 'none' is explicitly allowed + if ~strcmpi(model.std_exp_cond, 'none') + tmp_ind = strcmpi(model.std_exp_cond, model.timing{1}.names); + if ~any(tmp_ind); warning('ID:invalid_input', std_cond_war_msg); return; end + end elseif isnumeric(model.std_exp_cond) - if model.std_exp_cond < 1 || ... - model.std_exp_cond > numel(model.timing{1}.names) - warning('ID:invalid_input',std_cond_war_msg); return; - end - std_exp_cond.name = model.timing{1}.names(model.std_exp_cond); - std_exp_cond.ind = model.std_exp_cond; + % Must be one valid scalar integer index + if ~isscalar(model.std_exp_cond) || ~isfinite(model.std_exp_cond) || ... + model.std_exp_cond < 1 || model.std_exp_cond > numel(model.timing{1}.names) + warning('ID:invalid_input', std_cond_war_msg); return; + end + else + warning('ID:invalid_input',std_cond_war_msg); return; end clear std_cond_war_msg tmp_ind @@ -519,8 +517,15 @@ end end -if ~isfield(model.filter, 'down') || ~isnumeric(model.filter.down) - warning('ID:invalid_input', 'Filter structure needs a numeric ''down'' field.'); return; +if ~isfield(model.filter, 'down') + warning('ID:invalid_input', 'Filter structure needs a ''down'' field.'); return; +end + +valid_down = (isnumeric(model.filter.down) && isscalar(model.filter.down) && isfinite(model.filter.down) && model.filter.down > 0) ... + || (ischar(model.filter.down) && strcmpi(model.filter.down, 'none')); + +if ~valid_down + warning('ID:invalid_input', 'Filter structure needs a positive numeric ''down'' field or ''none''.'); return; end diff --git a/src/pspm_find_valid_fixations.m b/src/pspm_find_valid_fixations.m index 5e1ef793b..a2a60c242 100644 --- a/src/pspm_find_valid_fixations.m +++ b/src/pspm_find_valid_fixations.m @@ -20,17 +20,18 @@ % the circle around fixation point. Since this usage is currently % considered secondary, it still requires a valid pupil channel as % primary channel, even though unrelated to pupil analysis. -% In both usages, valid fixations can be outputted as additional channel. +% In both usages, an additional channel marking invalid fixations can be added. % By default, screen centre is assumed as fixation point. If an explicit % fixation point is given, the function assumes that the screen is % perpendicular to the vector from the eye to the fixation point (which % is approximately correct for large enough screen distance). % ● Format -% [sts, channel_index, fn] = pspm_find_valid_fixations(fn, bitmap, options) -% [sts, channel_index, fn] = pspm_find_valid_fixations(fn, circle_degree, distance, unit, options) +% [sts, pos_of_channel, fn] = pspm_find_valid_fixations(fn, bitmap, options) +% [sts, pos_of_channel, fn] = pspm_find_valid_fixations(fn, circle_degree, distance, unit, options) % ● Arguments -% * fn : The actual data file containing the eyelink recording with gaze -% data converted to cm. +% * fn : Data file containing the eye-tracking recording. +% The corresponding gaze channels must be available in +% distance units for fixation mode, or in distance or pixel units for bitmap mode. % * bitmap : A nxm matrix of the same size as the display, with 1 % for valid and 0 for invalid gaze points. IMPORTANT: the bitmap has to % be defined in terms of the eyetracker coordinate system, i.e. @@ -41,12 +42,13 @@ % * distance : Distance between eye and screen in length units. % * unit : Unit in which distance is given. % ┌────────options -% ├.fixation_point : A nx2 vector containing x and y of the fixation point (with respect +% ├.fixation_point : An nx2 matrix containing x and y of the fixation point (with respect % │ to the given resolution, and in the eyetracker coordinate system). % │ n should equal either 1 (constant fixation point) or the length -% │ of the actual data. If resolution is not defined the values are -% │ given in percent. Therefore (0.5 0.5) would correspond to the -% │ middle of the screen. Default is (0.5 0.5). Only taken into account +% │ of the actual data. If resolution is not defined, the +% │ coordinates are given as fractions of the screen dimensions. +% │ Therefore (0.5 0.5) would correspond to the middle of the screen. +% │ Default is (0.5 0.5). Only taken into account % │ if there is no bitmap. % ├────.resolution : Resolution with which the fixation point is defined (Maximum value % │ of the x and y coordinates). This can be the screen resolution in @@ -54,16 +56,18 @@ % │ in cm (e.g. (50 30)). Default is (1 1). Only taken into account % │ if there is no bitmap. % ├────.screen_dim : Only considered if .plot_gaze_coords is passed; used -% │ plot the gaze data and circle on the actual screen +% │ to plot the gaze data and circle on the actual screen % │ dimensions rather than using auto scaling. Input % │ should follow the format: [x_dim, y_dim]. % ├.plot_gaze_coords: Define whether to plot the gaze coordinates for visual % │ inspection of the validation process. Default is false. -% ├.channel_action: Define whether to add or replace the data. Default is -% │ 'add'. Possible values are 'add' or 'replace' +% ├.channel_action: Defines whether the processed channel is added as a new +% │ channel or replaces the selected input channel. Default +% │ is 'add'. Possible values are 'add' or 'replace'. % ├───.add_invalid: [0/1] If this option is enabled, an extra channel will be -% │ written containing information about the valid samples. -% │ Data points equal to 1 correspond to invalid fixation. +% │ added that marks invalid fixation samples. +% │ Data points equal to 1 correspond to invalid fixations, +% │ and data points equal to 0 correspond to valid fixations. % │ Default is not to add this channel. % └───────.channel: Choose channels in which the data should be set to NaN % during invalid fixations. This can be a channel @@ -89,7 +93,7 @@ % ● Developer % Additional i/o options for recursive calls are not included in the help. % (1) fn can be a data structure as permitted by pspm_load_data, -% (2) the output argument pos_of_channels is an index of the channel(s) +% (2) the output argument pos_of_channel is an index of the channel(s) % that was/were replaced or added % (3) The third output argument is required for recursive calls % ● History @@ -110,25 +114,27 @@ ' You have to either pass a bitmap or circle_degree, distance and unit',... ' to compute the valid fixations']); return; end -if numel(varargin{1}) > 1 +if numel(varargin) <= 2 && numel(varargin{1}) > 1 mode = 'bitmap'; bitmap = varargin{1}; - if ~ismatrix(bitmap) || (~isnumeric(bitmap) && ~islogical(bitmap)) - warning('ID:invalid_input', ['The bitmap must be a matrix and must',... - ' contain numeric or logical values.']); return; + if ~ismatrix(bitmap) || (~isnumeric(bitmap) && ~islogical(bitmap)) || any(~ismember(bitmap(:), [0 1])) + warning('ID:invalid_input', ['The bitmap must be a numeric or logical matrix containing only 0 and 1.']); return; end if numel(varargin) < 2 options = struct(); - options.mode = 'bitmap'; else options = varargin{2}; - options.mode = 'bitmap'; + if ~isstruct(options) + warning('ID:invalid_input', 'Options must be a struct.'); + return; + end end + options.mode = 'bitmap'; else mode = 'fixation'; if numel(varargin) < 3 warning('ID:invalid_input', ['Not enough input arguments.', ... - ' You have to set circle_degree, distance and unit',... + ' You have to provide circle_degree, distance and unit',... ' to compute the valid fixations']); return; end circle_degree = varargin{1}; @@ -146,11 +152,11 @@ options.mode = 'fixation'; end end - if ~isnumeric(circle_degree) - warning('ID:invalid_input', 'Circle_degree is not numeric.'); + if ~(isnumeric(circle_degree) && isscalar(circle_degree) && isreal(circle_degree) && isfinite(circle_degree) && circle_degree >= 0) + warning('ID:invalid_input', 'Circle_degree must be a finite, non-negative numeric scalar.'); return; - elseif ~isnumeric(distance) - warning('ID:invalid_input', 'Distance is not set or not numeric.'); + elseif ~(isnumeric(distance) && isscalar(distance) && isreal(distance) && isfinite(distance) && distance > 0) + warning('ID:invalid_input', 'Distance must be a finite, positive numeric scalar.'); return; elseif ~ischar(unit) warning('ID:invalid_input', 'Unit should be a char.'); @@ -164,6 +170,7 @@ [nsts,distance] = pspm_convert_unit(distance,unit ,'mm'); if nsts~=1 warning('ID:invalid_input', 'Failed to convert distance to mm.'); + return; end end end @@ -194,15 +201,25 @@ elseif strcmpi(mode, 'bitmap') channels_correct_units = find(~contains(channelunits_list, 'degree')); end + gazedata = struct('infos', alldata.infos, 'data', {alldata.data(channels_correct_units)}); [sts_gaze, gaze_x, gaze_y, eye] = pspm_load_gaze(gazedata, data.header.chantype); if sts_gaze < 1 warning('ID:invalid_input', ['Unable to perform gaze ', ... - 'validation. Cannot find gaze channels with distance ',... - 'unit values. Maybe you need to convert them with ', ... - 'pspm_convert_gaze()']); + 'validation. Cannot find corresponding gaze channels ', ... + 'with compatible units. Maybe you need to convert them ', ... + 'with pspm_convert_gaze().']); + return; + end + + if numel(gaze_x.data) ~= numel(gaze_y.data) || numel(gaze_x.data) ~= numel(data.data) + warning('ID:invalid_input', ['Pupil and corresponding gaze channels must have ', 'the same number of data points.']); + return; + end + if gaze_x.header.sr ~= gaze_y.header.sr || gaze_x.header.sr ~= data.header.sr + warning('ID:invalid_input', ['Pupil and corresponding gaze channels must have ', 'the same sampling rate.']); return; end @@ -213,13 +230,17 @@ case 'fixation' % expand fixation point to size of data fix_point = options.fixation_point; - if size(fix_point, 1) == 1 + if size(fix_point, 2) ~= 2 + warning('ID:invalid_input', ... + 'Fixation point must have exactly two columns for x and y coordinates.'); + return; + elseif size(fix_point, 1) == 1 fix_point = repmat(fix_point(:)', numel(gaze_x.data), 1); - elseif size(fix_point, 1) ~= numel(gaze_x) + elseif size(fix_point, 1) ~= numel(gaze_x.data) warning('ID:invalid_input', ['Fixation point has wrong ', ... 'dimensions - it should be 1x2 or nx2 where n is the ', ... 'number of gaze data points.']); - return + return; end % normalise fixation point to fraction of full screen @@ -398,16 +419,17 @@ if (rsts(1) < 1 && rsts(2) < 1) return; elseif (rsts(1) < 1 || rsts(2) < 1) + successful_idx = find(rsts > 0); pos_of_channel(rsts < 1) = []; + + [~, best_eye] = pspm_find_eye(channels{successful_idx}); + alldata.infos.source.best_eye = best_eye; else - % update best eye - eye_stat = Inf(1,numel(alldata.infos.source.eyesObserved)); - for i = 1:numel(alldata.infos.source.eyesObserved) - e_stat = alldata.infos.source.chan_stats(pos_of_channel); - eye_stat(i) = max(cellfun(@(x) x.nan_ratio, e_stat)); - end + eye_stat = cellfun(@(x) x.nan_ratio, alldata.infos.source.chan_stats(pos_of_channel)); + [~, min_idx] = min(eye_stat); - alldata.infos.source.best_eye = lower(alldata.infos.source.eyesObserved(min_idx)); + [~, best_eye] = pspm_find_eye(channels{min_idx}); + alldata.infos.source.best_eye = best_eye; end end diff --git a/src/pspm_glm.m b/src/pspm_glm.m index a94265449..4f5eb401c 100644 --- a/src/pspm_glm.m +++ b/src/pspm_glm.m @@ -1,4 +1,4 @@ - function [sts, glm] = pspm_glm(model, options) +function [sts, glm] = pspm_glm(model, options) % ● Description % pspm_glm specifies a within-subject general linear convolution model % (GLM) of predicted signals and calculates amplitude estimates for these diff --git a/src/pspm_init.m b/src/pspm_init.m index 26aff0d8b..aa50b2a2a 100644 --- a/src/pspm_init.m +++ b/src/pspm_init.m @@ -680,7 +680,7 @@ defaults.glm(end + 1) = struct(... 'modality', 'hp', ... 'modelspec', 'hp_rew', ... - 'cbf', struct('fhandle', @bf_hprf_rew, 'args', []), ... + 'cbf', struct('fhandle', @pspm_bf_hprf_rew, 'args', []), ... 'filter', struct('lpfreq', 0.5, 'lporder', 4, 'hpfreq', 0.015, 'hporder', 4, 'down', 10, 'direction', 'bi'), ... 'default', 0); % GLM for PS (fear-conditioning) @@ -756,13 +756,15 @@ 'cif', struct('fhandle', @pspm_bf_ldrf_gm, 'args', [0, 2.76 , 0.09 , 0.31],... % input function & default parameters 'lb', [0,0,0,0], 'ub', [0,Inf,Inf,Inf]),... % & the lower/upper bounds 'filter', struct('lpfreq', 'none', 'lporder', 0, ... % default filter - 'hpfreq', 'none', 'hporder', 0, 'down', 0, 'direction', 'bi')); + 'hpfreq', 'none', 'hporder', 0, 'down', 'none', 'direction', 'bi')); defaults.tam(2) = struct(... 'modality', 'pupil', ... 'modelspec', 'constriction',... 'cbf', struct('fhandle', @pspm_bf_lcrf_gm, 'args', [0.2, 3.24 , 0.18 , 0.43]),... - 'cif', struct('fhandle', @pspm_bf_lcrf_gm, 'args', [0, 2.76 , 0.09 , 0.31], 'lb', [0,0,0,0], 'ub', [0,Inf,Inf,Inf]),... - 'filter', struct('lpfreq', 'none', 'lporder', 0, 'hpfreq', 'none', 'hporder', 0, 'down', 0, 'direction', 'bi')); + 'cif', struct('fhandle', @pspm_bf_lcrf_gm, 'args', [0, 2.76 , 0.09 , 0.31], ... + 'lb', [0,0,0,0], 'ub', [0,Inf,Inf,Inf]),... + 'filter', struct('lpfreq', 'none', 'lporder', 0, ... + 'hpfreq', 'none', 'hporder', 0, 'down', 'none', 'direction', 'bi')); %% 9 FIRST LEVEL settings % 9.1 allowed first level model types defaults.first = {'glm', 'sf', 'dcm', 'tam'}; diff --git a/src/pspm_tam.m b/src/pspm_tam.m index f5008a211..4b2d6dd3e 100644 --- a/src/pspm_tam.m +++ b/src/pspm_tam.m @@ -1,4 +1,4 @@ -function tam = pspm_tam(model, options) +function [sts, tam] = pspm_tam(model, options) % ● Description % TAM stands for Trial Average Model and allows to fit models on trial-averaged data. % pspm_tam starts by extracting and averaging signal segments of length `model.window` @@ -12,7 +12,7 @@ % ├─────.timing: a multiple condition file name (single session) OR % │ a cell array of multiple condition file names OR % │ a struct (single session) with fields .names, .onsets, -% │ and (optional) .durations OR +% │ and (optional) duration OR % │ a cell array of struct OR % │ a struct with fields 'markerinfos', 'markervalues', % │ 'names' OR @@ -40,7 +40,7 @@ % ├───.baseline: [optional] allows to specify a baseline in 'seconds' which is % │ applied to the data before fitting the model. It has to % │ be positive and smaller than model.window. If no baseline -% │ specified, data will be baselined wrt. the first datapoint. +% │ specified, data will be baselined wrt. the first data-point. % │ DEFAULT: 0 % ├.std_exp_cond: [optional] allows to specify the standard experimental condition % │ as a string or an index in timing.names. @@ -83,6 +83,7 @@ % Introduced In PsPM 4.2 % Written in 2020 by Ivan Rojkov (University of Zurich) % Maintained in 2022 by Teddy +% Maintained in 2026 by Bernhard A. von Raußendorf % ● Developer % The fitting process is a residual least square minimisation where the % predicted value is calculated as following: @@ -114,6 +115,16 @@ pspm_init; end tam = struct(); +% %% tmp filter check for testing +% +% if isfield(model, 'filter') +% disp('Filter passed from batch:') +% disp(model.filter) +% else +% disp('Using default TAM filter:') +% global settings +% disp(settings.tam(strcmpi({settings.tam.modelspec}, model.modelspec)).filter) +% end %% 2 Check input % 2.1 check missing input -- @@ -126,12 +137,33 @@ return end +%% Load timing definitions +timing = model.timing; + +for iFile = 1:numel(timing) + if ischar(timing{iFile}) + timing{iFile} = load(timing{iFile}); + end +end + +%% Standard experimental condition +std_exp_cond = []; + +if ~(ischar(model.std_exp_cond) && strcmpi(model.std_exp_cond, 'none')) + if ischar(model.std_exp_cond) + std_exp_cond.ind = find(strcmpi( model.std_exp_cond, timing{1}.names), 1); + else + std_exp_cond.ind = model.std_exp_cond; + end + std_exp_cond.name = timing{1}.names{std_exp_cond.ind}; +end + %% Loading files fprintf('Computing Trial Average Model: %s \n', model.modelfile); -n_exp_cond = numel(model.timing{1}.names); % number of experimental conditions -n_file = numel(model.datafile); % number of files +n_exp_cond = numel(timing{1}.names); % number of experimental conditions +n_file = numel(model.datafile); % number of files % Loading data and sr fprintf('Getting data .'); @@ -142,17 +174,6 @@ % Filling up the data and the sampling rates y{iFile} = data.data(:); sr(iFile) = data.header.sr; - fprintf('.'); - - % If the timeunits is markers - if strcmpi(model.timeunits, 'markers') - [sts, data] = pspm_load_channel(model.datafile{iFile}, options.marker_chan{iFile}, 'marker'); - if sts < 1 - warning('ID:invalid_input','Could not load the specified marker channel.'); - return; - end - markers{iFile} = data.data; - end fprintf('.'); end @@ -161,25 +182,26 @@ oldsr = sr; % Checking if the sampling rate is the same for all samples. -if n_file > 1 && any(diff(sr) > 0) - if model.filter.down > min(sr)) ||... % if filter.down is less than the minimal sr - strcmpi(model.filter.down,'none') % if filter.down is none - model.filter.down = min(sr); - fprintf('\nSampling rate differs between sessions. Data will be downsampled.\n') +if n_file > 1 && any(diff(sr) ~= 0) % not any(diff(sr) > 0) + if (ischar(model.filter.down) && strcmpi(model.filter.down, 'none')) || ... % if filter.down is none + (isnumeric(model.filter.down) && model.filter.down > min(sr)) % if filter.down is less than the minimal sr + + model.filter.down = min(sr); + fprintf('\nSampling rate differs between sessions. Data will be downsampled.\n') % when and where? end else fprintf('\n'); end -%% Zscoring the data -if model.norm - fprintf('Zscoring ...\n') - n_file = numel(model.datafile); - for iFile = 1:n_file - % NANZSCORE found in src/VBA/stats&plots - [y{iFile},~,~] = nannorm(y{iFile}); - end -end +%% Zscoring the data -> will be done by pspm_extract_segments +% if model.norm +% fprintf('Zscoring ...\n') +% n_file = numel(model.datafile); +% for iFile = 1:n_file +% % NANZSCORE found in src/ext/VBA/stats&plots +% [y{iFile},~,~] = nanzscore(y{iFile}); %nannorm??? -> nanzscore?? +% end +% end %% Extracting segments fprintf('Extracting segments ...\n') @@ -188,20 +210,36 @@ extrsgopt.timeunits = model.timeunits; extrsgopt.length = model.window; % segments of 'model.window' time unit long extrsgopt.plot = 0; % do not plot mean value and std +extrsgopt.norm = model.norm; for k=1:n_file if strcmpi(model.timeunits, 'markers') - extrsgopt.marker_chan = markers(k); + % marker channel for this session + if iscell(options.marker_chan) + marker_chan = options.marker_chan{k}; + else + marker_chan = options.marker_chan; + end + + % pspm_extract_segments uses a different option name + fileopt = extrsgopt; + fileopt.marker_chan_num = marker_chan; + + % In file mode the raw pupil data are loaded again, + % so normalization must happen inside extract_segments. + [lsts, s] = pspm_extract_segments( 'file', model.datafile{k}, model.channel, timing{k}, fileopt); + else + [lsts, s] = pspm_extract_segments('data', y{k}, sr(k), timing{k}, extrsgopt); % wird es richtig gen znormed? + end - [lsts, s] = pspm_extract_segments('manual', y(k), sr(k), model.timing(k), extrsgopt); if lsts<1, warning('ID:error_extract_segments','An error occured in pspm_extract_segments.'); return; end for i=1:n_exp_cond - tmp_data.mean = s.segments{i,1}.mean; - tmp_data.std = s.segments{i,1}.std; - tmp_data.sem = s.segments{i,1}.sem; - tmp_data.t = s.segments{i,1}.t; + tmp_data.mean = s.segments{i}.mean; % why not eg 50x1 + tmp_data.std = s.segments{i}.std; % why not eg 50x1 + tmp_data.sem = s.segments{i}.sem; % why not eg 50x1 + tmp_data.t = s.segments{i}.t; % % a cell array of struct and of size (n_file x n_exp_cond) where each % line correspond to a given file and each column to an % experimental condition @@ -214,47 +252,59 @@ %% Downsample the data % if a filter was specified or if the data differ in sr -fprintf('Filtering ...\n') -for i = 1:n_exp_cond - for k = 1:n_file - model.filter.sr = sr(k); +fprintf('Filtering ...\n') % maybe only if there is a filtering? - [lsts, segm{k,i}, ~] = structfun(@(x) pspm_prepdata(x, model.filter),segm{k,i},'UniformOutput',false); - if any(structfun(@(x) x<1,lsts)), warning('ID:error_prepdata','An error occured in pspm_prepdata.'); return; end - clear new_sr lsts +fields = {'mean', 'std', 'sem'}; +for i = 1:n_exp_cond + for k = 1:n_file + model.filter.sr = sr(k); % adds the corresponding sr to the filer + + for f = 1:numel(fields) + field = fields{f}; + [lsts, segm{k,i}.(field), new_sr(k)] = pspm_prepdata(segm{k,i}.(field), model.filter); + if lsts < 1; warning('ID:error_prepdata', 'An error occured in pspm_prepdata.'); return; end + end + [llsts, segm{k,i}.t , new_sr(k) ] = pspm_downsample(segm{k,i}.t, sr(k), new_sr(k) ); % does not need to be filtered! + if llsts < 1; warning('ID:error_downsample', 'An error occured in pspm_downsample.'); return; end + end end -% changing the sampling rate -sr = model.filter.down*ones(size(sr)); +sr = new_sr; % maybe test upstream if all new_sr are the same! +% % changing the sampling rate +% if isnumeric(model.filter.down) +% sr = model.filter.down * ones(size(new_sr)); % needs to be tested why not the values of prepdata? +% end +clear new_sr %% Determining mean values fprintf('Preparing for fitting ...\n') baseline_index = floor(sr(1)*model.baseline)+1; -if exist('std_exp_cond','var') + +if ~isempty(std_exp_cond) tmp_data = [segm{:,std_exp_cond.ind}]; - std_exp_cond.data = nanmean([tmp_data.mean],2); - std_exp_cond.std = nanmean([tmp_data.std],2); - std_exp_cond.sem = nanmean([tmp_data.sem],2); + std_exp_cond.data = mean([tmp_data.mean],2, 'omitnan'); + std_exp_cond.std = mean([tmp_data.std],2, 'omitnan'); + std_exp_cond.sem = mean([tmp_data.sem],2, 'omitnan'); end for i=1:n_exp_cond tmp_data = [segm{:,i}]; - tmp_data_new.data = nanmean([tmp_data.mean],2); - tmp_data_new.std = nanmean([tmp_data.std],2); - tmp_data_new.sem = nanmean([tmp_data.sem],2); - tmp_data_new.t = nanmean([tmp_data.t],2); + tmp_data_new.data = mean([tmp_data.mean],2, 'omitnan'); + tmp_data_new.std = mean([tmp_data.std],2, 'omitnan'); + tmp_data_new.sem = mean([tmp_data.sem],2, 'omitnan'); + tmp_data_new.t = mean([tmp_data.t],2, 'omitnan'); % Subtracting the standard experimental condition - if exist('std_exp_cond','var') && i~=std_exp_cond.ind + if ~isempty(std_exp_cond) && i~=std_exp_cond.ind tmp_data_new.data = tmp_data_new.data - std_exp_cond.data; tmp_data_new.std = tmp_data_new.std + std_exp_cond.std; % the error adds up tmp_data_new.sem = tmp_data_new.sem + std_exp_cond.sem; % the error adds up @@ -264,14 +314,15 @@ tmp_data_new.data = tmp_data_new.data - tmp_data_new.data(baseline_index); % Dividing by the max value - if model.norm_max - [tmp_max,tmp_max_ind] = max(tmp_data_new.data); + if model.norm_max + [tmp_max,tmp_max_ind] = max(tmp_data_new.data); % this does not give us the first peak!!! + % maybe add if tmp_max == 0 tmp_data_new.data = tmp_data_new.data/tmp_max; - tmp_data_new.std = tmp_data_new.std + tmp_data_new.std(tmp_max_ind); % the error adds up - tmp_data_new.sem = tmp_data_new.sem + tmp_data_new.sem(tmp_max_ind); % the error adds up + tmp_data_new.std = tmp_data_new.std + tmp_data_new.std(tmp_max_ind); % the error adds up . why not /abs(tmp_max) + tmp_data_new.sem = tmp_data_new.sem + tmp_data_new.sem(tmp_max_ind); % the error adds up . why not /abs(tmp_max) end - mean{1,i} = tmp_data_new; + mean_data{1,i} = tmp_data_new; clear tmp_data tmp_data_new tmp_max tmp_max_ind end @@ -280,7 +331,7 @@ fprintf('Fitting ...\n') for i=1:n_exp_cond - raw_y = mean{1,i}.data; + raw_y = mean_data{1,i}.data; n = model.window; td = n / length(raw_y); @@ -299,7 +350,7 @@ % Minimization of RSS warning off all [~, fitted{1,i}.optargs, fitted{1,i}.fval, sts, fmincon_output] = ... - evalc('fmincon(RSS,model.if.args,[],[],[],[],model.if.lb,model.if.ub)'); + evalc('fmincon(RSS,model.if.args,[],[],[],[],model.if.lb,model.if.ub)'); % lb ub for opt. warning on all if sts == 0 warning('ID:fmincon',['During the fitting process, ''fmincon'' exceeded', ... @@ -320,7 +371,7 @@ % Calculating the predicted signal that will be included in the output structure fitted{1,i}.data = predicted_y(fitted{1,i}.optargs); % Cutting away tail - tmp_y = mean{1,i}.data; + tmp_y = mean_data{1,i}.data; fitted{1,i}.data(size(tmp_y,1)+1:end) = []; end @@ -336,9 +387,15 @@ tam.bf = model.bf; tam.if = model.if; +% Check if filtered +no_lp = ischar(model.filter.lpfreq) && strcmpi(model.filter.lpfreq, 'none'); +no_hp = ischar(model.filter.hpfreq) && strcmpi(model.filter.hpfreq, 'none'); +no_down = ischar(model.filter.down) && strcmpi(model.filter.down, 'none'); +filtered = ~(no_lp && no_hp && no_down); + % Collecting fitting data -tmp_mean = [mean{1,:}]; -tam.data.Y = {tmp_mean.data}; +tmp_mean = [mean_data{1,:}]; +tam.data.Y = {tmp_mean.data}; % measured data tam.data.X = {tmp_mean.t}; tam.data.std = {tmp_mean.std}; tam.data.sem = {tmp_mean.sem}; @@ -347,7 +404,7 @@ tam.data.normd = model.norm; tam.data.norm = model.norm; -if exist('std_exp_cond','var') +if ~isempty(std_exp_cond) tam.data.std_exp_cond.name = std_exp_cond.name; tam.data.std_exp_cond.ind = std_exp_cond.ind; else @@ -358,19 +415,19 @@ tmp_fitted = [fitted{1,:}]; tam.fit.Y = {tmp_fitted.data}; tam.fit.X = {tmp_mean.t}; -tam.fit.rss = {tmp_fitted.fval}; % RSS (residual sum square) +tam.fit.rss = {tmp_fitted.fval}; tam.fit.args = {tmp_fitted.optargs}; tam.fit.sr = num2cell(sr(:).'); tam.infos.duration = model.window; -tam.infos.durationinfo = 'duration in seconds'; +tam.infos.durationinfo = 'duration in seconds'; % not always true! -> timeunits='samples' tam.timing = model.timing; tam.modeltype = 'tam'; tam.modality = model.modality; -tam.names = model.timing{1}.names(:).'; +tam.names = timing{1}.names(:).'; % Saving structure savedata = struct('tam', tam); diff --git a/test/pspm_convert_hb2hp_test.m b/test/pspm_convert_hb2hp_test.m index 5585acebb..1df5ea2d2 100644 --- a/test/pspm_convert_hb2hp_test.m +++ b/test/pspm_convert_hb2hp_test.m @@ -75,6 +75,8 @@ function basic_conversion(this) [sts, infos, data, filestruct] = pspm_load_data(fn); this.verifyEqual(sts, 1); this.verifyEqual(data{outchannel}.header.chantype, 'hp'); + this.verifyEqual(data{outchannel}.header.units, 'ms'); + this.verifyEqual(data{outchannel}.header.sr, sr); end function too_strict_limits(this) @@ -82,11 +84,6 @@ function too_strict_limits(this) sr = 1; options = struct('limit_lower', 11, 'limit_upper', 11); - % options = struct('limit_lower', 10, 'limit_upper', 11); - % this.verifyWarning(@() pspm_convert_hb2hp(fn, sr, options), ... - % 'ID:too_strict_limits'); - - % this.verifyWarningFree(@() pspm_convert_hb2hp(this.input_filename, sr)); [sts, outchannel] = pspm_convert_hb2hp(fn, sr,options); this.verifyEqual(sts, 1); [sts, infos, data, filestruct] = pspm_load_data(fn); @@ -94,7 +91,15 @@ function too_strict_limits(this) this.verifyEqual(data{outchannel}.header.chantype, 'hp'); this.verifyTrue(any(isnan(data{outchannel}.data))) end +function invalid_channel_action(this) + fn = this.input_filename; + sr = 100; + options = struct(); + options.channel_action = 'abc'; + + this.verifyWarning( @() pspm_convert_hb2hp(fn, sr, options), 'ID:invalid_input'); +end function add_replace_channel_action(this) fn = this.input_filename; sr = 1; diff --git a/test/pspm_dcm_test.m b/test/pspm_dcm_test.m index 0d138e2c7..d6eed7da4 100644 --- a/test/pspm_dcm_test.m +++ b/test/pspm_dcm_test.m @@ -125,6 +125,7 @@ function valid_input(this) test_hra1_flex_cs(this); test_hra1_flex_cs_missing(this); test_hra1_flex_cs_nan(this); + test_hra1_flex_cs_inestimable(this); end end methods @@ -208,6 +209,29 @@ function test_hra1_flex_cs_nan(this) end end + function test_hra1_flex_cs_inestimable(this) + % find free filename + fn = pspm_find_free_fn(this.modelfile_prfx, '_inestimable.mat'); + + [df, ~] = this.get_hra_files(1); + [timing, eventnames, trialnames] = this.extract_hra_timings(1); + + model = struct( ... + 'modelfile', fn, ... + 'datafile', df, ... + 'timing', {timing}, ... + ... % force the last trial to be marked as inestimable + 'lasttrialcutoff', 1e6 ... + ); + + options = struct( ... + 'dispwin', 0, ... + 'trlnames', {trialnames}, ... + 'eventnames', {eventnames} ... + ); + + this.verifyWarning( @() pspm_dcm(model, options), 'ID:inestimable_trials'); + end function [timing, eventnames, trialnames] = ... extract_hra_timings(this, subject) [s, c] = this.get_hra_files(subject); diff --git a/test/pspm_tam_test.m b/test/pspm_tam_test.m new file mode 100644 index 000000000..bbcc42f00 --- /dev/null +++ b/test/pspm_tam_test.m @@ -0,0 +1,590 @@ +classdef pspm_tam_test < matlab.unittest.TestCase + +properties (TestParameter) + + nFiles = struct( ... + 'oneFile', 1, ... + 'twoFiles', 2); + + stdExpCond = struct( ... + 'byName', 'condition_A', ... + 'byIndex', 1); + +end + + +methods (Test) + +function testSingleConditionTrialAverage(testCase, nFiles) + + sr = 10; + duration = 50; + window = 5; + + t = (0:1/sr:window-1/sr)'; + response = exp(-((t - 2).^2) / (2 * 0.5^2)); + onsets = [10 20 30]; + + datafiles = testCase.createTestData(nFiles, sr, duration, {response}, {onsets}); + + timing.names = {'condition_1'}; + timing.onsets = {onsets}; + timing.durations = {zeros(size(onsets))}; + + modelfile = [tempname '.mat']; + testCase.addTeardown(@() pspm_tam_test.deleteIfExists(modelfile)); + + model = struct(); + model.modelfile = modelfile; + model.datafile = datafiles; + model.timing = repmat({timing}, 1, nFiles); + model.timeunits = 'seconds'; + model.window = window; + model.modelspec = 'dilation'; + model.modality = 'pupil'; + model.channel = 'pupil'; + model.norm = 0; + model.baseline = 0; + model.norm_max = 0; + model.filter = struct('lpfreq', 'none', 'lporder', 1, 'hpfreq', 'none', 'hporder', 1, 'down', 'none', 'direction', 'bi'); + + options.overwrite = 1; + + [sts, tam] = pspm_tam(model, options); + + testCase.assertEqual(sts, 1); + + expected = response - response(1); + + testCase.verifyEqual(tam.data.Y{1}, expected(:), 'AbsTol', 1e-10); + testCase.verifyEqual(tam.data.sr, repmat({sr}, 1, nFiles)); + +end + + +function testDownsamplingTo5Hz(testCase) + + sr = 10; + duration = 50; + window = 5; + + t = (0:1/sr:window-1/sr)'; + response = exp(-((t - 2).^2) / (2 * 0.5^2)); + onsets = [10 20 30]; + + datafiles = testCase.createTestData(1, sr, duration, {response}, {onsets}); + + timing.names = {'condition_1'}; + timing.onsets = {onsets}; + timing.durations = {zeros(size(onsets))}; + + modelfile = [tempname '.mat']; + testCase.addTeardown(@() pspm_tam_test.deleteIfExists(modelfile)); + + model = struct(); + model.modelfile = modelfile; + model.datafile = datafiles; + model.timing = {timing}; + model.timeunits = 'seconds'; + model.window = window; + model.modelspec = 'dilation'; + model.modality = 'pupil'; + model.channel = 'pupil'; + model.norm = 0; + model.baseline = 0; + model.norm_max = 0; + model.filter = struct('lpfreq', 'none', 'lporder', 0, 'hpfreq', 'none', 'hporder', 0, 'down', 5, 'direction', 'bi'); + + options.overwrite = 1; + + [sts, tam] = pspm_tam(model, options); + + testCase.assertEqual(sts, 1); + + expectedSize = [25 1]; + + testCase.verifySize(tam.data.Y{1}, expectedSize); + testCase.verifySize(tam.data.std{1}, expectedSize); + testCase.verifySize(tam.data.sem{1}, expectedSize); + testCase.verifySize(tam.data.X{1}, expectedSize); + testCase.verifyEqual(tam.data.sr{1}, 5); + +end + + +function testLowpassFilterWithoutDownsampling(testCase) + + sr = 10; + duration = 50; + window = 5; + + t = (0:1/sr:window-1/sr)'; + response = exp(-((t - 2).^2) / (2 * 0.5^2)) + 0.1 * sin(2*pi*4*t); + onsets = [10 20 30]; + + datafiles = testCase.createTestData(1, sr, duration, {response}, {onsets}); + + timing.names = {'condition_1'}; + timing.onsets = {onsets}; + timing.durations = {zeros(size(onsets))}; + + modelfile = [tempname '.mat']; + testCase.addTeardown(@() pspm_tam_test.deleteIfExists(modelfile)); + + model = struct(); + model.modelfile = modelfile; + model.datafile = datafiles; + model.timing = {timing}; + model.timeunits = 'seconds'; + model.window = window; + model.modelspec = 'dilation'; + model.modality = 'pupil'; + model.channel = 'pupil'; + model.norm = 0; + model.baseline = 0; + model.norm_max = 0; + model.filter = struct('lpfreq', 1, 'lporder', 1, 'hpfreq', 'none', 'hporder', 0, 'down', 'none', 'direction', 'bi'); + + options.overwrite = 1; + + [sts, tam] = pspm_tam(model, options); + + testCase.assertEqual(sts, 1); + + testCase.verifySize(tam.data.Y{1}, [50 1]); + testCase.verifySize(tam.data.std{1}, [50 1]); + testCase.verifySize(tam.data.sem{1}, [50 1]); + testCase.verifySize(tam.data.X{1}, [50 1]); + + testCase.verifyEqual(tam.data.sr{1}, 10); + testCase.verifyTrue(tam.data.filtered); + + expected_t = (1:50)' / sr; + testCase.verifyEqual(tam.data.X{1}, expected_t, 'AbsTol', 1e-12); + + unfiltered = response - response(1); + testCase.verifyGreaterThan(norm(tam.data.Y{1} - unfiltered), 1e-6); + +end + + +function testTwoSessions(testCase) + + sr = 10; + duration = 50; + window = 5; + + t = (0:1/sr:window-1/sr)'; + response1 = exp(-((t - 2).^2) / (2 * 0.5^2)); + response2 = 2 * response1; + onsets = [10 20 30]; + + datafiles = testCase.createTestData(2, sr, duration, {response1; response2}, {onsets}); + + timing.names = {'condition_1'}; + timing.onsets = {onsets}; + timing.durations = {zeros(size(onsets))}; + + modelfile = [tempname '.mat']; + testCase.addTeardown(@() pspm_tam_test.deleteIfExists(modelfile)); + + model = struct(); + model.modelfile = modelfile; + model.datafile = datafiles; + model.timing = {timing, timing}; + model.timeunits = 'seconds'; + model.window = window; + model.modelspec = 'dilation'; + model.modality = 'pupil'; + model.channel = 'pupil'; + model.norm = 0; + model.baseline = 0; + model.norm_max = 0; + model.filter = struct('lpfreq', 'none', 'lporder', 1, 'hpfreq', 'none', 'hporder', 1, 'down', 'none', 'direction', 'bi'); + + options.overwrite = 1; + + [sts, tam] = pspm_tam(model, options); + + testCase.assertEqual(sts, 1); + + expected = (response1 + response2) / 2; + expected = expected - expected(1); + + testCase.verifyEqual(tam.data.Y{1}, expected(:), 'AbsTol', 1e-10); + testCase.verifyEqual(tam.data.sr, {10, 10}); + testCase.verifyEqual(numel(tam.data.Y), 1); + +end + + +function testStandardExperimentalCondition(testCase, stdExpCond) + + sr = 10; + duration = 60; + window = 5; + + t = (0:1/sr:window-1/sr)'; + baseResponse = exp(-((t - 2).^2) / (2 * 0.5^2)); + + responseA = 0.5 * baseResponse; + responseB = 2 * baseResponse; + + onsetsA = [10 30]; + onsetsB = [20 40]; + + datafiles = testCase.createTestData(1, sr, duration, {responseA, responseB}, {onsetsA, onsetsB}); + + timing.names = {'condition_A', 'condition_B'}; + timing.onsets = {onsetsA, onsetsB}; + timing.durations = {zeros(size(onsetsA)), zeros(size(onsetsB))}; + + modelfile = [tempname '.mat']; + testCase.addTeardown(@() pspm_tam_test.deleteIfExists(modelfile)); + + model = struct(); + model.modelfile = modelfile; + model.datafile = datafiles; + model.timing = {timing}; + model.timeunits = 'seconds'; + model.window = window; + model.modelspec = 'dilation'; + model.modality = 'pupil'; + model.channel = 'pupil'; + model.norm = 0; + model.baseline = 0; + model.norm_max = 0; + model.std_exp_cond = stdExpCond; + model.filter = struct('lpfreq', 'none', 'lporder', 1, 'hpfreq', 'none', 'hporder', 1, 'down', 'none', 'direction', 'bi'); + + options.overwrite = 1; + + [sts, tam] = pspm_tam(model, options); + + testCase.assertEqual(sts, 1); + + expectedA = responseA - responseA(1); + + expectedB = responseB - responseA; + expectedB = expectedB - expectedB(1); + + testCase.verifyEqual(tam.data.Y{1}, expectedA(:), 'AbsTol', 1e-10); + testCase.verifyEqual(tam.data.Y{2}, expectedB(:), 'AbsTol', 1e-10); + + testCase.verifyEqual(tam.data.std_exp_cond.name, 'condition_A'); + testCase.verifyEqual(tam.data.std_exp_cond.ind, 1); + +end + + +function testMarkerTimeunits(testCase, nFiles) + + sr = 10; + duration = 50; + window = 5; + + t = (0:1/sr:window-1/sr)'; + response = exp(-((t - 2).^2) / (2 * 0.5^2)); + + markerTimes = [10 20 30]; + + datafiles = testCase.createTestData(nFiles, sr, duration, {response}, {markerTimes}, markerTimes); + + timing.names = {'condition_1'}; + timing.onsets = {[1 2 3]}; + timing.durations = {zeros(1,3)}; + + modelfile = [tempname '.mat']; + testCase.addTeardown(@() pspm_tam_test.deleteIfExists(modelfile)); + + model = struct(); + model.modelfile = modelfile; + model.datafile = datafiles; + model.timing = repmat({timing}, 1, nFiles); + model.timeunits = 'markers'; + model.window = window; + model.modelspec = 'dilation'; + model.modality = 'pupil'; + model.channel = 'pupil'; + model.norm = 0; + model.baseline = 0; + model.norm_max = 0; + model.filter = struct('lpfreq', 'none', 'lporder', 1, 'hpfreq', 'none', 'hporder', 1, 'down', 'none', 'direction', 'bi'); + + options.overwrite = 1; + options.marker_chan = 'marker'; + + [sts, tam] = pspm_tam(model, options); + + testCase.assertEqual(sts, 1); + + expected = response - response(1); + + testCase.verifyEqual(tam.data.Y{1}, expected(:), 'AbsTol', 1e-10); + testCase.verifyEqual(tam.data.sr, repmat({sr}, 1, nFiles)); + +end + + +function testZscoreNormalization(testCase) + + sr = 10; + duration = 50; + window = 5; + + t = (0:1/sr:window-1/sr)'; + response = exp(-((t - 2).^2) / (2 * 0.5^2)); + onsets = [10 20 30]; + + [datafiles, y] = testCase.createTestData(1, sr, duration, {response}, {onsets}); + y = y{1}; + + timing.names = {'condition_1'}; + timing.onsets = {onsets}; + timing.durations = {zeros(size(onsets))}; + + modelfile = [tempname '.mat']; + testCase.addTeardown(@() pspm_tam_test.deleteIfExists(modelfile)); + + model = struct(); + model.modelfile = modelfile; + model.datafile = datafiles; + model.timing = {timing}; + model.timeunits = 'seconds'; + model.window = window; + model.modelspec = 'dilation'; + model.modality = 'pupil'; + model.channel = 'pupil'; + model.norm = 1; + model.baseline = 0; + model.norm_max = 0; + model.filter = struct('lpfreq', 'none', 'lporder', 1, 'hpfreq', 'none', 'hporder', 1, 'down', 'none', 'direction', 'bi'); + + options.overwrite = 1; + + [sts, tam] = pspm_tam(model, options); + + testCase.assertEqual(sts, 1); + + mu = mean(y, 'omitnan'); + sigma = std(y, 0, 'omitnan'); + + expected = (response - mu) / sigma; + expected = expected - expected(1); + + testCase.verifyEqual(tam.data.Y{1}, expected(:), 'AbsTol', 1e-10); + testCase.verifyEqual(tam.data.norm, 1); + +end + + +function testBaselineEqualWindowIsInvalid(testCase) + + sr = 10; + duration = 50; + window = 5; + + t = (0:1/sr:window-1/sr)'; + response = exp(-((t - 2).^2) / (2 * 0.5^2)); + onsets = [10 20 30]; + + datafiles = testCase.createTestData(1, sr, duration, {response}, {onsets}); + + timing.names = {'condition_1'}; + timing.onsets = {onsets}; + timing.durations = {zeros(size(onsets))}; + + modelfile = [tempname '.mat']; + testCase.addTeardown(@() pspm_tam_test.deleteIfExists(modelfile)); + + model = struct(); + model.modelfile = modelfile; + model.datafile = datafiles; + model.timing = {timing}; + model.timeunits = 'seconds'; + model.window = window; + model.modelspec = 'dilation'; + model.modality = 'pupil'; + model.channel = 'pupil'; + model.norm = 0; + model.baseline = window; + model.norm_max = 0; + + options.overwrite = 1; + + testCase.verifyWarning(@() pspm_tam(model, options), 'ID:invalid_input'); + +end + + +function testThreeSessionsWithTimingFiles(testCase) + + sr = 10; + duration = 50; + window = 5; + nFiles = 3; + + t = (0:1/sr:window-1/sr)'; + response = exp(-((t - 2).^2) / (2 * 0.5^2)); + onsets = [10 20 30]; + + %% Create three PsPM data files + datafiles = testCase.createTestData( ... + nFiles, sr, duration, {response}, {onsets}); + + %% Create three timing files + timingfiles = cell(nFiles, 1); + + names = {'condition_1'}; + onsets = {[10 20 30]}; + durations = {zeros(1,3)}; + + for k = 1:nFiles + timingfiles{k} = [tempname '.mat']; + save(timingfiles{k}, 'names', 'onsets', 'durations'); + + testCase.addTeardown( @() pspm_tam_test.deleteIfExists(timingfiles{k})); + end + + %% Output model file + modelfile = [tempname '.mat']; + testCase.addTeardown( @() pspm_tam_test.deleteIfExists(modelfile)); + + %% TAM model + model = struct(); + model.modelfile = modelfile; + model.datafile = datafiles; + + % Important part of this regression test: + % model.timing contains FILENAMES, not timing structs. + model.timing = timingfiles; + + model.timeunits = 'seconds'; + model.window = window; + model.modelspec = 'dilation'; + model.modality = 'pupil'; + model.channel = 'pupil'; + model.norm = 0; + model.baseline = 0; + model.norm_max = 0; + + model.filter = struct( 'lpfreq', 'none', 'lporder', 1, 'hpfreq', 'none', 'hporder', 1, 'down', 'none', 'direction', 'bi'); + + options.overwrite = 1; + + %% Run TAM + [sts, tam] = pspm_tam(model, options); + + %% Verify + testCase.assertEqual(sts, 1); + + testCase.verifyEqual( tam.names, {'condition_1'}); + + testCase.verifyEqual( tam.timing, timingfiles); + + testCase.verifyEqual( tam.data.sr, repmat({sr}, 1, nFiles)); + +end + + + +end + + +methods (Access = private) + +function [datafiles, signals] = createTestData(testCase, nFiles, sr, duration, responses, onsets, markerTimes) + + % responses: + % 1 x nCond -> same responses for all files + % nFiles x nCond -> different responses between files + % + % onsets: + % 1 x nCond -> same onsets for all files + % nFiles x nCond -> different onsets between files + % + % markerTimes: + % optional; [] means no marker channel + + if nargin < 7 + markerTimes = []; + end + + if size(responses, 1) == 1 && nFiles > 1 + responses = repmat(responses, nFiles, 1); + end + + if size(onsets, 1) == 1 && nFiles > 1 + onsets = repmat(onsets, nFiles, 1); + end + + datafiles = cell(1, nFiles); + signals = cell(1, nFiles); + + for k = 1:nFiles + + y = zeros(duration * sr, 1); + + for iCond = 1:size(responses, 2) + + response = responses{k, iCond}; + condOnsets = onsets{k, iCond}; + + for iTrial = 1:numel(condOnsets) + + startSample = round(condOnsets(iTrial) * sr) + 1; + stopSample = startSample + numel(response) - 1; + + y(startSample:stopSample) = y(startSample:stopSample) + response; + + end + end + + signals{k} = y; + + data = cell(1, 1); + data{1}.header = struct( 'chantype', 'pupil', 'sr', sr, 'units', 'a.u.'); + data{1}.data = y; + + if ~isempty(markerTimes) + + if iscell(markerTimes) + thisMarkerTimes = markerTimes{k}; + else + thisMarkerTimes = markerTimes; + end + + data{2}.header = struct( 'chantype', 'marker', 'sr', 1, 'units', 'events'); + data{2}.data = thisMarkerTimes(:); + + end + + infos.duration = duration; + infos.durationinfo = 'Duration in seconds'; + infos.source = struct(); + + datafiles{k} = [tempname '.mat']; + save(datafiles{k}, 'data', 'infos'); + + testCase.addTeardown(@() pspm_tam_test.deleteIfExists(datafiles{k})); + + end + +end + +end + + +methods (Static, Access = private) + +function deleteIfExists(filename) + + if exist(filename, 'file') + delete(filename); + end + +end + +end + +end \ No newline at end of file