% FINAL SCRIPT FOR LPLASTIN RESULTS GENERATION
% clear all, clc, %close all force

% Model and measurement files
InputFile = 'Lplastin_model_update_mod.xlsx'; %% Network structure file
ContextsList = {'BT20','HCC38','MCF7', 'SkBR3'};
GlobalEdgesList = 'Lplastin_nofixed.xlsx'; 

MeasFileList = {};
for f = 1:length(ContextsList)
    MeasFileList = [MeasFileList,[char(ContextsList(f)) '_data_philippe.xlsx']];
end

estim = FalconMakeGlobalModel(InputFile, GlobalEdgesList, MeasFileList, ContextsList); %make the model

Np=length(estim.param_vector(1:end)); Nc=numel(ContextsList); Ns=length(estim.state_names);
R=[];
for c=1:(Nc):Np, R = [R;reshape(1:(Nc),1,Nc)+(c-1)]; end
N = 0;

estimSingle = FalconMakeGlobalModel(InputFile, InputFile, MeasFileList, ContextsList); %make the model
estimSingle.SSthresh = 0.001%estimSingle.SSthresh*1000000;
estimSingle.optRound = 4;
estimSingle = FalconOptimize(estimSingle);
estimSingle.NDatasets = 50;
estimSingle = FalconResample(estimSingle);
%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%
%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%

Costs_all = [];
ListLambda12 = power(2,[-100, -12, -11, -10, -9, -8, -7, -6, -5, -4, -3]);
ListLambdaCluster = power(2,[-100, -12, -11, -10, -9, -8, -7, -6, -5, -4, -3, -2, -1, 0, 1]);
AICs = zeros(length(ListLambda12), length(ListLambdaCluster));
MSEs = zeros(length(ListLambda12), length(ListLambdaCluster));
BICs = zeros(length(ListLambda12), length(ListLambdaCluster));
NPs = zeros(length(ListLambda12), length(ListLambdaCluster));
Times = zeros(length(ListLambda12), length(ListLambdaCluster));
Params = zeros(length(ListLambda12), length(ListLambdaCluster), length(estim.param_vector));

single_MSEs = [];
best_single_params = [];
liststd = []
for n = 1:length(MeasFileList)
    MeasFile = MeasFileList{n};
    estim = FalconMakeModel(InputFile, MeasFile);
    estim.SSthresh = estim.SSthresh*10000;
    estim.optRound = 4;
    estim = FalconOptimize(estim);
    [estim, StateValues, MSE] = FalconSimul(estim);
    single_MSEs = [single_MSEs; MSE];
    best_single_params = [best_single_params; estim.Results.Optimization.BestParams];
    estim.NDatasets = 20;
    estim = FalconResample(estim);
    
    liststd = [liststd; estim.Results.Resampling.OptimisedSD]
end

%%
% Builds a FALCON model for optimisation
estim = FalconMakeGlobalModel(InputFile, GlobalEdgesList, MeasFileList, ContextsList); %make the model
estim.SSthresh = estim.SSthresh*1000000;
%%%%
% Optimization
estim = FalconOptimize(estim);

%%% Re-simulate results based on the best optimised parameter set
[estim, StateValues, MSE] = FalconSimul(estim);


estim.Reg = 'PruneCluster';
estim.RegMatrix.Cluster = R;
            
%% Run 2D Regularization (about 50 hours)

l1=0;
Dims=[];
for L1=ListLambda12
    l2=0;
    l1=l1+1;
    for L2=ListLambdaCluster
        l2=l2+1;
        N=N+1;
        if AICs(l1,l2)==0
            Dims=[Dims;[L1,L2]];
            estim.Lambda = [L1, L2];
            % Optimization
            estim = FalconOptimize(estim);
            Costs_all = [Costs_all; ...
                min(estim.Results.Optimization.FittingCost), ...
                estim.Results.Optimization.BestParams, ...
                min(estim.Results.Optimization.FittingTime), ...
                estim.Results.Optimization.MSE, ...
                estim.Results.Optimization.BIC, ...
                estim.Results.Optimization.AIC, ...
                estim.Results.Optimization.Nparams];
            Idx=find(ismember(min(Costs_all(:,end-2)),Costs_all(:,end-2))); Idx = Idx(1);

            AICs(l1,l2)=Costs_all(Idx,end-3);
            MSEs(l1,l2)=Costs_all(Idx,end-2);
            BICs(l1,l2)=Costs_all(Idx,end-1);
            NPs(l1,l2)=Costs_all(Idx,end);
            Times(l1,l2)=fxt_all(Idx,end);
            Params(l1,l2,:)=bestx;

            save('run_LPL_Final')
        end
    end
