{
 "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 pathlib import Path\n",
    "from matplotlib import pyplot as plt\n",
    "from vbn_utils import formatFigure\n",
    "import vbn_utils as vbn\n",
    "import ccf_utils\n",
    "%matplotlib inline\n",
    "from functools import partial\n",
    "import decoding_utils as du\n",
    "from analysis_utils import exponential_convolve\n",
    "from scipy.stats import binned_statistic_2d\n",
    "from scipy.optimize import curve_fit"
   ]
  },
  {
   "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",
    "unit_table_with_rf_stats = \"/Volumes/programs/mindscope/workgroups/np-behavior/vbn_data_release/supplemental_tables/units_with_rf_stats.csv\"\n",
    "\n",
    "sessions_table_file = \"/Volumes/programs/mindscope/workgroups/np-behavior/vbn_data_release/supplemental_tables/master_sessions_table.csv\""
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "stims = pd.read_csv(stim_table_file)\n",
    "units = pd.read_csv(unit_table_file)\n",
    "structure_tree = pd.read_csv(\"/Volumes/programs/mindscope/workgroups/np-behavior/ccf_structure_tree_2017.csv\")\n",
    "units_rf = pd.read_csv(unit_table_with_rf_stats)\n",
    "rf_cols = [c for c in units_rf.columns if 'rf' in c]\n",
    "units = units.merge(units_rf[rf_cols + ['unit_id']], on='unit_id', how='left')\n",
    "rf_array_dir = \"/Volumes/programs/mindscope/workgroups/np-behavior/vbn_data_release/supplemental_tables/rfs/arrays\"\n",
    "all_rf_files = os.listdir(rf_array_dir)\n",
    "sessions = pd.read_csv(sessions_table_file)"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## put some basic numbers on dataset"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "quality = du.apply_unit_quality_filter(units, no_abnorm=True)\n",
    "quality_units = units[quality]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "quality = du.apply_unit_quality_filter(units)\n",
    "inregion = du.getUnitsInRegion(units, 'VISall')\n",
    "ctx = units.loc[quality&inregion]\n",
    "\n",
    "df = {'cell type': ['RS', 'FS', 'SST', 'VIP'], \n",
    "        'Criteria': ['Not optotagged; spike duration > 0.4 ms', \n",
    "                                                            'Not optotagged; spike duration < 0.4 ms',\n",
    "                                                            'Optotagged in SST-cre x Ai32 mouse',\n",
    "                                                            'Optotagged in VIP-cre x Ai32 mouse'], \n",
    "        'Count': [ctx['RS'].sum(), ctx['FS'].sum(), ctx['SST'].sum(), ctx['VIP'].sum()]}\n",
    "\n",
    "df = pd.DataFrame(df)\n",
    "display(df)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "mice = sessions.groupby('mouse_id').head(1).reset_index(drop=True)\n",
    "mice.value_counts(['genotype', 'sex'])"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## population RFs across sessions"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from notebook_utils import gaussian_2d"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "no_abnorm_sessions = sessions[sessions['abnormal_activity'].isnull() & sessions['abnormal_histology'].isnull()]['ecephys_session_id'].values\n",
    "quality = du.apply_unit_quality_filter(units, no_abnorm=True)\n",
    "\n",
    "rf_centers = {area: [] for area in ['VISp', 'VISl', 'VISal', 'VISpm', 'VISam', 'VISrl', 'LGd','LP']}\n",
    "for isess, session_id in enumerate(no_abnorm_sessions):\n",
    "    print(isess)\n",
    "    session_filter = units['ecephys_session_id'] == session_id\n",
    "    for area in ['VISp', 'VISl', 'VISal', 'VISpm', 'VISam', 'VISrl', 'LGd','LP']:\n",
    "        in_area = du.getUnitsInRegion(units, area)\n",
    "        sig_rf = units['p_value_rf'] < 0.05\n",
    "        on_screen = units['on_screen_rf']\n",
    "\n",
    "        selected = units[in_area & sig_rf & on_screen & quality & session_filter]\n",
    "\n",
    "        selected_rf_files= [r for r in all_rf_files if int(r.split('.npy')[0]) in selected['unit_id'].values]\n",
    "        if len(selected_rf_files) == 0:\n",
    "            continue\n",
    "        selected_rfs = []\n",
    "        for rf_file in selected_rf_files:\n",
    "            rf = np.load(os.path.join(rf_array_dir, rf_file))\n",
    "            selected_rfs.append(rf)\n",
    "\n",
    "        pop_rf = np.mean(selected_rfs, axis=0)\n",
    "        maxloc = np.unravel_index(np.argmax(pop_rf), pop_rf.shape)\n",
    "\n",
    "        initial_guess = (pop_rf.max(), maxloc[1], maxloc[0], 2, 2, 0, pop_rf.min())\n",
    "\n",
    "        try:\n",
    "            fit_params, pcov = curve_fit(gaussian_2d, np.meshgrid(range(9), range(9)), pop_rf.ravel(), p0=initial_guess, maxfev=10000)\n",
    "\n",
    "            rf_centers[area].append((fit_params[1]*10+10, 50-fit_params[2]*10)) #transform to degrees\n",
    "        except:\n",
    "            print(f'Failed to fit {area} in session {session_id}')\n",
    "            continue\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from matplotlib.ticker import MaxNLocator\n",
    "plt.rcParams['font.size'] = 14\n",
    "\n",
    "figs = []\n",
    "axes = []\n",
    "for area in ['VISp', 'VISl', 'VISal', 'VISpm', 'VISam', 'VISrl', 'LGd','LP']:\n",
    "    rf_centers[area] = np.array(rf_centers[area])\n",
    "    fig, ax = plt.subplots(1,2)\n",
    "    fig.set_size_inches(8, 4)\n",
    "    binnedarray, xedges, yedges, bins = binned_statistic_2d(rf_centers[area][:,0], rf_centers[area][:,1], \n",
    "                                                            rf_centers[area][:,0], statistic='count',\n",
    "                                                            bins=[range(5, 100, 10), range(-35, 60, 10)])\n",
    "    binnedarray = binnedarray/rf_centers[area].shape[0]\n",
    "    clims = [0, 0.25] if area == 'VISp' else [0, 0.15]\n",
    "    im = ax[1].imshow(binnedarray.T, extent=[5, 95, -35, 55],cmap='Greys', origin='lower')\n",
    "\n",
    "    print(f'{area}: {rf_centers[area].shape[0]} sessions')\n",
    "    ax[0].plot(rf_centers[area][:,0], rf_centers[area][:,1], 'ko', alpha=0.5)\n",
    "    ax[0].set_xlim(0, 100)\n",
    "    ax[0].set_ylim(-40, 60)\n",
    "\n",
    "    ax[1].tick_params(left=False, right=False, bottom=False, top=False, labelleft=False, labelbottom=False)\n",
    "\n",
    "\n",
    "    colorbar = fig.colorbar(im, orientation='horizontal', pad=0.1, aspect=25)  # Horizontal colorbar\n",
    "    colorbar.ax.set_position([ax[1].get_position().x0, ax[1].get_position().y1, ax[1].get_position().width, 0.05])  # Adjust position\n",
    "    colorbar.ax.xaxis.set_ticks_position('top')  # Move ticks to the top\n",
    "    colorbar.ax.xaxis.set_label_position('top')  # Move label to the top\n",
    "    colorbar.locator = MaxNLocator(nbins=3)  # Automatically determine 3 ticks\n",
    "    colorbar.update_ticks()  # Update the ticks after setting the locator\n",
    "\n",
    "    axes.append(ax)\n",
    "    figs.append(fig)\n",
    "\n",
    "for ia, area in enumerate(['VISp','VISl', 'VISal', 'VISpm', 'VISam', 'VISrl', 'LGd','LP']):\n",
    "    figs[ia].savefig(os.path.join(\"/Volumes/programs/mindscope/workgroups/np-behavior/VBN Manuscript/Figure 1\", f'{area}_session_rfs_test.pdf'))\n"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## area unit counts"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "quality = du.apply_unit_quality_filter(units, no_abnorm=True)\n",
    "quality_units = units[quality]\n",
    "\n",
    "quality_units = quality_units[quality_units['brain_division'] != 'not in list']\n",
    "quality_units.loc[quality_units['structure_acronym'].isin(['SCig', 'SCiw']).astype(bool), 'structure_acronym'] = 'SCm'\n",
    "quality_units.loc[quality_units['structure_acronym'].isin(['MGv', 'MGd', 'MGm']).astype(bool), 'structure_acronym'] = 'MG'\n",
    "\n",
    "areas_meet_threshold = quality_units['structure_acronym'].value_counts() > 1000\n",
    "areas_meet_threshold = areas_meet_threshold[areas_meet_threshold].index\n",
    "\n",
    "quality_units = quality_units[quality_units['structure_acronym'].isin(areas_meet_threshold)]\n",
    "quality_units['structure_acronym'].value_counts()\n",
    "brain_divisions = ['Thalamus', 'Isocortex', 'Hippocampal formation', 'Midbrain', 'Hypothalamus']\n",
    "\n",
    "fig, ax = plt.subplots()\n",
    "area_count = 0\n",
    "area_labels = []\n",
    "for bd in brain_divisions:\n",
    "    bd_units = quality_units[quality_units['brain_division'] == bd]\n",
    "    bd_areas = np.sort(bd_units['structure_acronym'].unique())\n",
    "    for area in bd_areas:\n",
    "        area_units = bd_units[bd_units['structure_acronym'] == area]\n",
    "        ax.bar(area_count, len(area_units), color=ccf_utils.get_area_color(area, structure_tree))\n",
    "        area_count += 1\n",
    "        area_labels.append(area)\n",
    "\n",
    "ax.set_xticks(np.arange(len(area_labels)))\n",
    "ax.set_xticklabels(area_labels, rotation=90)\n",
    "formatFigure(fig, ax)\n",
    "ax.set_ylabel('Number of units')\n",
    "fig.savefig(\"/Volumes/programs/mindscope/workgroups/np-behavior/VBN Manuscript/Figure 1/units_per_area_test.pdf\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## behavior characterization"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "good_sessions = sessions[sessions['abnormal_activity'].isnull() & sessions['abnormal_histology'].isnull()]\n",
    "good_stims = stims[stims['session_id'].isin(good_sessions['ecephys_session_id'].values)]\n",
    "good_stims = good_stims.merge(good_sessions[['ecephys_session_id', 'mouse_id']], left_on='session_id', right_on='ecephys_session_id')\n",
    "\n",
    "mouse_fas = {}\n",
    "mouse_hits = {}\n",
    "mouse_dprimes = {}\n",
    "for mouse in good_stims['mouse_id'].unique():\n",
    "    mouse_stims = good_stims[good_stims['mouse_id'] == mouse]\n",
    "    mouse_trials = mouse_stims.groupby('session_id').apply(lambda x: x.groupby('behavior_trial_id').first())\n",
    "    if (mouse_trials['engaged']&mouse_trials['catch']).sum() > 5:\n",
    "        mouse_fas[mouse] = (mouse_trials['engaged']&mouse_trials['false_alarm']).sum()/(mouse_trials['engaged']&mouse_trials['catch']).sum()\n",
    "        mouse_hits[mouse] = (mouse_trials['engaged']&mouse_trials['hit']).sum()/(mouse_trials['engaged']&mouse_trials['go']).sum()\n",
    "        mouse_dprimes[mouse] = vbn.calcDprime((mouse_trials['engaged']&mouse_trials['hit']).sum(),\n",
    "                                              (mouse_trials['engaged']&mouse_trials['miss']).sum(),\n",
    "                                              (mouse_trials['engaged']&mouse_trials['false_alarm']).sum(),\n",
    "                                              (mouse_trials['engaged']&mouse_trials['correct_reject']).sum())\n",
    "    else:\n",
    "        print(f'Mouse {mouse} does not have enough trials to calculate d-prime: {(mouse_trials[\"engaged\"]&mouse_trials[\"catch\"]).sum()} catch trials')"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "fig, axes = plt.subplots(1,2)\n",
    "fig.set_size_inches(5, 4)\n",
    "hits = []\n",
    "fas = []\n",
    "jitters = []\n",
    "for mouse in mouse_fas.keys():\n",
    "    jitter = np.random.rand()*0.2 - 0.1\n",
    "    jitters.append(jitter)\n",
    "    \n",
    "    mhit = mouse_hits[mouse]\n",
    "    mfa = mouse_fas[mouse]\n",
    "\n",
    "    hits.append(mhit)\n",
    "    fas.append(mfa)\n",
    "    axes[0].plot([0+jitter, 1+jitter], [mhit, mfa], '-', color='k',alpha=0.5)\n",
    "\n",
    "for im, mouse in enumerate(mouse_fas.keys()):\n",
    "    jitter = jitters[im]\n",
    "    \n",
    "    mhit = mouse_hits[mouse]\n",
    "    mfa = mouse_fas[mouse]\n",
    "\n",
    "    axes[0].plot([0+jitter, 1+jitter], [mhit, mfa], 'wo', mec='k',alpha=1, ms=10)\n",
    "\n",
    "axes[0].set_xlim(-0.3,1.3)\n",
    "axes[0].errorbar([0, 1], [np.mean(hits), np.mean(fas)], yerr=[np.std(hits)/(len(hits)**0.5), np.std(fas)/(len(fas)**0.5)], color='r', fmt='_', ms=10)\n",
    "axes[0].set_title(f'{scipy.stats.wilcoxon(hits, fas)} \\n n = {len(mouse_fas.keys())}')\n",
    "axes[0].set_xticks([0,1])\n",
    "axes[0].set_xticklabels(['Hits', 'False Alarms'])\n",
    "formatFigure(fig, axes[0], xLabel='Response Type', yLabel='Response Rate')\n",
    "\n",
    "\n",
    "for im, mouse in enumerate(mouse_fas.keys()):\n",
    "    jitter = jitters[im]\n",
    "    axes[1].plot([0+jitter], [mouse_dprimes[mouse]], 'wo', mec='k',alpha=1, ms=10)\n",
    "\n",
    "dprimes = [d for m,d in mouse_dprimes.items()]\n",
    "\n",
    "axes[1].set_xlim(-0.5,0.5)\n",
    "axes[1].set_xticks([0])\n",
    "axes[1].set_ylim(0, 3)\n",
    "axes[1].errorbar([0], [np.mean(dprimes)], yerr=[np.std(dprimes)/(len(dprimes)**0.5)], color='r', fmt='_', ms=10)\n",
    "\n",
    "formatFigure(fig, axes[1], yLabel='dprime')\n",
    "\n",
    "plt.tight_layout()\n",
    "fig.savefig(os.path.join(\"/Volumes/programs/mindscope/workgroups/np-behavior/VBN Manuscript/Figure 1\", 'hit_fa_rates.pdf'))\n",
    "\n",
    "np.mean(dprimes), np.std(dprimes)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "behavior_sessions = pd.read_csv(\"/Volumes/programs/mindscope/workgroups/np-behavior/vbn_data_release/vbn_s3_cache/visual-behavior-neuropixels-0.5.0/project_metadata/behavior_sessions.csv\")\n",
    "\n",
    "mice = good_sessions['mouse_id'].unique()\n",
    "\n",
    "mouse_passing_sessions = {}\n",
    "for mouse in mice:\n",
    "    mouse_sessions = behavior_sessions[behavior_sessions['mouse_id'] == mouse]\n",
    "\n",
    "    static_grating_sessions = np.sum(mouse_sessions['session_type'].str.contains('TRAINING_1'))\n",
    "    flashed_grating_sessions = np.sum(mouse_sessions['session_type'].str.contains('TRAINING_2'))\n",
    "    images_sessions =np.sum((mouse_sessions['session_type'].str.contains('TRAINING_3')) | \\\n",
    "                         (mouse_sessions['session_type'].str.contains('TRAINING_4')) | \\\n",
    "                        ((mouse_sessions['session_type'].str.contains('TRAINING_5')) & (mouse_sessions['session_type'].str.contains('epilogue'))))\n",
    "    \n",
    "    mouse_passing_sessions[mouse] = {'static_grating': static_grating_sessions, 'flashed_grating': flashed_grating_sessions, 'images': images_sessions}\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "fig, ax = plt.subplots()\n",
    "jitters = []\n",
    "\n",
    "for mouse in mouse_passing_sessions.keys():\n",
    "    jitter = np.random.rand()*0.3 - 0.15\n",
    "    jitters.append(jitter)\n",
    "    ax.plot(np.arange(3)+jitter, [mouse_passing_sessions[mouse]['static_grating'], \n",
    "                        mouse_passing_sessions[mouse]['flashed_grating'], \n",
    "                        mouse_passing_sessions[mouse]['images']], '-', color='k',alpha=0.5)\n",
    "\n",
    "for im, mouse in enumerate(mouse_passing_sessions.keys()):\n",
    "    jitter = jitters[im]\n",
    "    ax.plot(np.arange(3)+jitter, [mouse_passing_sessions[mouse]['static_grating'], \n",
    "                        mouse_passing_sessions[mouse]['flashed_grating'], \n",
    "                        mouse_passing_sessions[mouse]['images']], 'wo',mec='k',alpha=1, ms=10)\n",
    "\n",
    "\n",
    "ax.errorbar([0, 1, 2], [np.mean([mouse_passing_sessions[mouse][stage] for mouse in mouse_passing_sessions.keys()])\n",
    "                        for stage in ['static_grating', 'flashed_grating', 'images']], \n",
    "                    yerr=[np.std([mouse_passing_sessions[mouse][stage] for mouse in mouse_passing_sessions.keys()])/len(mouse_passing_sessions.keys())**0.5\n",
    "                        for stage in ['static_grating', 'flashed_grating', 'images']], color='r', fmt='_', ms=10)\n",
    "\n",
    "\n",
    "formatFigure(fig, ax)\n",
    "fig.savefig(\"/Volumes/programs/mindscope/workgroups/np-behavior/VBN Manuscript/Figure 1/training_times_test.pdf\")\n",
    "\n",
    "total_sessions_to_pass = []\n",
    "for mouse, mdata in mouse_passing_sessions.items():\n",
    "    total_sessions_to_pass.append(mdata['static_grating'] + mdata['flashed_grating'] + mdata['images'])\n",
    "np.mean(total_sessions_to_pass), np.std(total_sessions_to_pass)"
   ]
  }
 ],
 "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
}
