{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import pandas as pd\n",
    "import numpy as np\n",
    "import os\n",
    "from matplotlib import pyplot as plt\n",
    "import decoding_utils as du\n",
    "%matplotlib inline "
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import vbn_utils"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "vscode": {
     "languageId": "markdown"
    }
   },
   "outputs": [],
   "source": [
    "%load_ext autoreload\n",
    "%autoreload 2"
   ]
  },
  {
   "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",
    "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": [
    "units = pd.read_csv(unit_table_file)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "#Clean layers\n",
    "units['cortical_layer'].replace('3-Feb', '2/3', inplace=True)\n",
    "units['cortical_layer'].replace('6a', '6', inplace=True)\n",
    "units['cortical_layer'].replace('6b', '6', inplace=True)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "stim_table = pd.read_csv(stim_table_file)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "sessions_table = pd.read_csv(sessions_table_file)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "g_images = ['omitted'] + list(np.sort(stim_table[(stim_table['stimulus_name'].str.contains('_G_'))&\n",
    "                    (~stim_table['omitted'])&\n",
    "                    (~stim_table['image_name'].isin(['im083_r','im111_r']))]['image_name'].unique())) + ['im083_r','im111_r']\n",
    "\n",
    "h_images = ['omitted'] + list(np.sort(stim_table[(stim_table['stimulus_name'].str.contains('_H_'))&\n",
    "                    (~stim_table['omitted'])&\n",
    "                    (~stim_table['image_name'].isin(['im083_r','im111_r']))]['image_name'].unique())) + ['im083_r','im111_r']"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### Generated on HPC by 'run_image_decoding.py'"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "image_decoding_dir = \"/Volumes/programs/mindscope/workgroups/np-behavior/vbn_data_release/image_decoding\"\n",
    "decoding_files = os.listdir(image_decoding_dir)\n",
    "decoding_files = [os.path.join(image_decoding_dir, f) for f in decoding_files if f.endswith('nonchangeRS.npy')]\n",
    "\n",
    "session_ids = [int(os.path.basename(f).split('_')[0]) for f in decoding_files]\n",
    "\n",
    "missing_sessions = [s for s in sessions_table['ecephys_session_id'].values if s not in session_ids]\n",
    "missing_sessions"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "decoding_results_files = os.listdir(\"/Volumes/programs/mindscope/workgroups/np-behavior/vbn_data_release/image_decoding\")\n",
    "decoding_results_files = [os.path.join(\"/Volumes/programs/mindscope/workgroups/np-behavior/vbn_data_release/image_decoding\", d) for d in decoding_results_files]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "no_anomalies = sessions_table[(sessions_table['abnormal_activity'].isnull())&(sessions_table['abnormal_histology'].isnull())]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from notebook_utils import (\n",
    "    get_decoding_results_files, get_mouse_paired_indices\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "conditions = ['Familiar', 'Novel']"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "image_set = 'nonchangeRS' \n",
    "session_decoding_files = get_decoding_results_files(no_anomalies['ecephys_session_id'].values, decoding_results_files, image_set=image_set)\n",
    "decoding_results_dict = {}\n",
    "for decoding_file in session_decoding_files:\n",
    "\n",
    "    d = np.load(decoding_file, allow_pickle=True).item()\n",
    "    decoding_results_dict[decoding_file] = d\n",
    "\n",
    "decode_windows = d['decodeWindows']"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "gtoh_sessions = no_anomalies[((no_anomalies['image_set']=='G')&(no_anomalies['experience_level']=='Familiar')) | \n",
    "                            (((no_anomalies['image_set']=='H')&(no_anomalies['experience_level']=='Novel')))]\n",
    "\n",
    "htog_sessions = no_anomalies[((no_anomalies['image_set']=='H')&(no_anomalies['experience_level']=='Familiar')) | \n",
    "                            (((no_anomalies['image_set']=='G')&(no_anomalies['experience_level']=='Novel')))]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "regions = list(decoding_results_dict[decoding_file].keys())\n",
    "regions = [r for r in regions if r not in ['decodeWindows', 'hit', 'image_order']]\n",
    "unit_sub_sample_sizes = list(decoding_results_dict[decoding_file][regions[0]].keys())\n",
    "\n",
    "conditions = ['Familiar', 'Novel']\n",
    "\n",
    "region_timecourses = {}\n",
    "session_ids = {r:{c: {n:[] for n in unit_sub_sample_sizes} for c in conditions} for r in regions}\n",
    "for region in regions:\n",
    "    imagewise_recall_per_condition = {c:{} for c in conditions}\n",
    "    for condition in conditions:\n",
    "\n",
    "        condition_sessions = no_anomalies[no_anomalies['experience_level']==condition]\n",
    "        condition_decoding_files = get_decoding_results_files(condition_sessions['ecephys_session_id'].values, decoding_results_files, image_set=image_set)\n",
    "        \n",
    "        for n in unit_sub_sample_sizes:\n",
    "            bas = []\n",
    "            sess_ids = []\n",
    "            for cdf in condition_decoding_files:\n",
    "                d = decoding_results_dict[cdf]\n",
    "                ba = np.array(d[region][n]['imagewise_recall'])\n",
    "                #ba = np.array(d[region][n]['imagewise_precision']) #can also plot precision\n",
    "                if ba.size>0:\n",
    "                    bas.append(ba)\n",
    "                    sess_ids.append(int(os.path.basename(cdf).split('_')[0]))\n",
    "\n",
    "            ba_array = np.array(bas)\n",
    "            imagewise_recall_per_condition[condition][n] = ba_array\n",
    "            session_ids[region][condition][n] = sess_ids\n",
    "        \n",
    "    region_timecourses[region] = imagewise_recall_per_condition"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from notebook_utils import (\n",
    "    \n",
    "    get_sigmoidfit_midpoint,\n",
    "    \n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "for region in regions:\n",
    "    fig, axes = plt.subplots(1,len(unit_sub_sample_sizes))\n",
    "    fig.set_size_inches(18,6)\n",
    "    imagewise_recall_per_condition = region_timecourses[region]\n",
    "\n",
    "    for n, ax in zip(unit_sub_sample_sizes, axes):\n",
    "        colors = ['b', 'r']\n",
    "        counts = []\n",
    "        for ic, condition in enumerate(['Familiar', 'Novel']):\n",
    "            count = imagewise_recall_per_condition[condition][n].shape[0]\n",
    "            counts.append(count)\n",
    "            if count>0:\n",
    "                mean = np.mean(imagewise_recall_per_condition[condition][n][:, :, :6], axis=(0,2))\n",
    "                std = np.std(imagewise_recall_per_condition[condition][n][:, :, :6], axis=(0,2))\n",
    "                sem = std/count**0.5\n",
    "                time = np.arange(0, len(mean)*10, 10)\n",
    "                ax.plot(time, mean, colors[ic])\n",
    "                ax.fill_between(time, mean+sem, mean-sem, color=colors[ic], alpha=0.3)\n",
    "        \n",
    "        ax.set_title(f'{region} {n}, f: {counts[0]}, n: {counts[1]}')\n",
    "        ax.set_xlim([0,150])\n",
    "        ax.set_ylim([0,1])\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "#accuracy at end of decision window (100ms) and n=40\n",
    "hit_rates = {r: {c:{n:[] for n in unit_sub_sample_sizes} for c in ['Familiar', 'Novel']} for r in regions}\n",
    "for region in regions:\n",
    "    for n in [40,20]:\n",
    "        for condition in conditions:\n",
    "            timecourse = region_timecourses[region][condition][n]\n",
    "            count = timecourse.shape[0]\n",
    "            if count>0:\n",
    "                timecourse_means = np.mean(timecourse[:,:,:6], axis=2)\n",
    "                hrs = timecourse_means[:, 10]\n",
    "                hit_rates[region][condition][n] = hrs\n",
    "                "
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from notebook_utils import bh_multitest\n",
    "import scipy.stats"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "regions_to_plot = ['LGd', 'LP', 'VISp', 'VISl', 'VISrl', 'VISal', 'VISpm', 'VISam', 'VISall', 'SCMRN']\n",
    "fig, ax = plt.subplots()\n",
    "pvals = []\n",
    "for ir, region in enumerate(regions_to_plot):#regions_to_plot:\n",
    "\n",
    "    vals = [np.array(hit_rates[region][cond][40]) for cond in ['Familiar', 'Novel']]\n",
    "    pval = scipy.stats.ranksums(*vals, nan_policy='omit')[1]\n",
    "\n",
    "    pvals.append(pval)\n",
    "    \n",
    "    ax.plot(ir, np.nanmean(vals[0]), 'bo')\n",
    "    ax.plot(ir+0.1, np.nanmean(vals[1]), 'ro')\n",
    "\n",
    "    ax.errorbar(ir, np.nanmean(vals[0]), np.nanstd(vals[0])/(np.sum(~np.isnan(vals[0]))**0.5), color='b')\n",
    "    ax.errorbar(ir+0.1, np.nanmean(vals[1]), np.nanstd(vals[1])/(np.sum(~np.isnan(vals[1]))**0.5), color='r')\n",
    "\n",
    "ax.set_xticks(np.arange(len(regions_to_plot)))\n",
    "ax.set_xticklabels(regions_to_plot, rotation=90)\n",
    "ax.set_ylabel('Decoding accuracy')\n",
    "\n",
    "sig_after_correction = bh_multitest(pvals)[0]\n",
    "sigx_ind = np.where(sig_after_correction)[0]\n",
    "if len(sigx_ind)>0:\n",
    "    ax.text(sigx_ind[0], 1, '*')\n",
    "\n",
    "vbn_utils.formatFigure(fig, ax)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "colors = ['b', 'r']\n",
    "plt.rcParams['font.size'] = 20\n",
    "latencies = {r: {c:{n:[] for n in unit_sub_sample_sizes} for c in ['Familiar', 'Novel']} for r in regions}\n",
    "latency_sessions = {r: {c:{n:[] for n in unit_sub_sample_sizes} for c in ['Familiar', 'Novel']} for r in regions}\n",
    "\n",
    "for region in ['LGd', 'VISp', 'VISl', 'VISall']:\n",
    "    imagewise_recall_per_condition = region_timecourses[region]\n",
    "    plt.figure()\n",
    "    fig = plt.gcf()\n",
    "    fig.set_size_inches([8,6])\n",
    "    counts=[]\n",
    "    for n in [40]:\n",
    "        for ic, condition in enumerate(['Familiar', 'Novel']):\n",
    "            print(imagewise_recall_per_condition[condition][n].shape)\n",
    "            count = imagewise_recall_per_condition[condition][n].shape[0]\n",
    "            counts.append(count)\n",
    "            if count>0:\n",
    "                image_means = np.mean(imagewise_recall_per_condition[condition][n][:, :, :6], axis=(2))\n",
    "                lats = [get_sigmoidfit_midpoint(decode_windows, y)[0] for y in image_means]\n",
    "                ylats = [get_sigmoidfit_midpoint(decode_windows, y)[1] for y in image_means]\n",
    "                plt.plot(decode_windows, image_means.T, colors[ic], alpha=0.2)\n",
    "                plt.plot(lats, ylats, colors[ic]+'o', alpha=0.5, mec='none')\n",
    "                \n",
    "                latencies[region][condition][n] = lats\n",
    "                plt.plot(decode_windows, np.mean(image_means, axis=0), colors[ic], lw=2)\n",
    "                print(f'{region}: {np.median(lats)}')\n",
    "    ax = plt.gca()\n",
    "    ax.set_title(f'{region} {n} units, f: {counts[0]}, n: {counts[1]}')\n",
    "    ax.set_xlim([0, 120])\n",
    "    ax.set_xticks([20, 100])\n",
    "    ax.spines['bottom'].set_bounds(20, 100)\n",
    "    vbn_utils.formatFigure(fig, ax)\n",
    "    ax.set_xlabel('Time from stim start (ms)')\n",
    "    ax.set_ylabel('Decoding accuracy')\n",
    "\n",
    "    # Set the linewidth of all spines\n",
    "    for spine in ax.spines.values():\n",
    "        spine.set_linewidth(1.5)  # Set spine linewidth to 2 points\n",
    "\n",
    "    # Set the linewidth of ticks\n",
    "    ax.tick_params(width=1.5)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "latencies = {r: {c:{n:[] for n in unit_sub_sample_sizes} for c in ['Familiar', 'Novel']} for r in regions}\n",
    "latency_sessions = {r: {c:{n:[] for n in unit_sub_sample_sizes} for c in ['Familiar', 'Novel']} for r in regions}\n",
    "\n",
    "for region in regions:\n",
    "    imagewise_recall_per_condition = region_timecourses[region]\n",
    "\n",
    "    for n in [40,20]:#[10,20,40]:\n",
    "        for ic, condition in enumerate(['Familiar', 'Novel']):\n",
    "            count = imagewise_recall_per_condition[condition][n].shape[0]\n",
    "            counts.append(count)\n",
    "            if count>0:\n",
    "                image_means = np.mean(imagewise_recall_per_condition[condition][n][:, :, :6], axis=(2))\n",
    "                lats = [get_sigmoidfit_midpoint(decode_windows, y)[0] for y in image_means]\n",
    "                \n",
    "                latencies[region][condition][n] = lats\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "regions_to_plot = ['LGd', 'LP', 'VISp', 'VISl', 'VISrl', 'VISal', 'VISpm', 'VISam', 'VISall', 'SCMRN']\n",
    "region_means = []\n",
    "plt.figure()\n",
    "pvals = []\n",
    "for ir, region in enumerate(regions_to_plot):#regions:\n",
    "\n",
    "    vals = [np.array(latencies[region][cond][20]) for cond in ['Familiar', 'Novel']]\n",
    "    pvals.append(scipy.stats.mannwhitneyu(*vals, nan_policy='omit')[1])\n",
    "    # print(ir)\n",
    "    # print(np.nanmean(vals[0]))\n",
    "    plt.plot(ir, np.nanmean(vals[0]), 'bo')\n",
    "    plt.plot(ir+0.1, np.nanmean(vals[1]), 'ro')\n",
    "\n",
    "    plt.errorbar(ir, np.nanmean(vals[0]), np.nanstd(vals[0])/(np.sum(~np.isnan(vals[0]))**0.5), color='b')\n",
    "    plt.errorbar(ir+0.1, np.nanmean(vals[1]), np.nanstd(vals[1])/(np.sum(~np.isnan(vals[1]))**0.5), color='r')\n",
    "\n",
    "ax = plt.gca()\n",
    "ax.set_xticks(np.arange(len(regions_to_plot)))\n",
    "ax.set_xticklabels(regions_to_plot, rotation=90)\n",
    "ax.set_ylabel('Decoding latency (ms)')\n",
    "\n",
    "sig_after_correction = bh_multitest(pvals)[0]\n",
    "sigx_ind = np.where(sig_after_correction)[0]\n",
    "[ax.text(x, ax.get_ylim()[1], '*') for x in sigx_ind]\n",
    "vbn_utils.formatFigure(fig, ax)"
   ]
  }
 ],
 "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
}
