{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": 1,
   "id": "02bb45fc",
   "metadata": {},
   "outputs": [
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "/opt/anaconda3/envs/vbn_manuscript/lib/python3.8/site-packages/tqdm/auto.py:21: TqdmWarning: IProgress not found. Please update jupyter and ipywidgets. See https://ipywidgets.readthedocs.io/en/stable/user_install.html\n",
      "  from .autonotebook import tqdm as notebook_tqdm\n"
     ]
    }
   ],
   "source": [
    "import h5py\n",
    "import pandas as pd\n",
    "import numpy as np\n",
    "from matplotlib import pyplot as plt\n",
    "import vbn_utils\n",
    "import decoding_utils as du\n",
    "from analysis_utils import exponential_convolve\n",
    "import nwb_session_utils as nwb\n",
    "\n",
    "%matplotlib inline"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "9e17ee2e",
   "metadata": {},
   "source": [
    "## Data loading"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 2,
   "id": "3fd4ca51",
   "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\""
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 3,
   "id": "22edf518",
   "metadata": {},
   "outputs": [
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "/var/folders/7s/q3zz_qj910x07vwrdqkcgczr0000gp/T/ipykernel_34779/2052873693.py:1: DtypeWarning: Columns (120,123,126,129,132,135,138,141,144,147,150,153,156,159,162,165,174,177,180,183,192,195,198,201,228,229,236) have mixed types. Specify dtype option on import or set low_memory=False.\n",
      "  units = pd.read_csv(unit_table_file)\n",
      "/var/folders/7s/q3zz_qj910x07vwrdqkcgczr0000gp/T/ipykernel_34779/2052873693.py:4: DtypeWarning: Columns (32,34,35,43,45) have mixed types. Specify dtype option on import or set low_memory=False.\n",
      "  stim_table = pd.read_csv(stim_table_file)\n"
     ]
    }
   ],
   "source": [
    "units = pd.read_csv(unit_table_file)\n",
    "units['cortical_layer'] = units['cortical_layer'].replace('3-Feb','2/3') # 2/3 sometimes gets incorrectly reformatted as a date\n",
    "\n",
    "stim_table = pd.read_csv(stim_table_file)\n",
    "stim_table = stim_table.drop(columns='Unnamed: 0')\n",
    "\n",
    "active_tensor = h5py.File(active_tensor_file)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "1ae06f63",
   "metadata": {},
   "outputs": [],
   "source": [
    "structure_tree = pd.read_csv(\"/Volumes/programs/mindscope/workgroups/np-behavior/vbn_data_release/supplemental_tables/ccf_structure_tree_2017.csv\")"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "017dc38e",
   "metadata": {},
   "source": [
    "## Lick-triggered averages"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "b06ce9ed",
   "metadata": {},
   "source": [
    "### Compute lick-triggered responses (if pre-computed, load below)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "080d36b6",
   "metadata": {},
   "outputs": [],
   "source": [
    "unit_filter = du.apply_unit_quality_filter(units)\n",
    "unit_ids = units.loc[unit_filter]['unit_id'].values \n",
    "session_list = list(active_tensor.keys())\n",
    "\n",
    "lick_aligned_unit_summary = {u:[] for u in unit_ids}\n",
    "stim_filter = ['engaged', '~is_change', '~omitted', '~previous_omitted', 'flashes_since_change>5', 'lickbout_for_flash_during_response_window']\n",
    "unit_data, shuffle_data, unitIds = vbn_utils.unit_averaged_psth_lick_aligned(active_tensor_file, \n",
    "                                    stim_table, \n",
    "                                    session_list, \n",
    "                                    unit_ids, \n",
    "                                    *stim_filter, \n",
    "                                    baseline_length=750, \n",
    "                                    resp_window_length=1500)\n",
    "    "
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "6759df76",
   "metadata": {},
   "outputs": [],
   "source": [
    "for udata, sdata, uids in zip(unit_data, shuffle_data, unitIds):\n",
    "    if len(uids)>0:\n",
    "        for ud, sd, uid in zip(udata, sdata, uids):\n",
    "            umean = exponential_convolve(ud, 3, symmetrical=True)\n",
    "            time_above = np.convolve(umean>sd[2], np.ones(10)).max()\n",
    "            time_below = np.convolve(umean<sd[1], np.ones(10)).max()\n",
    "            passes = np.max((time_above, time_below))>9\n",
    "            lick_aligned_unit_summary[uid] = {'mean': umean, \n",
    "                                              'shuffle_mean': sd[0],\n",
    "                                              'shuffle_ci_low': sd[1],\n",
    "                                              'shuffle_ci_high': sd[2],\n",
    "                                              'pass': passes}"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "9e6049b2",
   "metadata": {},
   "outputs": [],
   "source": [
    "import pickle\n",
    "with open(\"/Volumes/programs/mindscope/workgroups/np-behavior/vbn_data_release/supplemental_tables/lick_aligned_unit_data.pkl\", 'wb') as file:\n",
    "    pickle.dump(lick_aligned_unit_summary, file)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "a6fc7802",
   "metadata": {},
   "source": [
    "### Load pre-computed lick-triggered data"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "977a7c14",
   "metadata": {},
   "outputs": [],
   "source": [
    "import pickle\n",
    "with open(\"/Volumes/programs/mindscope/workgroups/np-behavior/vbn_data_release/supplemental_tables/lick_aligned_unit_data.pkl\", 'rb') as file:\n",
    "    lick_aligned_unit_summary = pickle.load(file)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "a0a17c32",
   "metadata": {},
   "outputs": [],
   "source": [
    "lick_aligned_df = pd.DataFrame.from_dict(lick_aligned_unit_summary, orient='index')"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "184ef01d",
   "metadata": {},
   "outputs": [],
   "source": [
    "lick_aligned_df = lick_aligned_df.merge(units[['unit_id', 'cluster_labels_new', 'structure_acronym', 'cortical_layer']], left_index=True, right_on='unit_id')"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "90f9fe69",
   "metadata": {},
   "source": [
    "## Run-start triggered averages"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "cb2da03d",
   "metadata": {},
   "source": [
    "### Compute run-start triggered responses (if pre-computed, load below)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "6a71cfaf",
   "metadata": {},
   "outputs": [],
   "source": [
    "import warnings\n",
    "warnings.filterwarnings('ignore')\n",
    "session_list = units[units['no_anomalies']]['ecephys_session_id'].unique()\n",
    "acceleration_peths, deceleration_peths, unitIDs = nwb.unit_averaged_psth_time_aligned(session_list, alignment_time_func = 'running', \n",
    "                                    time_before=0.5, time_after=0.5, binsize=0.001)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "9e1f2c14",
   "metadata": {},
   "outputs": [],
   "source": [
    "running_df = {}\n",
    "for apeths, dpeths, uids in zip(acceleration_peths, deceleration_peths, unitIDs):\n",
    "    if len(uids)>0:\n",
    "        for apeth, dpeth, uid in zip(apeths, dpeths, uids):\n",
    "            running_df[uid] = {'acceleration': apeth, 'deceleration':dpeth}"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "c9f3daa8",
   "metadata": {},
   "outputs": [],
   "source": [
    "import pickle\n",
    "with open(\"/Volumes/programs/mindscope/workgroups/np-behavior/vbn_data_release/supplemental_tables/run_start_aligned_unit_data.pkl\", 'wb') as file:\n",
    "    pickle.dump(running_df, file)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "a6dbdaa1",
   "metadata": {},
   "source": [
    "### Load pre-computed run-start triggered data"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "d1ef687d",
   "metadata": {},
   "outputs": [],
   "source": [
    "import pickle\n",
    "with open(\"/Volumes/programs/mindscope/workgroups/np-behavior/vbn_data_release/supplemental_tables/run_start_aligned_unit_data.pkl\", 'rb') as file:\n",
    "    running_df = pickle.load(file)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "18eaf031",
   "metadata": {},
   "outputs": [],
   "source": [
    "running_aligned_df = pd.DataFrame.from_dict(running_df, orient='index')\n",
    "running_aligned_df = running_aligned_df.merge(units[['unit_id', 'cluster_labels_new', 'structure_acronym', 'cortical_layer']], left_index=True, right_on='unit_id')\n",
    "running_aligned_df.set_index('unit_id', inplace=True)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "88967911",
   "metadata": {},
   "source": [
    "## Figure panels"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "74adfdb2",
   "metadata": {},
   "outputs": [],
   "source": [
    "plt.rcParams['figure.dpi'] = 300"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "a01bc445",
   "metadata": {},
   "outputs": [],
   "source": [
    "from mpl_toolkits.axes_grid1.anchored_artists import AnchoredSizeBar\n",
    "import matplotlib.transforms as transforms\n",
    "import warnings\n",
    "warnings.filterwarnings(\"ignore\")\n",
    "\n",
    "plt.rcParams['font.size'] = 14\n",
    "\n",
    "def add_vertical_size_bar(ax, size, label, loc,\n",
    "                           pad=0, borderpad=0.1, sep=2, \n",
    "                           prop=None, barcolor=\"black\",\n",
    "                            **kwargs):\n",
    "    \"\"\"Adds a vertical scale bar to the axes.\"\"\"\n",
    "    \n",
    "    trans = transforms.Affine2D().rotate_deg(90) + ax.transData\n",
    "    size_bar = AnchoredSizeBar(trans,\n",
    "                             size, label, loc, \n",
    "                             pad=pad,\n",
    "                             color='black',\n",
    "                             frameon=False,\n",
    "                             size_vertical=0,\n",
    "                             )\n",
    "    ax.add_artist(size_bar)\n",
    "\n",
    "\n",
    "base_sub = False\n",
    "session_list = list(active_tensor.keys())\n",
    "plot_window = slice(0, 500)\n",
    "time = np.arange(-50, 750)[plot_window]\n",
    "\n",
    "clusters_to_plot = np.arange(1,9)\n",
    "for cluster in clusters_to_plot:\n",
    "\n",
    "    unit_filter = du.get_units_in_cluster(units, cluster, clustering='new') & \\\n",
    "                        du.apply_unit_quality_filter(units) \n",
    "    unit_ids = units.loc[unit_filter]['unit_id'].values\n",
    "    num_sessions = units.loc[unit_filter]['ecephys_session_id'].unique().size\n",
    "\n",
    "    fig, ax = plt.subplots(1,3)\n",
    "    fig.set_size_inches([7, 2.25])\n",
    "\n",
    "    \n",
    "    # Lick/No-lick stim-aligned comparison\n",
    "    stim_filters = (\n",
    "                ['engaged', '~is_change', '~omitted', '~previous_omitted', \n",
    "                'flashes_since_change>5', 'lickbout_for_flash_during_response_window', '~is_shared'],\n",
    "                \n",
    "                ['engaged', 'is_change', '~omitted', '~previous_omitted', \n",
    "                 '~lickbout_for_flash_during_response_window', '~is_shared'],\n",
    "\n",
    "                ['engaged', 'is_change', '~omitted', '~previous_omitted', \n",
    "                 'lickbout_for_flash_during_response_window', '~is_shared'],\n",
    "                    )\n",
    "    colors = ['k', 'coral', 'g'] \n",
    "    alphas = [0.5, 1, 1] \n",
    "    for stim_filter, color, alpha in zip(stim_filters, colors, alphas):\n",
    "        unit_data, unitIds = vbn_utils.unit_averaged_psth(active_tensor_file, \n",
    "                                    stim_table, \n",
    "                                    session_list, \n",
    "                                    unit_ids, \n",
    "                                    *stim_filter, \n",
    "                                    baseline_length=50, \n",
    "                                    resp_window_length=750)\n",
    "        concat = np.concatenate(unit_data)\n",
    "        concat = np.array([exponential_convolve(c, 3, symmetrical=True) for c in concat])\n",
    "        if concat.size==0:\n",
    "            continue\n",
    "        \n",
    "        if base_sub:\n",
    "            concat = concat - concat[:, :50].mean(axis=1)[:, None]\n",
    "\n",
    "        vbn_utils.mean_sem_plot(concat[:, plot_window]*1000, ax[0], time, color=color, alpha=alpha)\n",
    "\n",
    "\n",
    "    # Lick-aligned comparison\n",
    "    cluster_lick_aligned = lick_aligned_df[(lick_aligned_df['unit_id'].isin(unit_ids))]\n",
    "    means = np.array([c['mean'] for ic, c in cluster_lick_aligned.iterrows()])\n",
    "    shuffle_means = np.array([c['shuffle_mean'] for ic, c in cluster_lick_aligned.iterrows()])\n",
    "    vbn_utils.mean_sem_plot(means*1000, ax[1], np.arange(-500,500), color='k')\n",
    "    vbn_utils.mean_sem_plot(shuffle_means*1000, ax[1], np.arange(-500,500), color='gray', ls='dotted')\n",
    "\n",
    "\n",
    "    # Run aligned comparison\n",
    "    colors=['k', 'gray']\n",
    "    for ic, condition in enumerate(['acceleration', 'deceleration']):\n",
    "        means = running_aligned_df.loc[unit_ids][condition].values\n",
    "        means = np.array([exponential_convolve(m, 3, True) for m in means])\n",
    "        vbn_utils.mean_sem_plot(means, ax[2], np.arange(-500,500), color=colors[ic])\n",
    "        num_nonan = np.sum([~np.isnan(m[0]) for m in means])\n",
    "    print(f'cluster {cluster} count: {len(unit_ids)}; run aligned count: {num_nonan}')\n",
    "    \n",
    "    ymin = np.min([a.get_ylim()[0] for a in ax])\n",
    "    ymax = np.max([a.get_ylim()[1] for a in ax])\n",
    "\n",
    "    yminint = int(np.ceil(ymin))\n",
    "    a = ax[0]\n",
    "    a.spines['left'].set_bounds(yminint, yminint+5)\n",
    "    a.set_yticks([yminint, yminint+5])\n",
    "    a.spines['bottom'].set_bounds(0, 250)\n",
    "    a.set_xticks([0,250])\n",
    "    a.spines['top'].set_visible(False)\n",
    "    a.spines['right'].set_visible(False)\n",
    "    a.tick_params(labelleft=False,)\n",
    "\n",
    "    [a.set_ylim(ymin, ymax) for a in ax]\n",
    "    [a.axis('off') for a in ax[1:]]\n",
    "    [a.axvline(0, color='k', ls='dotted') for a in ax[1:]]\n",
    "\n",
    "\n",
    "    plt.tight_layout()\n",
    "    "
   ]
  },
  {
   "cell_type": "markdown",
   "id": "3e26247e",
   "metadata": {},
   "source": [
    "## GLM dropout analysis"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "87516546",
   "metadata": {},
   "outputs": [],
   "source": [
    "dropouts = pd.read_csv(\"/Volumes/programs/mindscope/workgroups/np-behavior/vbn_data_release/supplemental_tables/GLM_dropout_with_unit_info_active_passive.csv\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "0e271a17",
   "metadata": {},
   "outputs": [],
   "source": [
    "for region in ['VISall', 'VISp', ['LGd', 'LP'], 'SCMRN', 'Hipp']:\n",
    "    good_units = vbn_utils.get_unit_ids(dropouts, region)\n",
    "    region_dropouts = dropouts.set_index('unit_id').loc[good_units]\n",
    "\n",
    "    print(region)\n",
    "    print(region_dropouts[\"('variance_explained_full', 'Full')\"].describe())"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "bbab4f1e",
   "metadata": {},
   "outputs": [],
   "source": [
    "metric = 'absolute_change_from_full'\n",
    "for region in ['VISall', 'SCMRN']:\n",
    "    good_units = vbn_utils.get_unit_ids(dropouts, region)\n",
    "    region_dropouts = dropouts.set_index('unit_id').loc[good_units]\n",
    "\n",
    "    print(region, 'images', np.sum(region_dropouts[f\"('{metric}', 'all-images')\"].values<-0.01)/len(region_dropouts), len(region_dropouts))\n",
    "    print(region, 'licks', np.sum(region_dropouts[f\"('{metric}', 'licks')\"].values<-0.01)/len(region_dropouts))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "8b0113ab",
   "metadata": {},
   "outputs": [],
   "source": [
    "columns_to_plot = [f\"('{metric}', 'all-images')\",\n",
    "                   f\"('{metric}', 'omissions')\",\n",
    "                   f\"('{metric}', 'task')\",\n",
    "                   f\"('{metric}', 'pupil')\",\n",
    "                   f\"('{metric}', 'licks')\",\n",
    "                   f\"('{metric}', 'running')\"\n",
    "                   ]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "c9ecde37",
   "metadata": {},
   "outputs": [],
   "source": [
    "pt = dropouts[dropouts['cluster_labels_new'].isin(np.arange(9))].pivot_table(index='cluster_labels_new', values=columns_to_plot, aggfunc='mean')"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "a5a18464",
   "metadata": {},
   "outputs": [],
   "source": [
    "pt = pt[columns_to_plot]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "449c4c2b",
   "metadata": {},
   "outputs": [],
   "source": [
    "pt_col_normed = pt.div(pt.sum(axis=0), axis=1)\n",
    "pt_row_normed = pt.div(pt.sum(axis=1), axis=0)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "edc1ea69",
   "metadata": {},
   "outputs": [],
   "source": [
    "from matplotlib.pyplot import xticks\n",
    "\n",
    "toplot = pt_col_normed\n",
    "\n",
    "plt.rcParams.update({'font.size': 14})\n",
    "plt.figure(figsize=(5,5))\n",
    "plt.imshow(toplot.values, aspect='equal', cmap='Greys_r')\n",
    "\n",
    "plt.xticks(np.arange(toplot.shape[1]), [c.split(',')[1].replace(')', '').replace(\"'\", '') for c in toplot.columns], rotation=90)\n",
    "plt.yticks(np.arange(toplot.shape[0]), toplot.index.astype(int))\n",
    "plt.ylabel('Cluster')\n",
    "plt.xlabel('Kernel')\n",
    "cbar = plt.colorbar()\n",
    "cbar.set_label('Norm dropout score')"
   ]
  }
 ],
 "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": 5
}
