clear all; clc; close all;

addpath( 'Z:\Jongrok\Backup\HyperBrainX\DecSpace2\Zenodo\code\functions' );
addpath( 'Z:\Jongrok\Backup\HyperBrainX\toolbox\dPCA-master\matlab'   );
addpath('Z:\Jongrok\Backup\HyperBrainX\toolbox\ssPreprocessTool\fieldtrip-20230318');
ft_defaults;

neural_Dir = 'Z:\Jongrok\Backup\HyperBrainX\DecSpace2\Zenodo\data';
behave_Dir = 'Z:\Jongrok\Backup\HyperBrainX\DecSpace2\Zenodo\data\behavior';
result_Dir = 'Z:\Jongrok\Backup\HyperBrainX\DecSpace2\Zenodo\result';

% Basic parameters
params.nSub    = 19; % Number of subjects
params.nFreq   = 2; % Number of frequencies to be analyzed
params.nCat1   = 2; % Number of categories of first  stimulus (animate/inanimate)
params.nCat2   = 2; % Number of categories of second stimulus (animate/inanimate)
params.twin    = [-1  2.498]; % Epoched time considered for analysis
params.stpsize = 0.002; % step size, in secs (0.002 means no smoothing)
params.stw     = params.twin(1):params.stpsize:params.twin(2); % Time points
params.stw     = round(params.stw*10000)/10000;
params.Nwind   = size(params.stw,2); % Number of time points

% Channel selection
params.nChan = 128;
NuiElc = [1 32 17 43 44 48 49 56 107 113 114 119 120 125 126 127 128]; % Nuisance electrodes
params.valid_chan = setdiff(1:params.nChan,NuiElc); % 111 channels
clear NuiElc

% dPCA parameters
params.dPCA.combinedParams = {{1, [1 3]}, {2, [2 3]}, {3}, {[1 2], [1 2 3]}};
params.dPCA.margNames = {'Stim1', 'Stim2', 'Condition-independent', 'S1/S2 Interaction'};
params.dPCA.ifSimultaneousRecording = 1;
params.dPCA.stpsize     = 0.1; % in secs (0.002 means no smoothing)
params.dPCA.winsizeCell = 300;
params.dPCA.twin        = [-0.5 2]; % [-0.8 2.3];
params.dPCA.stw         = params.dPCA.twin(1):params.dPCA.stpsize:params.dPCA.twin(2); % 0.75; for whole trial range
params.dPCA.stw         = round(params.dPCA.stw*10000)/10000;
params.dPCA.Nwind       = length(params.dPCA.stw);
params.dPCA.LambdaIter  = 10;
params.dPCA.NoD         = 5; % Number of Dimension that dPCA extract for each experimental variable
params.dPCA.NoD_Dist    = 5; % Number of Dimension considered when calculating Euclidean distance

% Parallel process
NoW = 24;           % Number of workers (parfor)
p   = parpool(NoW); % open parpool