end

%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%
%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%
%% figures of regularization results
figure,
subplot(2,2,1)
imagesc(AICs), set(gca, 'xtick', 1:length(ListLambdaCluster)), set(gca, 'xticklabel', {ListLambdaCluster}), set(gca, 'xticklabelrotation', 90), title('AIC')
set(gca, 'ytick', 1:length(ListLambda12)), set(gca, 'yticklabel', {ListLambda12}), colorbar, colormap('hot'), xlabel('Uniformity'), ylabel('Pruning')
subplot(2,2,2)
imagesc(log(MSEs)), set(gca, 'xtick', 1:length(ListLambdaCluster)), set(gca, 'xticklabel', {ListLambdaCluster}), set(gca, 'xticklabelrotation', 90), title('log(MSE)')
set(gca, 'ytick', 1:length(ListLambda12)), set(gca, 'yticklabel', {ListLambda12}), colorbar, colormap('hot'), xlabel('Uniformity'), ylabel('Pruning')
subplot(2,2,3)
imagesc(BICs), set(gca, 'xtick', 1:length(ListLambdaCluster)), set(gca, 'xticklabel', {'No Reg', '-12', '-11', '-10', '-9', '-8', '-7', '-6', '-5', '-4', '-3', '-2', '-1', '0', '1'}), set(gca, 'xticklabelrotation', 90), title('Bayesian Information Criterion')
set(gca, 'ytick', 1:length(ListLambda12)), set(gca, 'yticklabel', {'No Reg', '-12', '-11', '-10', '-9', '-8', '-7', '-6', '-5', '-4', '-3'}), colorbar, colormap('hot'), xlabel('log_2(\lambda_{Uniformity})', 'Interpreter', 'latex'), ylabel('log_2(\lambda_{Pruning})', 'Interpreter', 'latex')
subplot(2,2,4)
imagesc(NPs), set(gca, 'xtick', 1:length(ListLambdaCluster)), set(gca, 'xticklabel', {ListLambdaCluster}), set(gca, 'xticklabelrotation', 90), title('Number of Params')
set(gca, 'ytick', 1:length(ListLambda12)), set(gca, 'yticklabel', {ListLambda12}), colorbar, colormap('hot'), xlabel('Uniformity'), ylabel('Pruning')
figure, 
imagesc(Times), set(gca, 'xtick', 1:length(ListLambdaCluster)), set(gca, 'xticklabel', {ListLambdaCluster}), set(gca, 'xticklabelrotation', 90), title('CPU time')
set(gca, 'ytick', 1:length(ListLambda12)), set(gca, 'yticklabel', {ListLambda12}), colorbar, colormap('hot'), xlabel('Uniformity'), ylabel('Pruning')
figure,
subplot(2,2,1)
surf(AICs), set(gca, 'xtick', 1:length(ListLambdaCluster)), set(gca, 'xticklabel', {ListLambdaCluster}), set(gca, 'xticklabelrotation', 90), title('AIC')
set(gca, 'ytick', 1:length(ListLambda12)), set(gca, 'yticklabel', {ListLambda12}), colorbar, colormap('hot'), xlabel('Uniformity'), ylabel('Pruning')
subplot(2,2,2)
surf(log(MSEs)), set(gca, 'xtick', 1:length(ListLambdaCluster)), set(gca, 'xticklabel', {ListLambdaCluster}), set(gca, 'xticklabelrotation', 90), title('log(MSE)')
set(gca, 'ytick', 1:length(ListLambda12)), set(gca, 'yticklabel', {ListLambda12}), colorbar, colormap('hot'), xlabel('Uniformity'), ylabel('Pruning')
subplot(2,2,3)
surf(BICs), set(gca, 'xtick', 1:length(ListLambdaCluster)), set(gca, 'xticklabel', {ListLambdaCluster}), set(gca, 'xticklabelrotation', 90), title('BIC')
set(gca, 'ytick', 1:length(ListLambda12)), set(gca, 'yticklabel', {ListLambda12}), colorbar, colormap('hot'), xlabel('Uniformity'), ylabel('Pruning')
subplot(2,2,4)
surf(NPs), set(gca, 'xtick', 1:length(ListLambdaCluster)), set(gca, 'xticklabel', {ListLambdaCluster}), set(gca, 'xticklabelrotation', 90), title('Number of Params')
set(gca, 'ytick', 1:length(ListLambda12)), set(gca, 'yticklabel', {ListLambda12}), colorbar, colormap('hot'), xlabel('Uniformity'), ylabel('Pruning')


