
import os
import csv
import numpy as np
import json
import matplotlib
matplotlib.use('Agg')

import matplotlib.pyplot as plt
import matplotlib.patches as patches
import seaborn as sns

import Spikes_fun 

# PSTH part 
dirname = 'figure2'
# dirname = 'prova_prob'
filename = dirname + '/All_options'
with open(filename + '.txt', 'r') as f:
    All_options = json.load(f)

filename = dirname + '/cells_is'
with open(filename + '.txt', 'r') as f:
    cell_ids = json.load(f)

filename = dirname + '/simulation_results'
with open(filename + '.txt', 'r') as f:
    dat = json.load(f)

"""
maxv = np.zeros(len(dat['Velocity_eye_tot']))
targets = np.array(dat['targets'])
for i in range(len(dat['Velocity_eye_tot'])):
    maxv[i] = max(dat['Velocity_eye_tot'][i])
unique_targets = np.unique(targets)
maxv_per_targets = []
errors_per_targets = []
errors = np.array(dat['Error'])
j = 1
for i in range(len(unique_targets)):
    mask_curr_unique_target = targets == unique_targets[i]
    maxv_per_targets.append(maxv[mask_curr_unique_target])
    plt.subplot(4,2,j)
    max_speed_last_trials = maxv_per_targets[i][-10:].mean()
    j += 1
    plt.title('target displacement = '+str(unique_targets[i]) + '\n' + 'max speed last 10 trials = '+ str(int(max_speed_last_trials)))
    plt.ylabel('max speed')
    plt.xlabel('learning trials')
    
    plt.plot(maxv_per_targets[i])

    mask_curr_unique_target = targets == unique_targets[i]
    errors_per_targets.append(errors[mask_curr_unique_target])
    plt.subplot(4,2,j)
    j += 1
    plt.title('error = '+str(unique_targets[i]))
    plt.ylabel('error')
    plt.xlabel('learning trials')
    
    plt.plot(errors_per_targets[i])

plt.tight_layout(pad=1.08)
plt.savefig('p.svg')
plt.close()"""

# f1 = plt.figure(figsize=(21, 5))
f1 = plt.figure(figsize=(120, 54))
gs = f1.add_gridspec(6, 9)

Saccade_time_before = 100
Saccade_time_after = 20
bin_smooth = 10


label_size = 100
label_size_tick = 100
legend_size = 75
title_size  = 100
lin_space_r = 20.0
lin_space_r_training = 7.5

fontweight_label = 'normal'
fontweight_title = 'bold'

time_tick_step = 20

Q0 = 0.25
Q1 = 0.75

n_trials_per_block = 10

# font_size = 2
trial_time = np.linspace(0,  Saccade_time_before + Saccade_time_after-1, Saccade_time_before + Saccade_time_after, dtype= int)
t_ref = (len(trial_time) - (len(dat['P_NO_CER']) + All_options['option_time_stimulation']['time_after_movement'] - All_options['option_time_stimulation']['time_movement_delay'] + Saccade_time_after))
trial_time -= t_ref

n_time_tick_step = int((80+20)/time_tick_step)
time_ticks = np.linspace(-20,80, n_time_tick_step+1, dtype = int)