%% dPCA analysis
for xSub = 1:params.nSub % Subjects
    for xFreq = 1:params.nFreq % Frequency

        if xFreq == 1
            load(fullfile(neural_Dir, sprintf('sub%d_alpha_filtered.mat', xSub)));
        elseif xFreq == 2
            load(fullfile(neural_Dir, sprintf('sub%d_delta_filtered.mat', xSub)));
        end

        % For subject 3: Remove neural data for run 6 due to missing behavioral data
        if xSub == 3
            MV(501:600) = [];
        end

        % Load behavior data
        participant_files = dir(fullfile(behave_Dir, sprintf('sub%d_behavior_run*.mat', xSub)));
        num_files = length(participant_files);

        clear condBhv
        for xRun = 1:num_files
            % Load the mat file for the current run of the participant
            condBhv(xRun) = load(fullfile(behave_Dir, participant_files(xRun).name));
        end

        clear data
        data.Cat1 = [];
        data.Cat2 = [];
        data.Same = [];
        data.Resp = [];
        data.Corr = [];
        data.RT   = [];
        for xRun = 1:size(condBhv, 2) % Block number
            data.Cat1 = [data.Cat1  condBhv(xRun).data.Cat1];
            data.Cat2 = [data.Cat2  condBhv(xRun).data.Cat2];
            data.Same = [data.Same  condBhv(xRun).data.Same];
            data.Resp = [data.Resp  condBhv(xRun).data.Resp];
            data.Corr = [data.Corr  condBhv(xRun).data.Corr];
            data.RT   = [data.RT    condBhv(xRun).data.RT  ];
        end
        clear condBhv


        % Data reconstruction
        rejectIdx = [];
        for xTr = 1:length(data.Cat1) % Number of trials
            if isempty(MV{xTr}) || data.Corr(xTr) == 0
                rejectIdx(end+1) = xTr;
            end
        end

        MV(rejectIdx) = [];
        data.Cat1(rejectIdx) = [];
        data.Cat2(rejectIdx) = [];
        data.Same(rejectIdx) = [];
        data.Resp(rejectIdx) = [];
        data.Corr(rejectIdx) = [];
        data.RT  (rejectIdx) = [];
        clear rejectIdx

        for xTr = 1:length(MV)
            data.MV(:,:,xTr) = MV{xTr};
            data.Trial(xTr)  = xTr;
        end
        clear MV


        NoT = size(data.MV, 3); % Number of trials
        ix = randperm(NoT);
        data.MV    = data.MV(:,:,ix);
        data.Cat1  = data.Cat1(  ix);
        data.Cat2  = data.Cat2(  ix);
        data.Same  = data.Same(  ix);
        data.Resp  = data.Resp(  ix);
        data.Corr  = data.Corr(  ix);
        data.RT    = data.RT(    ix);
        data.Trial = data.Trial( ix);

        clear ix_r
        for xTr = 1:length(ix)
            ix_r(xTr) = find(ix == xTr);
        end

        % Path indexing
        data.Path = [];
        for xTr = 1:length(data.Trial)
            if     data.Cat1(xTr) == 1 && data.Cat2(xTr) == 1
                data.Path(xTr)  = 1;
            elseif data.Cat1(xTr) == 1 && data.Cat2(xTr) == 2
                data.Path(xTr)  = 2;
            elseif data.Cat1(xTr) == 2 && data.Cat2(xTr) == 1
                data.Path(xTr)  = 3;
            elseif data.Cat1(xTr) == 2 && data.Cat2(xTr) == 2
                data.Path(xTr)  = 4;
            end
        end

        % Sorting according to path
        clear pathMat
        pathMat{1} = data.MV(:,:,find(data.Path == 1));
        pathMat{2} = data.MV(:,:,find(data.Path == 2));
        pathMat{3} = data.MV(:,:,find(data.Path == 3));
        pathMat{4} = data.MV(:,:,find(data.Path == 4));

        maxTriNum = max([size(pathMat{1}, 3) size(pathMat{2}, 3) size(pathMat{3}, 3) size(pathMat{4}, 3)]);

        % Elec x Stim1 x Stim2 x Time x NumTr
        data.MV_P = nan(params.nChan, params.nCat1, params.nCat2, params.Nwind, maxTriNum);

        data.MV_P(:,1,1,:,1:size(pathMat{1},3)) = pathMat{1};
        data.MV_P(:,1,2,:,1:size(pathMat{2},3)) = pathMat{2};
        data.MV_P(:,2,1,:,1:size(pathMat{3},3)) = pathMat{3};
        data.MV_P(:,2,2,:,1:size(pathMat{4},3)) = pathMat{4};

        data.numOfTrials        = zeros(params.nChan, params.nCat1, params.nCat2);
        data.numOfTrials(:,1,1) = size(pathMat{1},3);
        data.numOfTrials(:,1,2) = size(pathMat{2},3);
        data.numOfTrials(:,2,1) = size(pathMat{3},3);
        data.numOfTrials(:,2,2) = size(pathMat{4},3);

        data.MV_Pavg = nanmean(data.MV_P, 5);
        clear pathMat

        for xAxis = 1:3 % 1_Stim1, 2_Stim2, 3_Decision(Interaction)
            results.Cor(xAxis).mat{xFreq} = nan(params.nCat1, params.nCat2, params.dPCA.Nwind, params.Nwind, maxTriNum);
            results.DoP(xAxis).mat{xFreq} = nan(params.nCat1, params.nCat2, params.dPCA.Nwind, params.Nwind, maxTriNum);
        end

        for xTm = 1:params.dPCA.Nwind
            winCenter = find(params.stw == params.dPCA.stw(xTm));
            winRange  = (winCenter-params.dPCA.winsizeCell/2):(winCenter+params.dPCA.winsizeCell/2);

            Xtrial      = data.MV_P   (params.valid_chan,:,:,winRange,:);
            Xfull       = data.MV_Pavg(params.valid_chan,:,:,winRange,:);
            numOfTrials = data.numOfTrials(params.valid_chan,:,:);

            % default input parameters
            options = struct('numComps',        params.dPCA.NoD, ...
                'lambda',          0,                    ... % optimalLambda or 0
                'numRep',          100,                  ...
                'verbose',         'yes',                ...
                'decodingClasses', [],                   ...
                'timeSplits',      [],                   ...
                'timeParameter',   [],                   ...
                'notToSplit',      [],                   ...
                'filename',        [],                   ...
                'simultaneous',    true,                 ...
                'noiseCovType',    'pooled');

            options.combinedParams = params.dPCA.combinedParams;

            % find time marginalization
            timeComp = [];
            for k = 1:length(options.combinedParams)
                if options.combinedParams{k}{1} == length(size(Xfull))-1
                    timeComp = k;
                    break
                end
            end

            % set the number of components for the dPCA
            numCompsToUse = repmat(options.numComps, [1 length(options.combinedParams)]);
            numCompsToUse(timeComp) = 0;

            % set decoding classes
            dim = size(Xfull);
            neuronsConditions = squeeze(mean(numOfTrials, 1));

            loopSize = 2; % N-trials-leave-out cross validation
            loopPath = [0:loopSize:maxTriNum maxTriNum];

            parfor k = 1:(length(loopPath)-1)
                   [temp_Cor_Cond1{k}, ... % Correctness of decoding
                    temp_Cor_Cond2{k}, ...
                    temp_Cor_Inter{k}, ...
                    temp_DoP_Cond1{k}, ... % Degree of Proximity
                    temp_DoP_Cond2{k}, ...
                    temp_DoP_Inter{k}] = ...
                    dPCA_classification(Xtrial, (loopPath(k)+1):(loopPath(k+1)), numOfTrials, numCompsToUse, options, dim, params);
            end

            Cor_Cond1 = cat(4, temp_Cor_Cond1{:});
            Cor_Cond2 = cat(4, temp_Cor_Cond2{:});
            Cor_Inter = cat(4, temp_Cor_Inter{:});
            DoP_Cond1 = cat(4, temp_DoP_Cond1{:});
            DoP_Cond2 = cat(4, temp_DoP_Cond2{:});
            DoP_Inter = cat(4, temp_DoP_Inter{:});

            % 1_Stim1, 2_Stim2, 3_Interaction
            results.Cor(1).mat{xFreq}(:,:,xTm,winRange,:) = Cor_Cond1;
            results.Cor(2).mat{xFreq}(:,:,xTm,winRange,:) = Cor_Cond2;
            results.Cor(3).mat{xFreq}(:,:,xTm,winRange,:) = Cor_Inter;
            results.DoP(1).mat{xFreq}(:,:,xTm,winRange,:) = DoP_Cond1;
            results.DoP(2).mat{xFreq}(:,:,xTm,winRange,:) = DoP_Cond2;
            results.DoP(3).mat{xFreq}(:,:,xTm,winRange,:) = DoP_Inter;

            fprintf('Decoding (dPCA): Subject %2d/%2d  Time %2d/%2d\n', xSub, params.nSub, xTm, params.dPCA.Nwind);
            clear options numCompsToUse timeComp numOfTrials dim loopSize
            clear loopPath winCenter winRange Xfrial Xfull neuronsConditions
            clear temp_Cor_Cond1 temp_Cor_Cond2 temp_Cor_Inter
            clear temp_DoP_Cond1 temp_DoP_Cond2 temp_DoP_Inter
            clear Cor_Cond1 Cor_Cond2 Cor_Inter
            clear DoP_Cond1 DoP_Cond2 DoP_Inter
        end

        for xAxis = 1:3 % 1_Stim1, 2_Stim2, 3_Decision(Interaction)
            results.Cor(xAxis).mat{xFreq} = squeeze(nanmean(results.Cor(xAxis).mat{xFreq}, 3));
            results.DoP(xAxis).mat{xFreq} = squeeze(nanmean(results.DoP(xAxis).mat{xFreq}, 3));
        end

        clear data ix ix_r NoT maxTriNum
    end % xFreq

    save(fullfile(result_Dir, sprintf('results_sub%d.mat', xSub)), 'results', '-v7.3');
    clear results
