% MWF for AEP removal during speech production
% Update: 10.2024
%
% =========================================================================
%
% Script accompanying the methodological paper:
%   De Pretto, M., Kodrasi, I., Laganaro M. (submitted). ERP signals during
%   speech articulation: does auditory feedback mask other ongoing 
%   cognitive-motor processes?
%
%
% INPUTS
% - ERP .sef files
%   - During speech production (signal to be filtered)
%   - During listening task (for estimation of noise via multi-channel 
%     Wiener filter)
%
%
% OUTPUTS
% - ERP .sef files of the filtered signal
%
%
% FUNCTIONS CALLED
% - compute_gfp
% - open_sef
% - save_sef
% These functions may be found at https://doi.org/10.5281/zenodo.14794980
%
%
% Author: Michael De Pretto (Michael.DePretto@unige.ch)
%
% =========================================================================




%% PARAMETERS

FirstSubjCode   = '';
FilterStringPRO = ''; % For ERP files during production
FilterStringLIS = ''; % For ERP files during listening
FilesExtension  = '.sef';
Onset           = '340'; % 300 ms pre-onset + 40 ms post-onset
Offset          = '473'; % 300 ms pre-onset + 173 ms post-onset
Splittime       = '393'; % time when component 2 starts

PromptSetup = {'Exact code of the FIRST participant:',...
    'Filtering string for selecting ERP PRODUCTION files (optional)',...
    'Filtering string for selecting ERP LISTENING files (optional)',...
    'ERP files extension:',...
    'Onset [ms]:',...
    'Offset [ms]:',...
    'Split time [ms]:'};
PromptInputs    = inputdlg(PromptSetup,'Design',1,{FirstSubjCode,FilterStringPRO,FilterStringLIS,FilesExtension,Onset,Offset,Splittime});
FirstSubjCode   = PromptInputs{1};
% If filter string entered, add '*' for file selection
if ~isempty(PromptInputs{2})
    FilterStringPRO = strcat(PromptInputs{2},'*');
else
    FilterStringPRO = PromptInputs{2};
end
if ~isempty(PromptInputs{3})
    FilterStringLIS = strcat(PromptInputs{3},'*');
else
    FilterStringLIS = PromptInputs{3};
end
FilesExtension  = PromptInputs{4};
Onset           = str2double(PromptInputs{5});
Offset          = str2double(PromptInputs{6});
Splittime       = str2double(PromptInputs{7});



%% SELECT FILES

% Path of the upper folder containing the ERP PRODUCTION files
root_folder = uigetdir('title',...
    'Choose the path of your most upper folder containing the ERP PRODUCTION files');
cd(root_folder)
PRO_FilesList   = dir(['**/*' FilterStringPRO FilesExtension]);


% Saving folders
CIpath = fullfile(root_folder,'MWF');
if ~exist(CIpath, 'dir')
    mkdir(CIpath);
end


% Path of the upper folder containing the ERP LISTENING files
root_folder = uigetdir('title',...
    'Choose the path of your most upper folder containing the E-Prime LISTENING files');
cd(root_folder)
LIS_FilesList   = dir(['**/*' FilterStringLIS FilesExtension]);



%% DESIGN

% Identify subjects
SUBJ            = strings([length(PRO_FilesList),1]); % Create array for subject codes
SubjNameIndex   = strfind(PRO_FilesList(1).name,FirstSubjCode); % Index of Subject code in the file name

% Extract subject codes from files
for file = 1:length(PRO_FilesList)
    if SubjNameIndex+strlength(FirstSubjCode)-1 > length(PRO_FilesList(file).name)
        continue
    else
        SUBJ(file,:) = PRO_FilesList(file).name(SubjNameIndex:SubjNameIndex+strlength(FirstSubjCode)-1);
    end
end
SUBJ  = unique(SUBJ);
nSubj = length(SUBJ); % number of subjects



%% APPLY MULTI-CHANNEL WIENER FILTER