PSTH_part = True 
if PSTH_part:

    def my_smooth(Arr, bin):
        v = np.ones(bin, dtype=float)/bin
        Arr_conv = np.convolve(Arr, v, mode='same')
        return Arr_conv

    def resolve_sum(Arr1, Arr2):
        length_S = min([Arr1.size, Arr2.size])
        S = Arr1[0:length_S] + Arr2[0:length_S]

        return S

    def Get_subpop(Spk_bin, curr_cells_ids, semipop):

        def multiple_search(cells_ids,curr_cells):
            tot_bool = np.array([False]*len(cells_ids))
            for i in curr_cells:
                tot_bool = tot_bool | [cells_ids == i]

            return tot_bool

        if semipop!=0:
            semipop_mask = multiple_search(curr_cells_ids, semipop)
            new_spk_bin = Spk_bin[semipop_mask[0],:]
        else: 
            new_spk_bin = Spk_bin

        return new_spk_bin
        
    def Get_Saccade_trials(Population_activity, Time_end_saccade, Saccade_time_before, Saccade_time_after,bin_smooth):
        Population_activity_saccade = np.empty([len(Time_end_saccade), Saccade_time_before + Saccade_time_after])
        Population_activity = my_smooth(Population_activity, bin_smooth)
        for i in range(len(Time_end_saccade)):
            Population_activity_saccade[i,:] = Population_activity[int(Time_end_saccade[i]) - Saccade_time_before : int(Time_end_saccade[i]) +Saccade_time_after]
        
        return Population_activity_saccade
    
    

    curr_dir = os.getcwd()


    # semipop = cell_ids['PC_p']
    [new_spk_bin_PC01, cell_ids_PC_01] = Spikes_fun.Get_new_spk_bin(dirname,'PC_01')
    [new_spk_bin_PC02, cell_ids_PC_02] = Spikes_fun.Get_new_spk_bin(dirname,'PC_02')


    new_spk_bin_PC01 = new_spk_bin_PC01*1000.0
    new_spk_bin_PC02 = new_spk_bin_PC02*1000.0


    Time_end_saccade = All_options['Time_end_saccade']


    saccade_time = len(dat['P_NO_CER'])
    
    save_stuff_sync = True
    if save_stuff_sync: 
        all_ids = {
            'all_ids':np.concatenate([cell_ids_PC_01, cell_ids_PC_01], axis=0).tolist(),
            'all_populations':cell_ids,
            'trial_time':trial_time.tolist(),
            'Time_end_saccade':Time_end_saccade,
            }
        spk_to_save = np.concatenate([new_spk_bin_PC01, new_spk_bin_PC02], axis = 0)
        np.savetxt('data_spikes.csv', spk_to_save)
        with open('all_data.json' , 'w') as file1:
            json.dump(all_ids,file1)
    

    with sns.axes_style("darkgrid"):
        ax = f1.add_subplot(gs[0:2, 4:6])
        ax.tick_params(labelsize=label_size_tick)

        Population_activity = resolve_sum(np.mean(new_spk_bin_PC01, axis=0), np.mean(new_spk_bin_PC02, axis=0))/2.0
        Population_activity_saccade = Get_Saccade_trials(Population_activity, Time_end_saccade, Saccade_time_before, Saccade_time_after, bin_smooth)

        Initial_pop_trial = np.mean(Population_activity_saccade[1:n_trials_per_block+1,:], axis=0)
        End_pop_trial = np.mean(Population_activity_saccade[-n_trials_per_block:,:], axis=0)
        ax.plot(trial_time, Initial_pop_trial, linewidth=lin_space_r)
        ax.plot(trial_time, End_pop_trial, linewidth=lin_space_r)

        ax.set_title('c', loc='left', fontsize=title_size , fontweight=fontweight_title, pad=30.0)
        ax.set_ylabel('Simple Spike \n firing rate (Hz)', fontsize=label_size, fontweight=fontweight_label)
        ax.set_xlabel('time (ms)', fontsize=label_size, fontweight=fontweight_label)
        ax.set_xticks(time_ticks)
        plt.legend(('before training','after training','no cerebellum', 'target'), fontsize=legend_size)

    with sns.axes_style("darkgrid"):
        ax = f1.add_subplot(gs[0:2, 0:2])
        ax.tick_params(labelsize=label_size_tick)

        semipop = cell_ids['PC_p']
        spk_bin_PC01_p = Get_subpop(new_spk_bin_PC01, cell_ids_PC_01, semipop)
        spk_bin_PC02_p = Get_subpop(new_spk_bin_PC02, cell_ids_PC_02, semipop)
        Population_activity = resolve_sum(np.mean(spk_bin_PC01_p, axis=0), np.mean(spk_bin_PC02_p, axis=0))/2.0
        Population_activity_saccade = Get_Saccade_trials(Population_activity, Time_end_saccade, Saccade_time_before, Saccade_time_after,bin_smooth)
        

        Initial_pop_trial = np.mean(Population_activity_saccade[1:n_trials_per_block+1,:], axis=0)
        End_pop_trial = np.mean(Population_activity_saccade[-n_trials_per_block:,:], axis=0)
        ax.plot(trial_time, Initial_pop_trial, linewidth=lin_space_r)
        ax.plot(trial_time, End_pop_trial, linewidth=lin_space_r)

        ax.set_title('a', loc='left', fontsize=title_size , fontweight=fontweight_title, pad=30.0)
        ax.set_ylabel('Simple Spike \n firing rate (Hz)', fontsize=label_size, fontweight=fontweight_label)
        ax.set_xlabel('time (ms)', fontsize=label_size, fontweight=fontweight_label)
        ax.set_xticks(time_ticks)
        # plt.legend(('cerebellum before training','cerebellum after training','no cerebellum', 'target'), fontsize=legend_size)

    with sns.axes_style("darkgrid"):
        ax = f1.add_subplot(gs[0:2, 2:4])
        ax.tick_params(labelsize=label_size_tick)

        semipop = cell_ids['PC_n']
        spk_bin_PC01_p = Get_subpop(new_spk_bin_PC01, cell_ids_PC_01, semipop)
        spk_bin_PC02_p = Get_subpop(new_spk_bin_PC02, cell_ids_PC_02, semipop)
        Population_activity = resolve_sum(np.mean(spk_bin_PC01_p, axis=0), np.mean(spk_bin_PC02_p, axis=0))/2.0
        Population_activity_saccade = Get_Saccade_trials(Population_activity, Time_end_saccade, Saccade_time_before, Saccade_time_after,bin_smooth)
        

        Initial_pop_trial = np.mean(Population_activity_saccade[1:n_trials_per_block+1,:], axis=0)
        End_pop_trial = np.mean(Population_activity_saccade[-n_trials_per_block:,:], axis=0)
        ax.plot(trial_time, Initial_pop_trial, linewidth=lin_space_r)
        ax.plot(trial_time, End_pop_trial, linewidth=lin_space_r)

        ax.set_title('b', loc='left', fontsize=title_size , fontweight=fontweight_title, pad=30.0)
        ax.set_ylabel('Simple Spike \n firing rate (Hz)', fontsize=label_size, fontweight=fontweight_label)
        ax.set_xlabel('time (ms)', fontsize=label_size, fontweight=fontweight_label)
        ax.set_xticks(time_ticks)
        # plt.legend(('cerebellum before training','cerebellum after training','no cerebellum', 'target'), fontsize=legend_size)