%% set up best model

load('run_LPL_Final')
estimBase = estim;

ind = find(ismember(BICs, min(min(BICs))));
[I,J] = ind2sub([length(ListLambda12), length(ListLambdaCluster)], ind);
AllMeanStateValues = [];
AllCosts = [];
AllBestParams = [];

BestParams = squeeze(Params(I,J,:));
BestParams_orig = BestParams(R);
ParamsSpread = std(BestParams_orig,0,2);
RRParams = estim.param_vector(R);
SParams = RRParams(:,1); SParams = cellfun(@(x) x(1:end-5), SParams, 'UniformOutput', 0);

[SortedSpread, O] = sort(ParamsSpread, 'ascend');
Data_idx = estim.Output_idx(1,:);
ZeroParams = RRParams(BestParams_orig<0.01);
OneParams = RRParams(BestParams_orig>0.99);

Int = estim.Interactions;

for c = 1:size(ZeroParams, 1)
    Int(strcmp(Int(:,5),ZeroParams(c)), 5) = {'0'};
end

for c = 1:size(OneParams, 1)
    Int(strcmp(Int(:,5),OneParams(c)), 5) = {'1'};
end

for c = 1:size(ParamsSpread,1)
    if ParamsSpread(c)<0.01
        for cc = 1:size(RRParams, 2)
            Int(strcmp(Int(:,5), RRParams(c,cc)), 5) = SParams(c);
        end
    end
end

FalconInt2File(Int,'tempInt.txt')
FalconData2File(estim)

%% make final model
estim = FalconMakeModel('ShavenModel_final.xlsx', 'GlobalInputs.xls');
estim.SSthresh = estim.SSthresh*1000000;
estim.optRound = 10
% Optimization
estim = FalconOptimize(estim);

%%% Re-simulate results based on the best optimised parameter set
[estim, MeanStateVal_orig, MSE] = FalconSimul(estim);
%% Single model
estim3 = FalconMakeGlobalModel(InputFile, InputFile, MeasFileList, ContextsList); %make the model
estim3.SSthresh = estim3.SSthresh*1000000;
estim3 = FalconOptimize(estim3);
[estim3, StateValues, MSE] = FalconSimul(estim3);





%% bootstrap
BootResults = [];

for this_boot = 1:4
    estim_boot = FalconPreProcess(estim, 'bootstrap', 'rows')
    estim_boot = FalconPreProcess(estim_boot, 'normalize', [0 1])
    
    estim_boot = FalconOptimize(estim_boot)
    
    BootResults = [BootResults; estim_boot.Results.Optimization.BestParams];
    
end

%% resample
estim.NDatasets = 20;
estim = FalconResample(estim);