for subj = 1:nSubj
    disp(['Processing subject ', num2str(subj), ' out of ', num2str(nSubj)])
    
    %% Prepare PRODUCTION data
    
    for file = 1:length(PRO_FilesList)
        % open file and prepare saving paths
        if contains(PRO_FilesList(file).name,SUBJ(subj,1))
            PROfilename     = [PRO_FilesList(file).folder,'\',PRO_FilesList(file).name]; % full path and name of the file
            
            % open filename for reading
            [PROhdr,PROdata,PROevt] = open_sef(PROfilename);
            
            % Extract window of interest

            % Component 1
            PROonsetTF1  = round((Onset / 1000) / (1 / PROhdr.SamplingRate)) + 1;
            PROoffsetTF1 = round((Splittime / 1000) / (1 / PROhdr.SamplingRate)) + 1;
            nTF1         = PROoffsetTF1 - (PROonsetTF1 - 1);
            
            PROepoch1    = PROdata(PROonsetTF1:PROoffsetTF1,:)';
            PRObl1       = mean(PROepoch1,2); % Baseline
            Y1           = PROepoch1 - PRObl1; % Baseline correction on whole window
            estRyy1      = 1/nTF1 * (Y1 * Y1');

            % Component 2
            PROonsetTF2  = round((Splittime / 1000) / (1 / PROhdr.SamplingRate)) + 1;
            PROoffsetTF2 = round((Offset / 1000) / (1 / PROhdr.SamplingRate)) + 1;
            nTF2         = PROoffsetTF2 - (PROonsetTF2 - 1);
            
            PROepoch2    = PROdata(PROonsetTF2:PROoffsetTF2,:)';
            PRObl2       = mean(PROepoch2,2); % Baseline
            Y2           = PROepoch2 - PRObl2; % Baseline correction on whole window
            estRyy2      = 1/nTF2 * (Y2 * Y2');
        end
    end
    

    %% Prepare LISTENING data
    
    for file = 1:length(LIS_FilesList)
        % open file and prepare saving paths
        if contains(LIS_FilesList(file).name,SUBJ(subj,1))
            LISfilename     = [LIS_FilesList(file).folder,'\',LIS_FilesList(file).name]; % full path and name of the file
            
            % open filename for reading
            [LIShdr,LISdata,LISevt] = open_sef(LISfilename);
            
            % Extract window of interest

            % Component 1
            LISonsetTF1  = round((Onset / 1000) / (1 / LIShdr.SamplingRate));
            LISoffsetTF1 = round((Splittime / 1000) / (1 / LIShdr.SamplingRate));
            nTF1         = LISoffsetTF1 - (LISonsetTF1 - 1);
            
            LISepoch1    = LISdata(LISonsetTF1:LISoffsetTF1,:)';
            LISbl1       = mean(LISepoch1,2); % Baseline
            F1           = LISepoch1 - LISbl1; % Baseline correction on whole window
            estRff1      = 1/nTF1 * (F1 * F1');

            % Component 2
            LISonsetTF2  = round((Splittime / 1000) / (1 / LIShdr.SamplingRate));
            LISoffsetTF2 = round((Offset / 1000) / (1 / LIShdr.SamplingRate));
            nTF2         = LISoffsetTF2 - (LISonsetTF2 - 1);
            
            LISepoch2    = LISdata(LISonsetTF2:LISoffsetTF2,:)';
            LISbl2       = mean(LISepoch2,2); % Baseline
            F2           = LISepoch2 - LISbl2; % Baseline correction on whole window
            estRff2      = 1/nTF2 * (F2 * F2');
        end
    end
    
    estRdd1 = estRyy1 - estRff1;
    
    % force estRdd1 to be positive (semi) definite
    [U,V] = eig(estRdd1);
    V(V<0) = 0;
    estRdd1 = U*V*U';
    
    
    estRdd2 = estRyy2 - estRff2;
    
    % force estRdd2 to be positive (semi) definite
    [U,V] = eig(estRdd2);
    V(V<0) = 0;
    estRdd2 = U*V*U';

    
    mult = 1;
    estRyy1 = estRyy1+10^(-mult)*eye(size(estRyy1)); % regularizing the cov matrix for numerical instability before inversion
    estW1 = (1 ./ estRyy1) * estRdd1;
    
    mult = 1;
    estRyy2 = estRyy2+10^(-mult)*eye(size(estRyy2)); % regularizing the cov matrix for numerical instability before inversion
    estW2 = (1 ./ estRyy2) * estRdd2;

    
    %% Recompute the signal
    
    estSIG1 = zeros(size(Y1));
    for y = 1:size(Y1,2)
        estSIG1(:,y) = estW1' * Y1(:,y);
    end

    estSIG2 = zeros(size(Y2));
    for y = 1:size(Y2,2)
        estSIG2(:,y) = estW2' * Y2(:,y);
    end
    
    % Rescale the data
    GFP_estSIG1      = compute_gfp(estSIG1(:,1)'); % GFP of the first data point of the estimated signal
    estSIG1          = estSIG1 / GFP_estSIG1; % Noramlize the estimated signal
    GFP_PROdata1     = compute_gfp(PROdata(PROonsetTF1,:) - PRObl1'); % GFP of the original signal at first data point of the processed data (BL corrected)
    estSIG1          = estSIG1 * GFP_PROdata1; % rescale the estimated data to match the unprocessed data

    GFP_estSIG2      = compute_gfp(estSIG2(:,1)'); % GFP of the first data point of the estimated signal
    estSIG2          = estSIG2 / GFP_estSIG2; % Noramlize the estimated signal
    GFP_PROdata2     = compute_gfp(PROdata(PROonsetTF2,:) - PRObl2'); % GFP of the original signal at first data point of the processed data (BL corrected)
    estSIG2          = estSIG2 * GFP_PROdata2; % rescale the estimated data to match the unprocessed data
    
    
    % Baseline uncorrection!
    estSIG1 = estSIG1 + PRObl1;
    estSIG2 = estSIG2 + PRObl2;
    estSIG = [estSIG1 estSIG2];
    
    
    % Insert filtered part in the production data
    NEWdata = PROdata;
    NEWdata(PROonsetTF1-1:PROoffsetTF2,:) = estSIG';
    
    SmoothTrans = movmean(NEWdata(PROonsetTF1-8:PROonsetTF1+8,:),9,'Endpoints','discard'); % for AEP C1 onset
    NEWdata(PROonsetTF1-4:PROonsetTF1+4,:) = SmoothTrans;
    
    SmoothTrans = movmean(NEWdata(PROoffsetTF1-8:PROoffsetTF1+8,:),9,'Endpoints','discard'); % for Splittime (C1-C2)
    NEWdata(PROoffsetTF1-4:PROoffsetTF1+4,:) = SmoothTrans;
    
    SmoothTrans = movmean(NEWdata(PROoffsetTF2-8:PROoffsetTF2+8,:),9,'Endpoints','discard'); % for AEP C2 offset
    NEWdata(PROoffsetTF2-4:PROoffsetTF2+4,:) = SmoothTrans;
    
    
    %% Save data
    
    NEWevt = [PROevt; [{PROonsetTF1-1} {PROoffsetTF2-1} {"Filtered"}]]; % -1 because Cartool starts at 0
    savefilename = strcat(CIpath,'\',SUBJ(subj,1),'_filtered.sef');
    save_sef(savefilename,NEWdata,PROhdr.SamplingRate,PROhdr.Channels,NEWevt)
end

disp('done')