# f1.set_tight_layout(True)
legend_size = 75
kin_part = True
if kin_part:

    n_trials = 10
    line_wd = 2.0
    
    """def resolve_array(V):
        Vmax = np.zeros(len(V))
        for i in range(len(V)):
            Vmax[i] = len(V[i])
        V_trials = np.zeros(Vmax.max()) """

    V_trials = np.array(dat['Velocity_eye_tot'])[1:,:]
    Error = dat['Error'][1:]
    Vmax_trials = np.max(V_trials, axis=1)[0:]

    def GetTarray(array, trial_time, t_ref, last_point = 0):
        Tarray = np.zeros(trial_time.shape[0])
        Tarray[t_ref:array.shape[0]+t_ref] = array
        Tarray[array.shape[0]+t_ref:] = last_point
        return Tarray   


    # trial_time = np.linspace(0,V_trials.shape[1]-1,V_trials.shape[1])
    all_trials = np.linspace(0,V_trials.shape[0]-1,V_trials.shape[0])

    P_trials = np.array(dat['Position_eye_tot'])[1:]
    Pmax_trials = np.max(P_trials, axis=1)


    with sns.axes_style("darkgrid"):
        ax = f1.add_subplot(gs[2:4, 0:3])
        ax.tick_params(labelsize=label_size_tick)
        
        first_trials = np.mean(V_trials[0:n_trials,:], axis=0)
        last_trials = np.mean(V_trials[-n_trials:,:], axis=0)
        no_cerebellum = np.array(dat['V_NO_CER'])

        

        ax.plot(trial_time, GetTarray(first_trials, trial_time, t_ref), linewidth=lin_space_r)
        ax.plot(trial_time, GetTarray(last_trials, trial_time, t_ref), linewidth=lin_space_r)
        ax.plot(trial_time, GetTarray(no_cerebellum, trial_time, t_ref), linewidth=lin_space_r, linestyle='dashed', color=[0,0,0])
        ax.set_title('d', loc='left', fontsize=title_size , fontweight=fontweight_title, pad=30.0)
        ax.set_ylabel('speed (deg/sec)', fontsize=label_size, fontweight=fontweight_label)
        ax.set_xlabel('time (ms)', fontsize=label_size, fontweight=fontweight_label)
        ax.set_xticks(time_ticks)
        plt.legend(('before training','after training','no cerebellum', 'target'), fontsize=legend_size)

        

        ax = f1.add_subplot(gs[2:4, 3:6])
        
        first_trials = np.mean(P_trials[0:n_trials,:], axis=0)
        last_trials = np.mean(P_trials[-n_trials:,:], axis=0)
        no_cerebellum = np.array(dat['P_NO_CER'])


        ax.plot(trial_time, GetTarray(first_trials, trial_time, t_ref, last_point = first_trials[-1]), linewidth=lin_space_r)
        ax.plot(trial_time, GetTarray(last_trials, trial_time, t_ref,  last_point = last_trials[-1]), linewidth=lin_space_r)
        ax.plot(trial_time, GetTarray(no_cerebellum, trial_time, t_ref,  last_point = no_cerebellum[-1]), linewidth=lin_space_r, linestyle='dashed', color=[0,0,0])
        rect = patches.Rectangle((trial_time[0], 9), len(trial_time)-1, 2, linewidth=1, edgecolor=None, facecolor='pink', alpha=0.7)
        ax.add_patch(rect)
        ax.set_title('e', loc='left', fontsize=title_size , fontweight=fontweight_title, pad=30.0)
        ax.set_xlabel('time (ms)', fontsize=label_size, fontweight=fontweight_label)
        ax.set_xticks(time_ticks)
        ax.set_ylabel('position (deg)', fontsize=label_size, fontweight=fontweight_label)
        ax.tick_params(labelsize=label_size_tick)
        plt.legend(('before training','after training','no cerebellum', 'target'), fontsize=legend_size)

        kin_data = {
            'before_training': first_trials.tolist(),
            'after_training' : last_trials.tolist(),
            'no_cerebellum': no_cerebellum.tolist()}

        with open('kin_data.txt', 'w') as f:
            json.dump(kin_data, f)



    with sns.axes_style("darkgrid"):
        ax = f1.add_subplot(gs[4:6, 0:3])
        ax.tick_params(labelsize=label_size_tick)

        ax.plot(all_trials, Vmax_trials, linewidth=lin_space_r_training, color=[0.1, 0.1, 0.1] )
        print(Vmax_trials[-10:].mean())
        ax.set_title('f', loc='left', fontsize=title_size , fontweight=fontweight_title, pad=30.0)
        ax.set_xlabel('learning trials', fontsize=label_size, fontweight=fontweight_label)
        ax.set_ylabel('peack speed \n (deg/sec)', fontsize=label_size, fontweight=fontweight_label)
        
        ax = f1.add_subplot(gs[4:6, 3:6])

        ax.tick_params(labelsize=label_size_tick)
        ax.plot(all_trials, Error, linewidth=lin_space_r_training, color=[0.1, 0.1, 0.1] )
        ax.set_title('g', loc='left', fontsize=title_size , fontweight=fontweight_title, pad=30.0)
        ax.set_xlabel('learning trials', fontsize=label_size, fontweight=fontweight_label)
        ax.set_ylabel('foveal error (deg)', fontsize=label_size, fontweight=fontweight_label)
        
        