end % xSub



%% Group analysis
clear results
for xSub = 1:params.nSub % Subject
    results_SN = load(fullfile(result_Dir, sprintf('results_sub%d.mat', xSub)));
    results_SN = results_SN.results;

    for xFreq = 1:params.nFreq % Frequency
        clear tempMat 
        for xCond = 1:3 % 1_Stim1, 2_Stim2, 3_Decision(Interaction)
            tempMat.Cor(xCond).mat = permute(results_SN.Cor(xCond).mat{xFreq}, [3 1 2 4]);
            tempMat.Cor(xCond).dim = size(tempMat.Cor(xCond).mat); % Time x Stim1 x Stim2 x Trials
            tempMat.Cor(xCond).mat = reshape(tempMat.Cor(xCond).mat, [tempMat.Cor(xCond).dim(1)  tempMat.Cor(xCond).dim(2)*tempMat.Cor(xCond).dim(3)*tempMat.Cor(xCond).dim(4)]);
            results.Cor(xCond).mat{xFreq}(xSub,:) = gaussianFilter(nanmean(tempMat.Cor(xCond).mat, 2), 5);
        end
        clear tempMat
    end % xFreq
    fprintf('Subject %2d/%2d\n', xSub, params.nSub);
    clear results_SN tempMat
end % xSub