save('run_LPL_Final2')
%% Bootstrap figure
figure, hold on, 
sinaplot(BootResults), hold on,
set(gca, 'xtick', 1:size(BootResults, 2))
set(gca, 'xticklabel', estim.Results.Optimization.ParamNames)
set(gca, 'xticklabelrotation', 90)
ylim([0 1])
set(gca, 'XGrid', 'on');
title('Distribution of optimised parameters after bootstrapping')

%% to have BICs by state
estimBest = FalconKONodes_fast(estimBest);

%% flux analysis
M = zeros(size(Nodes, 1), length(Params))
Nodes = estim.Results.Optimization.StateValueAll
Params = estim.Results.Optimization.BestParams
Fluxes = []


for cell_line = 1:4
    M = zeros(size(Nodes, 1), length(Params))
    for this_cond = 1:20
        for this_param = 1:85
            param_name = estim.param_vector{this_param}
            int_idx = find(ismember(estim.Interactions(:,5), param_name))
            if length(int_idx) > 1
                int_idx = int_idx(cell_line)
            end
            upstream_node = estim.Interactions{int_idx, 2}
            upstream_node_val = estim.Results.Optimization.StateValueAll(this_cond, find(ismember(estim.state_names, upstream_node)))
            M(this_cond, this_param) = upstream_node_val * Params(this_param);

        end
    end
    Fluxes(:,:,cell_line) = M;
end

for cell_line = 1:4
    figure, hold on,
    data = Fluxes(:,:,cell_line)-mean(Fluxes,3);
    isdif = sum(abs(data))>0.01;
    data = data(:, isdif)
    imagesc(data, [-0.2 0.2]), colorbar,
    set(gca, 'xtick', 1:size(data, 2))
    set(gca, 'xticklabel', estim.param_vector(isdif))
    set(gca, 'xticklabelrotation', 90)
    set(gca, 'ytick', 1:size(data, 1))
    set(gca, 'yticklabel', cellstr([estim.Annotation; 'mean']))
    axis([0.5, size(data, 2)+0.5, 0.5, size(data, 1)+0.5])
    colormap('hot')
    title(['Flux differential ', ContextsList(cell_line)])
    
end




%%

estim_i = FalconMakeModel('ShavenModel_Inhib.xlsx', 'GlobalInputs_ZeroInhib.xls');
estim_i.SSthresh = estim_i.SSthresh;

estim_i = FalconOptimize(estim_i)

x = estim_i.Results.Optimisation.BestStates

num_plots = size(estim.Output, 2);
NLines = ceil(sqrt(num_plots));
NCols = ceil(num_plots/NLines);

% Plot molecular profiles
for counter = 1:num_plots
    h1 = figure; hold on

    % Plot simulated data on top %figure, 
    h = bar(1:size(x, 1), diag(x(:, estim.Output_idx(1, counter))), 'stacked')
    set(h(:), 'facecolor', [0.8 0.95 0.95])
    set(h(1:5:size(x, 1)),'facecolor','g')
    % Figure adjustment
    axis([0 size(x,1)+1 0 1.1])
    set(gca, 'XTick',  1:size(x, 1))
    set(gca, 'XTickLabel', {'ctrl Neg', 'ctrl mTORCi', 'ctrl MEKi', 'ctrl mTORCi+MEKi' ...
        'ctrl SGKi+MEKi', 'EGF Neg', 'EGF mTORCi', 'EGF MEKi', 'EGF mTORCi+MEKi', ...
        'EGF SGKi+MEKi', 'HGF Neg', 'HGF mTORCi', 'HGF MEKi', 'HGF mTORCi+MEKi', ...
        'HGF SGKi+MEKi', 'IGF Neg' 'IGF mTORCi', 'IGF MEKi', 'IGF mTORCi+MEKi', ...
        'IGF SGKi+MEKi', 'PMA Neg', 'PMA mTORCi', 'PMA MEKi', 'PMA mTORCi+MEKi', 'PMA SGKi+MEKi'})
    set(gca, 'XTickLabelRotation', 45)
    set(gca,'fontsize',15)
    set(gca,'XGrid','on')
    t = title(estim.state_names(estim.Output_idx(1,counter)));
    set(t,'fontsize',25)
    hold off
end