raster_part = True
if raster_part:


    def my_smooth(Arr, bin):
        v = np.ones(bin, dtype=float)/bin
        Arr_conv = np.convolve(Arr, v, mode='same')
        return Arr_conv

    def resolve_sum(Arr1, Arr2):
        length_S = min([Arr1.size, Arr2.size])
        S = Arr1[0:length_S] + Arr2[0:length_S]

        return S

    def resolve_concatenate(Arr1, Arr2, dim=0):
        Shape_dim = [1, 0]
        length_S = min([Arr1.shape[Shape_dim[dim]], Arr2.shape[Shape_dim[dim]]])
        S = np.concatenate([Arr1[:,0:length_S], Arr2[:,0:length_S]])

        return S

    def Get_subpop(Spk_bin, curr_cells_ids, semipop):

        def multiple_search(cells_ids,curr_cells):
            tot_bool = np.array([False]*len(cells_ids))
            for i in curr_cells:
                tot_bool = tot_bool | [cells_ids == i]

            return tot_bool

        if semipop!=0:
            semipop_mask = multiple_search(curr_cells_ids, semipop)
            new_spk_bin = Spk_bin[semipop_mask[0],:]
        else: 
            new_spk_bin = Spk_bin

        return new_spk_bin
        
    def Get_Saccade_trials_2D(Population_activity, Time_end_saccade, Saccade_time_before, Saccade_time_after):
        Population_activity_saccade = np.empty([Population_activity.shape[0], len(Time_end_saccade), Saccade_time_before + Saccade_time_after])
        # Population_activity = my_smooth(Population_activity, bin_smooth)
        for i in range(len(Time_end_saccade)):
            Population_activity_saccade[:,i,:] = Population_activity[:, int(Time_end_saccade[i]) - Saccade_time_before : int(Time_end_saccade[i]) +Saccade_time_after]
        
        return np.array(Population_activity_saccade)

    def From_array01_2_eveplot(Arr, t_ref):
        N = Arr.shape[0]
        data = [None] * N
        for i in range(N):
            data[i] = np.where(Arr[i])[0] + t_ref

        return data

    def Rearrange_array01(Arr):
        n_events_per_cells = np.sum(Arr, axis = 1)
        idx = np.argsort(n_events_per_cells)
        Arr_rearragend = Arr[idx,:]
        return Arr_rearragend
        
    
    N_MLI = 50
    saccade_time = len(dat['P_NO_CER'])

    # semipop = cell_ids['PC_p']
    [new_spk_bin_PC01, cell_ids_PC_01] = Spikes_fun.Get_new_spk_bin(dirname,'PC_01')
    [new_spk_bin_PC02, cell_ids_PC_02] = Spikes_fun.Get_new_spk_bin(dirname,'PC_02')
    
    [new_spk_bin_BC, cell_ids_BC] = Spikes_fun.Get_new_spk_bin(dirname,'_BC')
    [new_spk_bin_SC, cell_ids_SC] = Spikes_fun.Get_new_spk_bin(dirname,'_SC')

    [new_spk_bin_GR, cell_ids_GR] = Spikes_fun.Get_new_spk_bin(dirname,'_GR')
    [new_spk_bin_MF, cell_ids_MF] = Spikes_fun.Get_new_spk_bin(dirname,'_MF')


    new_spk_bin_PC01_p = Get_subpop(new_spk_bin_PC01, cell_ids_PC_01, cell_ids['PC_p'])
    new_spk_bin_PC02_p = Get_subpop(new_spk_bin_PC02, cell_ids_PC_02, cell_ids['PC_p'])

    new_spk_bin_PC01_n = Get_subpop(new_spk_bin_PC01, cell_ids_PC_01, cell_ids['PC_n'])
    new_spk_bin_PC02_n = Get_subpop(new_spk_bin_PC02, cell_ids_PC_02, cell_ids['PC_n'])


    Time_end_saccade = All_options['Time_end_saccade'][1:]


    Time_end_saccade = All_options['Time_end_saccade'][1:]

    new_spk_bin_p = resolve_concatenate(new_spk_bin_PC01_p, new_spk_bin_PC02_p)
    Population_activity_saccade_p = Get_Saccade_trials_2D(new_spk_bin_p, Time_end_saccade, Saccade_time_before, Saccade_time_after)[:,1:,:]
    # Population_activity_saccade_p = np.array(Population_activity_saccade_p)

    new_spk_bin_n = resolve_concatenate(new_spk_bin_PC01_n, new_spk_bin_PC02_n)
    Population_activity_saccade_n = Get_Saccade_trials_2D(new_spk_bin_n, Time_end_saccade, Saccade_time_before, Saccade_time_after)[:,1:,:]
    # Population_activity_saccade_n = np.array(Po

    new_spk_bin_MLI = resolve_concatenate(new_spk_bin_BC[0:int(N_MLI/2),:], new_spk_bin_SC[0:int(N_MLI/2),:])
    Population_activity_saccade_MLI = Get_Saccade_trials_2D(new_spk_bin_MLI, Time_end_saccade, Saccade_time_before, Saccade_time_after)[:,1:,:]

    Population_activity_saccade_GR = Get_Saccade_trials_2D(new_spk_bin_GR, Time_end_saccade, Saccade_time_before, Saccade_time_after)[:,1:,:]
    Population_activity_saccade_MF = Get_Saccade_trials_2D(new_spk_bin_MF, Time_end_saccade, Saccade_time_before, Saccade_time_after)[:,1:,:]
    
    Color_PC_pause = '#783800'
    Color_PC_burst = '#004166'
    Color_BC = '#126E27'
    Color_SC = '#82E670'
    Color_GR = '#6E5239'
    Color_MF = '#3B3B3B'

    color_movement_start = '#E6852C'
    lin_space_r = 4.0
    


    with sns.axes_style("white"):

        ax = f1.add_subplot(gs[:, 6:])
        
        Initial_pop_trial_p = Rearrange_array01(Population_activity_saccade_p[:,1,:])
        Initial_pop_trial_n = Rearrange_array01(Population_activity_saccade_n[:,1,:])

        Initial_pop_trial_GR = Rearrange_array01(Population_activity_saccade_GR[:,1,:])
        Initial_pop_trial_MLI= Rearrange_array01(Population_activity_saccade_MLI[:,1,:])

        Initial_pop_trial_MF = Rearrange_array01(Population_activity_saccade_MF[:,1,:])
        

        colors_pause = [[Color_PC_pause] for i in range(Initial_pop_trial_p.shape[0])]
        colors_burst = [[Color_PC_burst] for i in range(Initial_pop_trial_n.shape[0])]
        colors_BC = [[Color_BC] for i in range(int(N_MLI/2))]
        colors_SC = [[Color_SC] for i in range(int(N_MLI/2))]
        colors_GR = [[Color_GR] for i in range(Initial_pop_trial_GR.shape[0])]
        colors_MF = [[Color_MF] for i in range(Initial_pop_trial_MF.shape[0])]
        

        Initial_pop_trial_p = Rearrange_array01(Population_activity_saccade_p[:,-1,:])
        Initial_pop_trial_n = Rearrange_array01(Population_activity_saccade_n[:,-1,:])

        Initial_pop_trial_GR = Rearrange_array01(Population_activity_saccade_GR[:,-1,:])
        Initial_pop_trial_MLI= Rearrange_array01(Population_activity_saccade_MLI[:,-1,:])

        Initial_pop_trial_MF = Rearrange_array01(Population_activity_saccade_MF[:,-1,:])
        

        colors_pause = [[Color_PC_pause] for i in range(Initial_pop_trial_p.shape[0])]
        colors_burst = [[Color_PC_burst] for i in range(Initial_pop_trial_n.shape[0])]
        colors_BC = [[Color_BC] for i in range(int(N_MLI/2))]
        colors_SC = [[Color_SC] for i in range(int(N_MLI/2))]
        colors_GR = [[Color_GR] for i in range(Initial_pop_trial_GR.shape[0])]
        colors_MF = [[Color_MF] for i in range(Initial_pop_trial_MF.shape[0])]
        

        colors = colors_MF + colors_GR + colors_SC + colors_BC + colors_pause + colors_burst
        Initial_pop_trial = resolve_concatenate(Initial_pop_trial_MF, Initial_pop_trial_GR)
        Initial_pop_trial = resolve_concatenate(Initial_pop_trial, Initial_pop_trial_MLI)
        Initial_pop_trial = resolve_concatenate(Initial_pop_trial, Initial_pop_trial_p)
        Initial_pop_trial = resolve_concatenate(Initial_pop_trial, Initial_pop_trial_n)

        # Initial_pop_trial = Initial_pop_trial - t_ref
        Initial_pop_trial_list = From_array01_2_eveplot(Initial_pop_trial, -t_ref)
        lineoffsets_eventplot = 1.0
        ax.eventplot(Initial_pop_trial_list, colors=colors, linelengths=0.5,linewidths = 15.0, lineoffsets=lineoffsets_eventplot)

        # TEXT
        text_x = trial_time[-1] + 15 
        text_kwargs = dict(ha='center', va='center', fontsize=60, color='k', fontweight='bold')
        curr_point_text = 0

        # MF 
        text_y = curr_point_text + len(colors_MF)/(1/(lineoffsets_eventplot)*2)
        ax.text(text_x, text_y,'MF', **text_kwargs)
        curr_point_text +=  len(colors_MF)/lineoffsets_eventplot

        # GR
        text_y = curr_point_text + len(colors_GR)/(1/(lineoffsets_eventplot)*2)
        ax.text(text_x, text_y,'GrC', **text_kwargs)
        curr_point_text +=  len(colors_GR)/lineoffsets_eventplot

        # SC
        text_y = curr_point_text + len(colors_SC)/(1/(lineoffsets_eventplot)*2)
        ax.text(text_x, text_y,'SC', **text_kwargs)
        curr_point_text +=  len(colors_SC)/lineoffsets_eventplot

        # BC
        text_y = curr_point_text + len(colors_BC)/(1/(lineoffsets_eventplot)*2)
        ax.text(text_x, text_y,'BC', **text_kwargs)
        curr_point_text +=  len(colors_BC)/lineoffsets_eventplot

        # PC busrt
        text_y = curr_point_text + len(colors_pause)/(1/(lineoffsets_eventplot)*2)
        ax.text(text_x, text_y,'burst PC', **text_kwargs)
        curr_point_text +=  len(colors_pause)/lineoffsets_eventplot

        # PC pause
        text_y = curr_point_text + len(colors_burst)/(1/(lineoffsets_eventplot)*2)
        ax.text(text_x, text_y,'pause PC', **text_kwargs)
        curr_point_text +=  len(colors_burst)/lineoffsets_eventplot

        # ax.axvline(x = 0.0, linewidth=lin_space_r*2, linestyle='dashed', color=color_movement_start)
        ax.axvline(x =  trial_time[-1], linewidth=0.5, linestyle='solid', color='k')
        
        

        ax.set_title('h', loc='left', fontsize=title_size , fontweight=fontweight_title, pad=30.0)
        rect = patches.Rectangle((0, 0), len(no_cerebellum)-1, len(Initial_pop_trial_list),
            linewidth=1, edgecolor=None, facecolor='black', alpha=0.1)
        ax.add_patch(rect)

        # ax.set_xlim( t_ref -10, Saccade_time_before - t_ref + Saccade_time_after)
        ax.set_xlim( trial_time[0], trial_time[-1] + 30)
        ax.set_ylim( 0, len(Initial_pop_trial_list))
        ax.set_ylabel('Cells ids', fontsize=label_size, fontweight=fontweight_label)
        ax.set_xlabel('time (ms)', fontsize=label_size, fontweight=fontweight_label)
        
        ax.set_yticks([])
        xticks_raster = [-25, 0, 25, 50, 75]
        ax.set_xticks(time_ticks)
        ax.tick_params(labelsize=label_size_tick)
        # plt.legend(('cerebellum before training','cerebellum after training','no cerebellum', 'target'), fontsize=legend_size)
        ax2 = ax.twinx()
        ax2.set_ylabel('Cells types', fontsize=label_size, fontweight=fontweight_label)

# f1.set_tight_layout(True)
f1.tight_layout(pad = 10.0)
dirname = 'p1'
if not os.path.exists('Plots/' + dirname ):
    os.makedirs('Plots/' + dirname )

f1.savefig('Plots/' + dirname + '/panel1.svg', format='svg', dpi=300)
f1.savefig('Plots/' + dirname + '/panel1.png', format='png', dpi=300)

import pickle
pickle.dump(f1, open(('Plots/p1/panel1.p'), 'wb'))
