{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import h5py\n",
    "import pandas as pd\n",
    "import numpy as np\n",
    "import os\n",
    "from matplotlib import pyplot as plt\n",
    "\n",
    "import math\n",
    "import warnings\n",
    "import scipy.optimize\n",
    "import scipy.stats\n",
    "import sklearn\n",
    "from sklearn.svm import LinearSVC\n",
    "from sklearn.decomposition import PCA\n",
    "from scipy.stats import wilcoxon\n",
    "\n",
    "from scipy.spatial import distance\n",
    "from scipy.linalg import norm\n",
    "\n",
    "from sklearn.metrics import balanced_accuracy_score\n",
    "from sklearn.discriminant_analysis import LinearDiscriminantAnalysis as LDA\n",
    "from sklearn.linear_model import LogisticRegression"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import matplotlib as mpl\n",
    "mpl.rcParams['pdf.fonttype'] = 42"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "%load_ext autoreload\n",
    "%autoreload 2"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import decoding_utils as du"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "#Paths to all of the useful supplemental tables and tensors\n",
    "active_tensor_file = \"/Volumes/programs/mindscope/workgroups/np-behavior/vbn_data_release/supplemental_tables/vbnAllUnitSpikeTensor.hdf5\"\n",
    "\n",
    "stim_table_file = \"/Volumes/programs/mindscope/workgroups/np-behavior/vbn_data_release/supplemental_tables/master_stim_table_no_filter.csv\"\n",
    "unit_table_file = \"/Volumes/programs/mindscope/workgroups/np-behavior/vbn_data_release/supplemental_tables/master_units_with_responsiveness.csv\"\n",
    "\n",
    "sessions_table_file = \"/Volumes/programs/mindscope/workgroups/np-behavior/vbn_data_release/vbn_s3_cache/visual-behavior-neuropixels-0.5.0/project_metadata/ecephys_sessions.csv\""
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "stim_table = pd.read_csv(stim_table_file)\n",
    "unit_table = pd.read_csv(unit_table_file)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "unitData = h5py.File(active_tensor_file)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "sessions = pd.read_csv(sessions_table_file)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from notebook_utils import (\n",
    "    make_mono_colormap,\n",
    "    mean_paired_image_mat_from_stim_table,\n",
    "    beh_mat_from_stim_table,\n",
    "    skip_diag_masking,\n",
    "    get_session_engaged_dprime,\n",
    "    get_session_engaged_hit_count,\n",
    "    get_experience_session_id_for_mouse,\n",
    "    get_omission_mean,\n",
    "    get_post_omission_mean,\n",
    "    get_shared_change_mean,\n",
    "    get_private_change_mean,\n",
    "    get_shared_catch_mean,\n",
    "    get_private_catch_mean,\n",
    "    get_shared_nonchange_mean,\n",
    "    get_private_nonchange_mean,\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "save_dir = \"/Volumes/programs/mindscope/workgroups/np-behavior/vbn_data_release/supplemental_tables/flash_decoding_metrics\""
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "regions = ('VISall', 'VISp', 'VISl', 'VISrl', 'VISal', 'VISpm', 'VISam', 'LGd', 'LP', 'SCMRN', 'Hipp')"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Generate flash metrics"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "def decodeSingleClass(sp, y, nCrossVal, model, unitSampleSize):\n",
    "    '''\n",
    "    sp: units x trials\n",
    "    y: labels \n",
    "    nCrossVal: number of cross validation splits to run\n",
    "    '''\n",
    "    nUnits = sp.shape[0]\n",
    "    warnings.filterwarnings('ignore')\n",
    "    summary = {sampleSize: \n",
    "                {metric: [] for metric in  ('trainAccuracy',\n",
    "                                            'featureWeights',\n",
    "                                            'accuracy',\n",
    "                                            'prediction',\n",
    "                                            'confidence', \n",
    "                                            'balanced_accuracy')}\n",
    "                for sampleSize in unitSampleSize}\n",
    "    \n",
    "    for sampleSize in unitSampleSize:\n",
    "        if nUnits < sampleSize:\n",
    "            continue\n",
    "        if sampleSize>1:\n",
    "            if sampleSize==nUnits:\n",
    "                nSamples = 1\n",
    "                unitSamples = [np.arange(nUnits)]\n",
    "            else:\n",
    "                # >99% chance each neuron is chosen at least once\n",
    "                nSamples = int(math.ceil(math.log(0.01)/math.log(1-sampleSize/nUnits)))\n",
    "                unitSamples = [np.random.choice(nUnits,sampleSize,replace=False) for _ in range(nSamples)]\n",
    "        elif sampleSize==1:\n",
    "            nSamples = nUnits\n",
    "            unitSamples = [[i] for i in range(nUnits)]\n",
    "        else: #if sampleSize<1, just run 1 iteration with all the units available\n",
    "            nSamples = 1\n",
    "            unitSamples = [np.arange(nUnits)]\n",
    "\n",
    "        for metric in summary[sampleSize]:\n",
    "            summary[sampleSize][metric].append([])\n",
    "\n",
    "        for unitSamp in unitSamples:\n",
    "            cv = du.trainDecoder(model,sp[unitSamp].T,y,nCrossVal)\n",
    "            summary[sampleSize]['trainAccuracy'][-1].append(np.mean(cv['train_score']))\n",
    "            summary[sampleSize]['featureWeights'][-1].append(np.mean(cv['coef'],axis=0).squeeze())\n",
    "            summary[sampleSize]['accuracy'][-1].append(np.mean(cv['test_score']))\n",
    "            summary[sampleSize]['prediction'][-1].append(cv['predict'])\n",
    "            summary[sampleSize]['confidence'][-1].append(cv['decision_function'])\n",
    "            summary[sampleSize]['balanced_accuracy'][-1].append(sklearn.metrics.balanced_accuracy_score(y.astype(bool), cv['predict'].astype(bool)))\n",
    "\n",
    "        for metric in summary[sampleSize]:\n",
    "            if metric == 'prediction':\n",
    "                summary[sampleSize][metric][-1] = scipy.stats.mode(summary[sampleSize][metric][-1],axis=0)[0][0]\n",
    "            elif metric != 'featureWeights':\n",
    "                summary[sampleSize][metric][-1] = np.median(summary[sampleSize][metric][-1],axis=0)\n",
    "    \n",
    "    warnings.filterwarnings('default')\n",
    "    return summary"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "unitSampleSize = [20,40,-1]\n",
    "unitSampleNames = [u if u>=1 else 'all' for u in unitSampleSize]\n",
    "respWin = slice(20,100)\n",
    "\n",
    "for isess, session in sessions.iterrows():\n",
    "    print(isess)\n",
    "\n",
    "    sessionId = session['ecephys_session_id']\n",
    "\n",
    "    stim = stim_table[(stim_table['session_id']==sessionId) & stim_table['active']].reset_index()\n",
    "    im_ids = stim['image_name'].values\n",
    "    previous_im_ids = np.insert(stim['image_name'].values, 0, 'omitted')[:-1]\n",
    "    previous_im_ids_skipping_omitted = []\n",
    "    for ind, id in enumerate(previous_im_ids):\n",
    "        if not id=='omitted':\n",
    "            previous_im_ids_skipping_omitted.append(id)\n",
    "        else:\n",
    "            if ind==0:\n",
    "                previous_im_ids_skipping_omitted.append(np.nan)\n",
    "            else:\n",
    "                previous_im_ids_skipping_omitted.append(previous_im_ids[ind-1])\n",
    "\n",
    "    previous_im_ids_skipping_omitted = np.array(previous_im_ids_skipping_omitted)\n",
    "\n",
    "    stim['previous_image_skipping_omitted'] = previous_im_ids_skipping_omitted\n",
    "    unique_images = [im for im in np.unique(stim['image_name'].values) if im not in ['im083_r', 'im111_r', 'omitted']] + ['im083_r', 'im111_r']\n",
    "\n",
    "    units = unit_table.set_index('unit_id').loc[unitData[str(sessionId)]['unitIds'][:]]\n",
    "    spikes = unitData[str(sessionId)]['spikes']\n",
    "    highQuality = du.apply_unit_quality_filter(units, no_abnorm=False)\n",
    "\n",
    "    for region in regions:\n",
    "        inRegion = du.getUnitsInRegion(units,region)\n",
    "        final_unit_filter = highQuality & inRegion\n",
    "        \n",
    "        cols_to_add = [f'{metric}_{region}_{unitSampleName}' for unitSampleName in unitSampleNames for metric in ['previous_image_confidence', 'change_confidence']]\n",
    "        stim[cols_to_add] = np.nan\n",
    "        \n",
    "        if np.sum(final_unit_filter)<10:\n",
    "            continue\n",
    "\n",
    "        sp = np.zeros((final_unit_filter.sum(),spikes.shape[1],spikes.shape[2]),dtype=bool)\n",
    "        for i,u in enumerate(np.where(final_unit_filter)[0]):\n",
    "            sp[i]=spikes[u,:,:]\n",
    "\n",
    "        flash_response = sp[:, :, respWin].mean(axis=2)\n",
    "\n",
    "        image_decoder_summary = {}\n",
    "        for im in unique_images:\n",
    "            y = im_ids == im\n",
    "            model = LinearSVC(C=1.0,max_iter=int(1e4),class_weight='balanced')\n",
    "            results = decodeSingleClass(flash_response, y, 5, model, unitSampleSize)\n",
    "            image_decoder_summary[im] = results\n",
    "\n",
    "        model = LinearSVC(C=1.0,max_iter=int(1e4), class_weight='balanced')\n",
    "        change_results = decodeSingleClass(flash_response, stim['is_change'].values, 5, model, unitSampleSize)\n",
    "\n",
    "        for unitSample, unitSampleName in zip(unitSampleSize, unitSampleNames):\n",
    "            for i, row in stim.iterrows():\n",
    "                previous_stim_id = row['previous_image_skipping_omitted']\n",
    "                if previous_stim_id in image_decoder_summary:\n",
    "                    previous_results = image_decoder_summary[previous_stim_id]\n",
    "                    if len(previous_results[unitSample]['confidence'])>0:\n",
    "                        confidence = previous_results[unitSample]['confidence'][0][i]\n",
    "                        stim.at[i, f'previous_image_confidence_{region}_{unitSampleName}'] = confidence\n",
    "\n",
    "            if len(change_results[unitSample]['confidence'])>0:\n",
    "                stim[f'change_confidence_{region}_{unitSampleName}'] = change_results[unitSample]['confidence'][0]\n",
    "\n",
    "    cols_to_save = ['session_id', 'is_change', 'previous_image_skipping_omitted', 'image_name'] + \\\n",
    "                    [f'previous_image_confidence_{region}_{unitSampleName}' for region in regions for unitSampleName in unitSampleNames] + \\\n",
    "                    [f'change_confidence_{region}_{unitSampleName}' for region in regions for unitSampleName in unitSampleNames]\n",
    "    stim[cols_to_save].to_csv(os.path.join(save_dir, f'{sessionId}_responseWin_20to100.csv'), index=False)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "save_dir = \"/Volumes/programs/mindscope/workgroups/np-behavior/vbn_data_release/supplemental_tables/flash_decoding_metrics\"\n",
    "\n",
    "stim_files = [os.path.join(save_dir, f) for f in os.listdir(save_dir) if 'responseWin_20to100' in f]\n",
    "dfs = []\n",
    "for stim_file in stim_files:\n",
    "    df = pd.read_csv(stim_file)\n",
    "    stim = stim_table[stim_table['session_id']==int(os.path.basename(stim_file).split('.')[0].split('_')[0])]\n",
    "    df = df.merge(stim.reset_index(), left_index=True, right_index=True, suffixes=('', '_stimtable'))     \n",
    "    \n",
    "    dfs.append(df)\n",
    "\n",
    "stimtable_with_flash_metrics = pd.concat(dfs)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "stimtable_with_flash_metrics = stimtable_with_flash_metrics.drop(columns=[c for c in stimtable_with_flash_metrics if 'Unnamed' in c])\n",
    "stimtable_with_flash_metrics.head()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "def trainModel(model,X,y,nSplits):\n",
    "    classVals = np.unique(y)\n",
    "    nClasses = len(classVals)\n",
    "    nSamples = len(y)\n",
    "    cv = {'estimator': [sklearn.base.clone(model) for _ in range(nSplits)]}\n",
    "    cv['train_balanced_accuracy'] = []\n",
    "    cv['test_balanced_accuracy'] = []\n",
    "    cv['predict'] = np.full(nSamples, '', dtype='O')\n",
    "    cv['predict_proba'] = np.full((nSamples,nClasses),np.nan)\n",
    "    cv['coef'] = []\n",
    "    modelMethods = dir(model)\n",
    "    trainInd,testInd = du.getTrainTestSplits(y,nSplits,hasClasses=False)\n",
    "    for num_iter, (estimator,train,test) in enumerate(zip(cv['estimator'],trainInd,testInd)):\n",
    "        estimator.fit(X[train],y[train])\n",
    "        cv['train_balanced_accuracy'].append(balanced_accuracy_score(y[train], estimator.predict(X[train])))\n",
    "        cv['test_balanced_accuracy'].append(balanced_accuracy_score(y[test], estimator.predict(X[test])))\n",
    "        cv['predict'][test] = estimator.predict(X[test])\n",
    "        for method in ('predict_proba',):\n",
    "            if method in modelMethods:\n",
    "                cv[method][test] = getattr(estimator,method)(X[test])\n",
    "        for attr in ('coef_',):\n",
    "            if attr in estimator.__dict__:\n",
    "                cv[attr[:-1]].append(getattr(estimator,attr))\n",
    "    return cv"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "session_model_results = {s:{} for s in stimtable_with_flash_metrics['session_id'].unique()}\n",
    "areas_to_run = ('VISall', 'VISp', 'VISl', 'VISrl', 'VISal', 'VISpm', 'VISam', 'LGd', 'LP', 'SCMRN', 'Hipp')\n",
    "metrics_to_run = ['previous_image_confidence', 'change_confidence']\n",
    "samplesizes = ['20', '40', 'all']\n",
    "\n",
    "columns_to_predict = ['is_change', 'lickbout_for_flash_during_response_window']\n",
    "filters = [lambda x: [True]*len(x), lambda x: ~x['is_change']]\n",
    "\n",
    "warnings.filterwarnings(\"ignore\")\n",
    "for isess, session in enumerate(stimtable_with_flash_metrics['session_id'].unique()):\n",
    "    if isess%10==0:\n",
    "        print(isess)\n",
    "    stim = stimtable_with_flash_metrics[stimtable_with_flash_metrics['session_id']==session]\n",
    "    stim = stim[stim['engaged'] & \\\n",
    "                stim['no_abnorm'] & \\\n",
    "                ~stim['grace_period_after_hit']]\n",
    "                \n",
    "    for samplesize in samplesizes:\n",
    "        for area in areas_to_run:\n",
    "            for metrics in metrics_to_run:\n",
    "                \n",
    "                if isinstance(metrics, str):\n",
    "                        metrics = metrics + '_' + area + '_' + samplesize if not 'lick' in metrics else metrics\n",
    "                        metrics_name = metrics\n",
    "                else:\n",
    "                    metrics = [m+'_'+area if not 'lick' in m else m for m in metrics]\n",
    "                    metrics_name = 'combo'\n",
    "                    for m in metrics_to_run:\n",
    "                        metrics_name = metrics_name + '__' + m\n",
    "                \n",
    "                for column_to_predict, filter in zip(columns_to_predict, filters):\n",
    "                    \n",
    "                    curated_table = stim[filter(stim)]\n",
    "                    if len(curated_table)==0:\n",
    "                        continue\n",
    "\n",
    "                    X = curated_table[metrics].to_numpy().astype(float)\n",
    "                    no_nan_inds = [irow for irow, row in enumerate(X) if not np.any(np.isnan(row))]\n",
    "                    if len(no_nan_inds)==0:\n",
    "                        session_model_results[session].update({'train_balanced_accuracy' + '_' + metrics_name: np.nan, \n",
    "                                                                'test_balanced_accuracy' + '_' + metrics_name: np.nan})\n",
    "                        continue\n",
    "                    #X_nonan = X[no_nan_inds]\n",
    "                    X_nonan = X[no_nan_inds].reshape(-1,1)\n",
    "\n",
    "                    X_nonan = (X_nonan - X_nonan.mean(axis=0))/X_nonan.std(axis=0)\n",
    "                    y = curated_table[column_to_predict].values\n",
    "                    y_nonan = y[no_nan_inds]\n",
    "                    model = LogisticRegression(random_state=0, solver='liblinear', class_weight='balanced')\n",
    "                    res = trainModel(model, X_nonan, y_nonan, 5,)\n",
    "\n",
    "                    session_model_results[session].update({f'{column_to_predict}_train_balanced_accuracy_{metrics_name}': np.mean(res['train_balanced_accuracy']),\n",
    "                                                            f'{column_to_predict}_test_balanced_accuracy_{metrics_name}': np.mean(res['test_balanced_accuracy'])})"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "sess_change_df = pd.DataFrame.from_dict(session_model_results, orient='index')\n",
    "sess_change_df = sess_change_df.merge(sessions, left_index=True, right_on='ecephys_session_id').set_index('ecephys_session_id')"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from analysis_utils import multiple_comparisons, comparison_matrix"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "areas_to_run = ('VISall', 'VISp', 'VISl', 'VISrl', 'VISal', 'VISpm', 'VISam', 'LGd', 'LP', 'SCMRN', 'Hipp')\n",
    "\n",
    "cols = ['is_change', 'lickbout_for_flash_during_response_window']\n",
    "samplesize = 'all'\n",
    "area_pvalues_change_vs_image = {}\n",
    "area_change_conf_values = {c:{} for c in cols}\n",
    "area_previous_image_conf_values = {c:{} for c in cols}\n",
    "for col in cols:\n",
    "    fig, ax = plt.subplots()\n",
    "    fig.suptitle(col)\n",
    "    for ia, area in enumerate(areas_to_run):\n",
    "        previous_image_conf = sess_change_df[col+'_test_balanced_accuracy_previous_image_confidence_' + area + '_' + samplesize].agg(['mean', 'sem'])\n",
    "        change_conf = sess_change_df[col+'_test_balanced_accuracy_change_confidence_' + area + '_' + samplesize].agg(['mean', 'sem'])\n",
    "\n",
    "        previous_image_conf_values = sess_change_df[col+'_test_balanced_accuracy_previous_image_confidence_' + area + '_' + samplesize].dropna()\n",
    "        change_conf_values = sess_change_df[col+'_test_balanced_accuracy_change_confidence_' + area + '_' + samplesize].dropna()\n",
    "        area_change_conf_values[col][area] = change_conf_values\n",
    "        area_previous_image_conf_values[col][area] = previous_image_conf_values\n",
    "        area_pvalues_change_vs_image[area] = wilcoxon(previous_image_conf_values, change_conf_values, nan_policy='omit')[1]\n",
    "\n",
    "        ax.plot(ia, previous_image_conf['mean'], 'ko')\n",
    "        ax.errorbar(ia, previous_image_conf['mean'], previous_image_conf['sem'], color='k')\n",
    "        \n",
    "        ax.plot(ia, change_conf['mean'], 'ro')\n",
    "        ax.errorbar(ia, change_conf['mean'], change_conf['sem'], color='r')\n",
    "\n",
    "    ax.set_xticks(np.arange(len(areas_to_run)))\n",
    "    ax.set_xticklabels(areas_to_run)\n",
    "    ax.set_ylabel(f'balanced accuracy')\n",
    "    ax.legend(['previous_image_confidence', 'change_confidence'])\n",
    "\n",
    "    #Correct for multiple comparisons\n",
    "    area_pvalues_change_vs_image = multiple_comparisons(area_pvalues_change_vs_image)\n",
    "\n",
    "    maxy = ax.get_ylim()[1]\n",
    "    for ia, area in enumerate(areas_to_run):\n",
    "        if area_pvalues_change_vs_image[area]<0.05:\n",
    "            ax.text(ia, maxy, '*', horizontalalignment='center', verticalalignment='center')\n",
    "    "
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "for col in cols:\n",
    "    for conf_values, name in zip([area_change_conf_values[col], area_previous_image_conf_values[col]], ['change confidence', 'image confidence']):\n",
    "        ps,sigs = comparison_matrix(*list(conf_values.values()), test_func=scipy.stats.ranksums)\n",
    "        fig, ax = plt.subplots()\n",
    "        fig.suptitle(name)\n",
    "        # Create a mask for coordinates that are < 0.05\n",
    "        mask = ps < 0.05\n",
    "        # Plot the modified array\n",
    "        plt.xticks(range(len(ps)), list(conf_values.keys()), rotation=90)\n",
    "        plt.yticks(range(len(ps)), list(conf_values.keys()))\n",
    "        # Draw red asterisks at coordinates where ps < 0.05\n",
    "        ax.plot(*np.where(mask), 'w*')\n",
    "        im = ax.imshow(ps, clim=(0,1))\n",
    "        plt.xlabel('Area')\n",
    "        plt.ylabel('Area')\n",
    "        plt.colorbar(im, label='corrected p-value')\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "experience_level = ['Familiar', 'Novel']\n",
    "cols = ['is_change', 'lickbout_for_flash_during_response_window']\n",
    "samplesize = 'all'\n",
    "for experience in experience_level:\n",
    "    for col in cols:\n",
    "        fig, ax = plt.subplots()\n",
    "        fig.suptitle(f'{col} {experience}')\n",
    "        for ia, area in enumerate(areas_to_run):\n",
    "            df_to_use = sess_change_df[sess_change_df['experience_level']==experience]\n",
    "            previous_image_conf = df_to_use[col+'_test_balanced_accuracy_previous_image_confidence_' + area + '_' + samplesize].agg(['mean', 'sem'])\n",
    "            change_conf = df_to_use[col+'_test_balanced_accuracy_change_confidence_' + area + '_' + samplesize].agg(['mean', 'sem'])\n",
    "\n",
    "            # ax.violinplot(changelda, [ia])\n",
    "            ax.plot(ia, previous_image_conf['mean'], 'ko')\n",
    "            ax.errorbar(ia, previous_image_conf['mean'], previous_image_conf['sem'], color='k')\n",
    "            \n",
    "            ax.plot(ia, change_conf['mean'], 'ro')\n",
    "            ax.errorbar(ia, change_conf['mean'], change_conf['sem'], color='r')\n",
    "\n",
    "        ax.set_xticks(np.arange(len(areas_to_run)))\n",
    "        ax.set_xticklabels(areas_to_run)\n",
    "        ax.set_ylabel(f'balanced accuracy')\n",
    "        ax.legend(['previous_image_confidence', 'change_confidence'])"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "image_name_dict = {\n",
    "    'H': ['im005_r', 'im024_r', 'im034_r', 'im087_r', 'im104_r','im114_r','im083_r','im111_r'],\n",
    "    'G': ['im012_r', 'im036_r', 'im044_r','im047_r','im078_r','im115_r','im083_r', 'im111_r']\n",
    "}"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "plt.rcParams['font.size'] = 14\n",
    "area = 'VISam'\n",
    "samplesize = 'all'\n",
    "metrics_to_plot = ['previous_image_confidence', 'change_confidence']\n",
    "colors = [(1,0,1), (0, 1, 1)]\n",
    "for metric_name, sign, color in zip(metrics_to_plot, [-1,1], colors):\n",
    "    cmap = make_mono_colormap((1,1,1), color).reversed()\n",
    "    metric_name = f'{metric_name}_{area}_{samplesize}' if metric_name != 'lickbout_for_flash_during_response_window' else metric_name\n",
    "\n",
    "    fig, ax = plt.subplots()\n",
    "    fig.suptitle(f'{metric_name}')\n",
    "    images, counts, vals = mean_paired_image_mat_from_stim_table(stimtable_with_flash_metrics, metric_name, experience='Novel', image_set='H')\n",
    "    vals = vals*sign\n",
    "    vals_normed = (vals-vals.min())/np.max(vals-vals.min())\n",
    "    im = ax.imshow(vals_normed, cmap='inferno')\n",
    "    images = [im.replace('_r', '') for im in images]  # Remove '_r' from image names for better readability\n",
    "    ax.set_xticks(np.arange(len(images)))\n",
    "    ax.set_xticklabels(images, rotation=90)\n",
    "    ax.set_yticks(np.arange(len(images)))\n",
    "    ax.set_yticklabels(images)\n",
    "    plt.colorbar(im)\n",
    "    images, counts, response_vals = mean_paired_image_mat_from_stim_table(stimtable_with_flash_metrics, 'lickbout_for_flash_during_response_window', experience='Novel', image_set='H')\n",
    "    images = [im.replace('_r', '') for im in images]  # Remove '_r' from image names for better readability\n",
    "    fig, ax = plt.subplots()\n",
    "    fig.suptitle(metric_name)\n",
    "    ax.plot(response_vals.flatten(), vals.flatten(), 'k.')\n",
    "\n",
    "    nodiag = ~np.eye(response_vals.shape[0], dtype=bool)\n",
    "\n",
    "    print(metric_name, np.corrcoef(response_vals[nodiag].flatten(), vals[nodiag].flatten())[0,1]**2)\n",
    "\n",
    "\n",
    "fig, ax = plt.subplots()\n",
    "fig.suptitle('licks')\n",
    "im = ax.imshow(response_vals, cmap='inferno', clim=(0,1))\n",
    "ax.set_xticks(np.arange(len(images)))\n",
    "ax.set_xticklabels(images, rotation=90)\n",
    "ax.set_yticks(np.arange(len(images)))\n",
    "ax.set_yticklabels(images)\n",
    "plt.colorbar(im)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import warnings\n",
    "warnings.filterwarnings(\"ignore\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "area = 'VISam'\n",
    "samplesize = 'all'\n",
    "metrics_to_plot = ['previous_image_confidence', 'change_confidence']\n",
    "\n",
    "experience='Novel'\n",
    "image_set='H'\n",
    "\n",
    "sessions = stimtable_with_flash_metrics[(stimtable_with_flash_metrics['experience_level']==experience)&(stimtable_with_flash_metrics['image_set']==image_set)]['session_id'].unique()\n",
    "\n",
    "colors = [(1,0,1), (0, 1, 1)]\n",
    "metric_lick_correlations = {}\n",
    "metric_mats = {}\n",
    "for session in sessions: \n",
    "    for metric_name, sign, color in zip(metrics_to_plot, [-1,1], colors):\n",
    "        cmap = make_mono_colormap((1,1,1), color).reversed()\n",
    "        metric_name = f'{metric_name}_{area}_{samplesize}' if metric_name != 'lickbout_for_flash_during_response_window' else metric_name\n",
    "        if metric_name not in metric_lick_correlations:\n",
    "            metric_lick_correlations[metric_name] = []\n",
    "            metric_mats[metric_name] = []\n",
    "\n",
    "        fig, ax = plt.subplots()\n",
    "        fig.suptitle(f'{metric_name}')\n",
    "        images, counts, vals = mean_paired_image_mat_from_stim_table(stimtable_with_flash_metrics, metric_name, experience='Novel', image_set='H', session_id=session)\n",
    "        vals = vals*sign\n",
    "        vals_normed = (vals-vals.min())/np.max(vals-vals.min())\n",
    "        metric_mats[metric_name].append(vals_normed)\n",
    "        \n",
    "        im = ax.imshow(vals_normed, cmap='inferno')\n",
    "        plt.colorbar(im)\n",
    "        images, counts, response_vals = mean_paired_image_mat_from_stim_table(stimtable_with_flash_metrics, 'lickbout_for_flash_during_response_window', experience='Novel', image_set='H', session_id= session)\n",
    "        fig, ax = plt.subplots()\n",
    "        fig.suptitle(metric_name)\n",
    "        ax.plot(response_vals.flatten(), vals.flatten(), 'k.')\n",
    "\n",
    "        nodiag = ~np.eye(response_vals.shape[0], dtype=bool)\n",
    "\n",
    "        print(metric_name, np.corrcoef(response_vals[nodiag].flatten(), vals[nodiag].flatten())[0,1]**2)\n",
    "        metric_lick_correlations[metric_name].append(np.corrcoef(response_vals[nodiag].flatten(), vals[nodiag].flatten())[0,1]**2)\n",
    "\n",
    "\n",
    "fig, ax = plt.subplots()\n",
    "fig.suptitle('licks')\n",
    "im = ax.imshow(response_vals, cmap='inferno', clim=(0,1))\n",
    "plt.colorbar(im)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "for metric in metric_mats:\n",
    "    mean_across_sessions = np.nanmean(metric_mats[metric], axis=0)\n",
    "    fig, ax = plt.subplots()\n",
    "    fig.suptitle(metric)\n",
    "    im = ax.imshow(mean_across_sessions, cmap='inferno', clim=(0,1))\n",
    "    plt.colorbar(im)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "fig, ax = plt.subplots()\n",
    "fig.set_size_inches(6,6)\n",
    "ax.plot(metric_lick_correlations['previous_image_confidence_VISam_all'], metric_lick_correlations['change_confidence_VISam_all'], 'ko', alpha=0.5)\n",
    "ax.set_aspect('equal')\n",
    "ax.plot([0,1], [0,1], 'k--')\n",
    "ax.set_xlim(-0.05,0.2)"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Lick rates for holdovers during familiar and novel sessions"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "sessions_table = pd.read_csv(sessions_table_file)\n",
    "good_sessions = sessions_table[sessions_table['abnormal_histology'].isnull() & sessions_table['abnormal_activity'].isnull()]\n",
    "sessions_to_analyze = good_sessions['ecephys_session_id'].unique()\n",
    "\n",
    "session_beh_mats = {}\n",
    "beh_mat_array = []\n",
    "count_array = []\n",
    "session_labels = []\n",
    "experience_labels = []\n",
    "image_set_labels = []\n",
    "for _, session in good_sessions.iterrows():\n",
    "    session_id = session['ecephys_session_id']\n",
    "    mouse_id = session['mouse_id']\n",
    "    experience = session['experience_level']\n",
    "\n",
    "    ims, counts, beh_mat = beh_mat_from_stim_table(stim_table, session_id)\n",
    "    image_set = 'G' if 'im036_r' in ims else 'H'\n",
    "    \n",
    "    beh_mat_array.append(beh_mat)\n",
    "    count_array.append(counts)\n",
    "    session_labels.append(session_id)\n",
    "    experience_labels.append(experience)\n",
    "    image_set_labels.append(image_set)\n",
    "\n",
    "\n",
    "    session_beh_mats[session_id] = {'beh_mat': beh_mat, 'counts': counts, 'images': ims}\n",
    "    "
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "mouse_session_counts = good_sessions.value_counts('mouse_id')\n",
    "mice_with_fam_and_nov_sessions = mouse_session_counts[mouse_session_counts > 1].index.values\n",
    "mice_with_fam_and_nov_sessions\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "im_83_hit_rates = {'Familiar':[], 'Novel':[], 'Familiar_from_holdover':[], 'Novel_from_holdover':[], }\n",
    "im_111_hit_rates = {'Familiar':[], 'Novel':[], 'Familiar_from_holdover':[], 'Novel_from_holdover':[], }\n",
    "\n",
    "for mouse in mice_with_fam_and_nov_sessions:\n",
    "\n",
    "    for experience in ['Familiar', 'Novel']:\n",
    "\n",
    "        session_id = get_experience_session_id_for_mouse(sessions_table, mouse, experience)\n",
    "\n",
    "        beh_mat = session_beh_mats[session_id]['beh_mat']\n",
    "        counts = session_beh_mats[session_id]['counts']\n",
    "        ims = session_beh_mats[session_id]['images']\n",
    "\n",
    "        beh_mat_no_diag = skip_diag_masking(beh_mat)\n",
    "        counts_no_diag = skip_diag_masking(counts)\n",
    "\n",
    "\n",
    "        for ind, imrates in zip([6,7], [im_83_hit_rates[experience], im_111_hit_rates[experience]]):\n",
    "\n",
    "            ind_count = np.nansum(counts_no_diag[:, ind])\n",
    "            if (ind_count > 10) & (get_session_engaged_dprime(stim_table, session_id)>1) & (get_session_engaged_hit_count(stim_table, session_id)>50):\n",
    "\n",
    "                imrates.append(np.nanmean(beh_mat_no_diag[:, ind]))\n",
    "            \n",
    "            else:\n",
    "                imrates.append(np.nan)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import vbn_utils"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "fig, ax = plt.subplots()\n",
    "ax.plot(im_83_hit_rates['Familiar'], im_83_hit_rates['Novel'], 'ko')\n",
    "ax.plot(im_111_hit_rates['Familiar'], im_111_hit_rates['Novel'], 'ko', mfc='w')\n",
    "ax.set_aspect('equal')\n",
    "ax.plot([0,1], [0,1], 'k--')\n",
    "ax.set_xlabel('Hit rate during Familiar session')\n",
    "ax.set_ylabel('Hit rate during Novel session')\n",
    "\n",
    "vbn_utils.formatFigure(fig, ax, )"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "#Add dprime and hit count to sessions table\n",
    "dprimes = []\n",
    "hit_counts = []\n",
    "for _, session in sessions.iterrows():\n",
    "\n",
    "    session_id = session['ecephys_session_id']\n",
    "    dprime = get_session_engaged_dprime(stim_table, session_id)\n",
    "    hit_count = get_session_engaged_hit_count(stim_table, session_id)\n",
    "\n",
    "    dprimes.append(dprime)\n",
    "    hit_counts.append(hit_count)\n",
    "\n",
    "\n",
    "sessions['engaged_dprime'] = dprimes\n",
    "sessions['engaged_hitcount'] = hit_counts"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "good_sessions = sessions[sessions['abnormal_histology'].isnull() & sessions['abnormal_activity'].isnull()]\n",
    "good_behavior_sessions = good_sessions[(good_sessions['engaged_dprime']>=1)&(good_sessions['engaged_hitcount']>=50)]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "areas_to_run = ('VISall', 'VISp', 'VISl', 'VISrl', 'VISal', 'VISpm', 'VISam', 'LGd', 'LP', 'SCMRN')\n",
    "metrics_to_run = ['previous_image_confidence', 'change_confidence']\n",
    "samplesizes = ['20', '40', 'all']\n",
    "\n",
    "response_rate_summary = {s:{f'{image_set}_{experience_level}_{metric}_{area}_{samplesize}_{response_type}': np.nan \\\n",
    "                            for image_set in ['G', 'H'] \\\n",
    "                            for experience_level in ['Familiar', 'Novel'] \\\n",
    "                            for metric in ['previous_image_confidence', 'change_confidence'] \\\n",
    "                            for area in areas_to_run \\\n",
    "                            for samplesize in samplesizes \\\n",
    "                            for response_type in ['private_nonchange', 'shared_nonchange', 'private_hit', 'shared_hit', 'private_fa', 'shared_fa', 'omission', 'postomission']}\n",
    "                        for s in good_behavior_sessions['ecephys_session_id'].values}\n",
    "\n",
    "for isess, session in good_behavior_sessions.iterrows():\n",
    "    print(isess)\n",
    "    experience_level = session['experience_level']\n",
    "    image_set = session['image_set']\n",
    "    session_id = session['ecephys_session_id']\n",
    "    session_stim_table = stimtable_with_flash_metrics[stimtable_with_flash_metrics['session_id']==session_id]\n",
    "    \n",
    "    for metric in metrics_to_run:\n",
    "        for area in areas_to_run:\n",
    "            for samplesize in samplesizes:\n",
    "                metric_name = f'{metric}_{area}_{samplesize}'\n",
    "                response_rate_summary[session_id][image_set + '_' + experience_level + '_' + metric_name + '_private_nonchange'] = get_private_nonchange_mean(session_stim_table, metric_name)\n",
    "                response_rate_summary[session_id][image_set + '_' + experience_level + '_' + metric_name + '_private_hit'] = get_private_change_mean(session_stim_table, metric_name)\n",
    "                response_rate_summary[session_id][image_set + '_' + experience_level + '_' + metric_name + '_private_fa'] = get_private_catch_mean(session_stim_table, metric_name)\n",
    "                response_rate_summary[session_id][image_set + '_' + experience_level + '_' + metric_name + '_shared_nonchange'] = get_shared_nonchange_mean(session_stim_table, metric_name)\n",
    "                response_rate_summary[session_id][image_set + '_' + experience_level + '_' + metric_name + '_shared_hit'] = get_shared_change_mean(session_stim_table, metric_name)\n",
    "                response_rate_summary[session_id][image_set + '_' + experience_level + '_' + metric_name + '_shared_fa'] = get_shared_catch_mean(session_stim_table, metric_name)\n",
    "                response_rate_summary[session_id][image_set + '_' + experience_level + '_' + metric_name + '_omission'] = get_omission_mean(session_stim_table, metric_name)\n",
    "                response_rate_summary[session_id][image_set + '_' + experience_level + '_' + metric_name + '_postomission'] = get_post_omission_mean(session_stim_table, metric_name)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "response_rate_df = pd.DataFrame.from_dict(response_rate_summary, orient='index')"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "response_rate_df.to_csv(os.path.join(save_dir, 'previous_image_and_change_confidence_response_metrics.csv'))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "aggregated_labels = ['Familiar_shared', 'Familiar_private', 'Novel_shared', 'Novel_private', 'Novel', 'Familiar']\n",
    "response_categories = ['nonchange', 'hit', 'fa', 'omission', 'postomission']\n",
    "samplesize = 'all'\n",
    "aggregations = (\n",
    "                ['Familiar', 'private', 'nonchange'],\n",
    "                ['Familiar', 'shared', 'nonchange'],\n",
    "                ['Novel', 'private', 'nonchange'],\n",
    "                ['Novel', 'shared', 'nonchange'],\n",
    "                ['Familiar', 'private', 'hit'],\n",
    "                ['Familiar', 'shared', 'hit'],\n",
    "                ['Novel', 'private', 'hit'],\n",
    "                ['Novel', 'shared', 'hit'],\n",
    "                ['Familiar', 'private', 'fa'],\n",
    "                ['Familiar', 'shared', 'fa'],\n",
    "                ['Novel', 'private', 'fa'],\n",
    "                ['Novel', 'shared', 'fa'],\n",
    "                ['Familiar', '', 'omission'],\n",
    "                ['Novel', '', 'omission'],\n",
    "                ['Familiar', '', 'postomission'],\n",
    "                ['Novel', '', 'postomission'],\n",
    "                )\n",
    "\n",
    "\n",
    "\n",
    "cols_to_aggregate = []\n",
    "for metric in ['previous_image_confidence', 'change_confidence']:\n",
    "    for area in areas_to_run:\n",
    "        for aggregation in aggregations:\n",
    "                cols = [c for c in response_rate_df.columns if \\\n",
    "                        aggregation[0] in c and \\\n",
    "                        aggregation[1] in c and \\\n",
    "                        f'_{aggregation[2]}' in c and \\\n",
    "                        metric in c and f'_{area}_' in c and f'_{samplesize}_' in c]\n",
    "                if len(cols)>0:\n",
    "                    cols_to_aggregate.append(cols)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "labels = []\n",
    "values = []\n",
    "means = []\n",
    "sems = []\n",
    "for cols in cols_to_aggregate:\n",
    "    labels.append(cols[0][2:])\n",
    "    vals = response_rate_df[cols].to_numpy().flatten()\n",
    "    vals = vals[~np.isnan(vals)]\n",
    "\n",
    "    values.append(vals)\n",
    "    means.append(np.mean(vals))\n",
    "    sems.append(np.std(vals)/len(vals)**0.5)\n",
    "\n",
    "labels = np.array(labels)\n",
    "values = np.array(values)\n",
    "means = np.array(means)\n",
    "sems = np.array(sems)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "for area in areas_to_run:\n",
    "    for metric, sign in zip(['previous_image_confidence', 'change_confidence'], [-1,1]):\n",
    "        fig, ax = plt.subplots()\n",
    "        fig.suptitle(metric + ' ' + area)\n",
    "        inds = [i for i, label in enumerate(labels) if metric in label and f'_{area}_' in label]\n",
    "        metric_labels = [l.replace(f'_{metric}', '').replace(f'_{area}', '').replace(f'_{samplesize}', '') for l in labels[inds]]\n",
    "        metric_values = values[inds] * sign\n",
    "        ax.violinplot(metric_values, showmedians=True)\n",
    "        \n",
    "        ax.set_xticks(np.arange(len(metric_labels))+1)\n",
    "        ax.set_xticklabels(metric_labels, rotation=90)\n",
    "plt.tight_layout()"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Plot flash responses in PCA space to illustrate different strategies"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "sessionId = good_behavior_sessions.iloc[0]['ecephys_session_id']\n",
    "region = 'VISall'\n",
    "\n",
    "stim = stim_table[(stim_table['session_id']==sessionId) & stim_table['active']].reset_index()\n",
    "im_ids = stim['image_name'].values\n",
    "previous_im_ids = np.insert(stim['image_name'].values, 0, 'omitted')[:-1]\n",
    "previous_im_ids_skipping_omitted = []\n",
    "for ind, id in enumerate(previous_im_ids):\n",
    "    if not id=='omitted':\n",
    "        previous_im_ids_skipping_omitted.append(id)\n",
    "    else:\n",
    "        if ind==0:\n",
    "            previous_im_ids_skipping_omitted.append(np.nan)\n",
    "        else:\n",
    "            previous_im_ids_skipping_omitted.append(previous_im_ids[ind-1])\n",
    "\n",
    "previous_im_ids_skipping_omitted = np.array(previous_im_ids_skipping_omitted)\n",
    "\n",
    "stim['previous_image_skipping_omitted'] = previous_im_ids_skipping_omitted\n",
    "unique_images = [im for im in np.unique(stim['image_name'].values) if im not in ['im083_r', 'im111_r', 'omitted']] + ['im083_r', 'im111_r']\n",
    "\n",
    "units = unit_table.set_index('unit_id').loc[unitData[str(sessionId)]['unitIds'][:]]\n",
    "spikes = unitData[str(sessionId)]['spikes']\n",
    "highQuality = du.apply_unit_quality_filter(units, no_abnorm=False)\n",
    "\n",
    "\n",
    "inRegion = du.getUnitsInRegion(units,region)\n",
    "final_unit_filter = highQuality & inRegion\n",
    "\n",
    "cols_to_add = [f'{metric}_{region}_{unitSampleName}' for unitSampleName in unitSampleNames for metric in ['previous_image_confidence', 'change_confidence']]\n",
    "stim[cols_to_add] = np.nan\n",
    "\n",
    "sp = np.zeros((final_unit_filter.sum(),spikes.shape[1],spikes.shape[2]),dtype=bool)\n",
    "for i,u in enumerate(np.where(final_unit_filter)[0]):\n",
    "    sp[i]=spikes[u,:,:]\n",
    "\n",
    "flash_response = sp[:, :, respWin].mean(axis=2)\n",
    "im_id_ints = np.array([np.where(im==np.unique(im_ids))[0][0] for im in im_ids])"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### Plot colored by image ID"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "save_dir = \"/Volumes/programs/mindscope/workgroups/np-behavior/VBN Manuscript/computation_of_change\""
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "images_to_plot = [1,2]\n",
    "im_id_to_color = images_to_plot[0]\n",
    "to_color = im_id_ints == im_id_to_color\n",
    "not_to_color = np.isin(im_id_ints, images_to_plot) & ~to_color\n",
    "\n",
    "lda_clf = LDA()\n",
    "non_omitted = ~(stim['omitted'].astype(bool).values)\n",
    "lda_clf.fit(flash_response[:, to_color | not_to_color].T, to_color[to_color | not_to_color])\n",
    "reduced_image = lda_clf.transform(flash_response.T)\n",
    "\n",
    "lda_change_clf = LDA()\n",
    "lda_change_clf.fit(flash_response[:, to_color | not_to_color].T, stim[to_color | not_to_color]['is_change'])\n",
    "reduced_change = lda_change_clf.transform(flash_response.T)\n",
    "change_preds = lda_change_clf.predict(flash_response.T)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "point_on_boundary = np.zeros(len(lda_change_clf.coef_[0]))\n",
    "point_on_boundary[0] = -lda_change_clf.intercept_/lda_change_clf.coef_[0][0]\n",
    "lda_change_boundary = lda_change_clf.transform(point_on_boundary.reshape(1,-1))\n",
    "\n",
    "point_on_boundary = np.zeros(len(lda_clf.coef_[0]))\n",
    "point_on_boundary[0] = -lda_clf.intercept_/lda_clf.coef_[0][0]\n",
    "lda_image_boundary = lda_clf.transform(point_on_boundary.reshape(1,-1))\n",
    "\n",
    "lda_change_boundary[0], lda_image_boundary[0]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "reduced = np.hstack([reduced_image, reduced_change])"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# find a change from the first image to the second image\n",
    "eligible_changes = stim[(stim['previous_image_skipping_omitted']==np.unique(im_ids)[images_to_plot[0]])&(stim['image_name']==np.unique(im_ids)[images_to_plot[1]])]\n",
    "eligible_change_inds = eligible_changes.index.values"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "images_used = np.unique(im_ids)[images_to_plot[0]]+'_'+np.unique(im_ids)[images_to_plot[1]]\n",
    "images_used"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# Plot colored by image ID in 2D\n",
    "num_to_show = 1\n",
    "ims_to_show = np.zeros(len(im_ids), dtype=bool)\n",
    "ims_to_show[::num_to_show] = True\n",
    "\n",
    "eligible_change_ind = eligible_change_inds[3]\n",
    "\n",
    "fig, ax = plt.subplots()\n",
    "fig.suptitle('colored_by_image')\n",
    "fig.patch.set_facecolor('white')\n",
    "\n",
    "fig.set_size_inches(8,8)\n",
    "ax.set_facecolor('white')\n",
    "\n",
    "# Remove grid lines\n",
    "ax.grid(False)\n",
    "\n",
    "ax.scatter(reduced[non_omitted & to_color & ims_to_show & ~(stim['is_change']), 0], \n",
    "           reduced[non_omitted & to_color & ims_to_show & ~(stim['is_change']), 1], \n",
    "           c='teal', alpha=0.1, edgecolors='none', s=50)\n",
    "ax.scatter(reduced[non_omitted & not_to_color & ims_to_show & ~(stim['is_change']), 0], \n",
    "           reduced[non_omitted & not_to_color & ims_to_show & ~(stim['is_change']), 1], \n",
    "           c='gray', alpha=0.1, edgecolors='none', s=50)\n",
    "ax.scatter(reduced[non_omitted & not_to_color & (stim['is_change']) & ims_to_show, 0], \n",
    "           reduced[non_omitted & not_to_color & (stim['is_change']) & ims_to_show, 1], \n",
    "           c='w', alpha=0.3, edgecolors='gray',facecolors='w', s=40)\n",
    "ax.scatter(reduced[non_omitted & to_color & (stim['is_change']) & ims_to_show, 0], \n",
    "           reduced[non_omitted & to_color & (stim['is_change']) & ims_to_show, 1], \n",
    "           c='w', alpha=0.3, edgecolors='teal',facecolors='w', s=40)\n",
    "\n",
    "ax.scatter(reduced[eligible_change_ind, 0],reduced[eligible_change_ind, 1], c='r')\n",
    "ax.scatter(reduced[eligible_change_ind+1, 0],reduced[eligible_change_ind+1, 1], c='r')\n",
    "\n",
    "for nback in range(1,4):\n",
    "    ax.scatter(reduced[eligible_change_ind-nback, 0],reduced[eligible_change_ind-nback, 1], c='teal', alpha=1/nback)\n",
    "\n",
    "\n",
    "ax.axhline(lda_change_boundary[0], color='k', linestyle='--')\n",
    "ax.axvline(lda_image_boundary[0], color='k', linestyle='--')\n",
    "\n",
    "fig.savefig(os.path.join(save_dir, f'{sessionId}_{images_used}_test.pdf'))"
   ]
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "vbn_manuscript",
   "language": "python",
   "name": "python3"
  },
  "language_info": {
   "codemirror_mode": {
    "name": "ipython",
    "version": 3
   },
   "file_extension": ".py",
   "mimetype": "text/x-python",
   "name": "python",
   "nbconvert_exporter": "python",
   "pygments_lexer": "ipython3",
   "version": "3.8.20"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 2
}