%% Visualization (Decoding accuracy)
condTitle = {'Stim1','Stim2','Decision'};
color_mat{1} = [230 49 51];
color_mat{2} = [75 138 190];
color_mat{3} = [95 182 92];

for xFreq = 1:params.nFreq % Frequency
    clear graph
    for xCond = 1:3 % 1_Stim1, 2_Stim2, 3_Decision(Interaction)
        graph.mat{xCond} = results.Cor(xCond).mat{xFreq}*100;
        graph.avg{xCond} = gaussianFilter(squeeze(mean(graph.mat{xCond}, 1)), 10);
        graph.Eb{xCond}  = std(graph.mat{xCond}, 1)/sqrt(params.nSub);
    end

    figure1 = figure; hold on;
    plot(params.stw, graph.avg{1}, 'color', color_mat{1}/255, 'LineWidth', 1.5);
    plot(params.stw, graph.avg{2}, 'color', color_mat{2}/255, 'LineWidth', 1.5);
    plot(params.stw, graph.avg{3}, 'color', color_mat{3}/255, 'LineWidth', 1.5);

    legend({'Stim1', 'Stim2', 'Decision'}, 'AutoUpdate', 'off');
    shadedErrorBar(params.stw, graph.avg{1}, graph.Eb{1}, {'color', color_mat{1}/255}, 0.5);
    shadedErrorBar(params.stw, graph.avg{2}, graph.Eb{2}, {'color', color_mat{2}/255}, 0.5);
    shadedErrorBar(params.stw, graph.avg{3}, graph.Eb{3}, {'color', color_mat{3}/255}, 0.5);

    plot(params.stw, 50*ones(length(params.stw)), 'k');
    plot([0   0],   [48 55], 'k--');
    plot([0.5 0.5], [48 55], 'k--');
    plot([1   1],   [48 55], 'k--');
    plot([1.5 1.5], [48 55], 'k--');
    set(gca,'XLim',[-0.5 2])
    set(gca,'YLim',[49 55])
    xlabel('Time from 1st stimulus onset (s)');
    ylabel('Decoding accuracy (%)');
    if xFreq == 1
        title('Alpha band (8-15hz)');
    elseif xFreq == 2
        title('Delta band (1-4hz)');
    end
    clear graph figure1
end


