{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "743554d6",
   "metadata": {},
   "outputs": [],
   "source": [
    "import pandas as pd\n",
    "import numpy as np\n",
    "import os\n",
    "from matplotlib import pyplot as plt\n",
    "import scipy.stats\n",
    "from decoding_utils import apply_condition_filter\n",
    "import vbn_utils\n",
    "import warnings\n",
    "warnings.filterwarnings(\"ignore\")\n",
    "%matplotlib inline"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "ac28c7a6",
   "metadata": {},
   "outputs": [],
   "source": [
    "from notebook_utils import (\n",
    "    beh_mat_from_stim_table,\n",
    "    mean_beh_mat_across_sessions,\n",
    "    get_omission_response_rate,\n",
    "    get_post_omission_response_rate,\n",
    "    get_shared_hit_rate,\n",
    "    get_private_hit_rate,\n",
    "    get_shared_fa_rate,\n",
    "    get_private_fa_rate,\n",
    "    get_shared_nonchange_response_rate,\n",
    "    get_private_nonchange_response_rate,\n",
    "    skip_diag_masking,\n",
    "    get_session_engaged_dprime,\n",
    "    get_session_engaged_hit_count,\n",
    "    get_experience_session_id_for_mouse,\n",
    "    calculate_metric_for_selection,\n",
    "    paired_image_mat_from_stim_table,\n",
    ")"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "9bb57319",
   "metadata": {},
   "source": [
    "## Data loading"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "14ce2bab",
   "metadata": {},
   "outputs": [],
   "source": [
    "#Paths to all of the useful supplemental tables and tensors\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/supplemental_tables/master_sessions_table.csv\""
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "368da5f6",
   "metadata": {},
   "outputs": [],
   "source": [
    "units = pd.read_csv(unit_table_file)\n",
    "unnamedcols = [c for c in units.columns if 'Unnamed' in c]\n",
    "units = units.drop(columns=unnamedcols)\n",
    "\n",
    "stim_table = pd.read_csv(stim_table_file)\n",
    "stim_table = stim_table.drop(columns='Unnamed: 0')\n",
    "\n",
    "sessions_table = pd.read_csv(sessions_table_file)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "8ade2bcd",
   "metadata": {},
   "outputs": [],
   "source": [
    "good_sessions = sessions_table[sessions_table['abnormal_histology'].isnull() & sessions_table['abnormal_activity'].isnull()]\n",
    "good_behavior_sessions = good_sessions[(good_sessions['engaged_dprime']>=1)&(good_sessions['engaged_hitcount']>=50)]"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "acc9f1a3",
   "metadata": {},
   "source": [
    "## Hit rates for holdover images on Familiar and Novel days"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "e1d43b70",
   "metadata": {},
   "outputs": [],
   "source": [
    "plt.rcParams['figure.dpi'] = 150\n",
    "plt.rcParams['savefig.dpi'] = 300\n",
    "plt.rcParams['font.size'] = 10\n",
    "plt.rcParams['pdf.fonttype'] = 42"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "bfc44cb7",
   "metadata": {},
   "outputs": [],
   "source": [
    "import warnings\n",
    "warnings.filterwarnings(\"ignore\")\n",
    "\n",
    "for training_trajectory in ['GGH', 'GHG', 'HHG',]:\n",
    "    beh_mat_dict = {'Familiar':[], 'Novel':[]}\n",
    "    image_set_dict = {}\n",
    "    for experience_level in ['Familiar', 'Novel']:\n",
    "        sessions_to_analyze = good_sessions[apply_condition_filter(good_sessions, experience_level, training_trajectory)]['ecephys_session_id'].unique()\n",
    "\n",
    "        ims, counts, beh_mats = mean_beh_mat_across_sessions(stim_table, sessions_to_analyze)\n",
    "        beh_mat_dict[experience_level] = beh_mats\n",
    "        image_set_dict[experience_level] = ims[0]\n",
    "\n",
    "        \n",
    "    fig, axes = plt.subplots(1,3)\n",
    "    fig.set_size_inches(10,5)\n",
    "    fig.suptitle(f'{training_trajectory}: {len(sessions_to_analyze)} sessions')\n",
    "    for experience_level, ax in zip(['Familiar', 'Novel'], axes[:2]):\n",
    "        beh_mats = beh_mat_dict[experience_level]\n",
    "        im = ax.imshow(np.nanmean(beh_mats, axis=0), clim=[0,1])\n",
    "        ax.set_xlabel('Change Image')\n",
    "        ax.set_ylabel('Pre-change Image')\n",
    "        ax.set_xticks(np.arange(8))\n",
    "        ax.set_xticklabels(image_set_dict[experience_level], rotation=90)\n",
    "        ax.set_yticks(np.arange(8))\n",
    "        ax.set_yticklabels(image_set_dict[experience_level])\n",
    "        \n",
    "    holdover_responses_83 = {'Familiar':[], 'Novel':[]}\n",
    "    holdover_responses_111 = {'Familiar':[], 'Novel':[]}\n",
    "    for experience_level in ['Familiar', 'Novel']: \n",
    "        beh_mats = beh_mat_dict[experience_level]\n",
    "        beh_mats = np.array([skip_diag_masking(b) for b in beh_mats])\n",
    "        holdover_responses_83[experience_level].append(np.nanmean(beh_mats[:, :, 6], axis = 0))\n",
    "        holdover_responses_111[experience_level].append(np.nanmean(beh_mats[:, :, 7], axis = 0))\n",
    "\n",
    "    axes[2].plot(holdover_responses_83['Familiar'], holdover_responses_83['Novel'], 'ko')\n",
    "    axes[2].plot(holdover_responses_111['Familiar'], holdover_responses_111['Novel'], 'ko', markerfacecolor='w', markeredgewidth=2)\n",
    "    axes[2].plot([0,1], [0,1], 'k--')\n",
    "    axes[2].set_aspect('equal')\n",
    "    axes[2].set_xlabel('Hit rate during familiar sessions')\n",
    "    axes[2].set_ylabel('Hit rate during novel sessions')\n",
    "\n",
    "    plt.tight_layout()\n",
    "    fig.savefig(os.path.join(\"/Volumes/programs/mindscope/workgroups/np-behavior/VBN Manuscript\", f'behavioral_response_matrix_{training_trajectory}_test.png'))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "fd82c85c",
   "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"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "a937d877",
   "metadata": {},
   "outputs": [],
   "source": [
    "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,
   "id": "ed7c1911",
   "metadata": {},
   "outputs": [],
   "source": [
    "beh_mat_array = np.array(beh_mat_array)\n",
    "beh_mat_no_diag_array = np.array([skip_diag_masking(b) for b in beh_mat_array])\n",
    "\n",
    "count_array = np.array(count_array)\n",
    "count_no_diag_array = np.array([skip_diag_masking(b) for b in count_array])\n",
    "\n",
    "good_behavior_session_filter = np.array([(get_session_engaged_dprime(stim_table, session_id)>1) & \\\n",
    "                                (get_session_engaged_hit_count(stim_table, session_id)>50) for session_id in session_labels])\n",
    "\n",
    "experience_labels = np.array(experience_labels)\n",
    "session_labels = np.array(session_labels)\n",
    "image_set_labels = np.array(image_set_labels)\n"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "332be96e",
   "metadata": {},
   "source": [
    "## Hit rates on holdover images for Novel days for novel/familiar pre-images"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "25143c58",
   "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",
    "                #print(ind_count, get_session_engaged_dprime(stim_table, session_id), get_session_engaged_hit_count(stim_table, session_id))\n",
    "                imrates.append(np.nanmean(beh_mat_no_diag[:, ind]))\n",
    "            \n",
    "            else:\n",
    "                imrates.append(np.nan)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "776ae232",
   "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'], 'ro')\n",
    "ax.set_aspect('equal')\n",
    "ax.legend(['im083', 'im111'])\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"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "d01b2e45",
   "metadata": {},
   "outputs": [],
   "source": [
    "shared_to_shared_hit_rate = beh_mat_no_diag_array[(good_behavior_session_filter) & (experience_labels=='Novel')][:, 6, 6:]\n",
    "nonshared_to_shared_hit_rate = beh_mat_no_diag_array[(good_behavior_session_filter) & (experience_labels=='Novel')][:, :6, 6:]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "f3fa6eed",
   "metadata": {},
   "outputs": [],
   "source": [
    "null_shared_to_shared = []\n",
    "null_nonshared_to_shared = []\n",
    "for iteration in range(100):\n",
    "    # Shuffle rows of each beh matrix\n",
    "    beh_mat_no_diag_array_shuff = np.array([np.random.permutation(b) for b in beh_mat_no_diag_array[(good_behavior_session_filter) & (experience_labels=='Novel')]])\n",
    "    stos = np.nanmean(beh_mat_no_diag_array_shuff[:, 6, 6:], axis=1)\n",
    "    ntos = np.nanmean(beh_mat_no_diag_array_shuff[:, :6, 6:], axis=(1,2))\n",
    "    null_shared_to_shared.extend(np.nanmean(beh_mat_no_diag_array_shuff[:, 6, 6:], axis=1))\n",
    "    null_nonshared_to_shared.extend(np.nanmean(beh_mat_no_diag_array_shuff[:, :6, 6:], axis=(1,2)))\n",
    "\n",
    "null_shared_to_shared = np.array(null_shared_to_shared)\n",
    "null_nonshared_to_shared = np.array(null_nonshared_to_shared)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "97f47a1d",
   "metadata": {},
   "outputs": [],
   "source": [
    "fig, ax = plt.subplots()\n",
    "ax.plot(null_shared_to_shared, null_nonshared_to_shared, 'ko', alpha=0.002)\n",
    "ax.plot(np.mean(shared_to_shared_hit_rate, axis=1), np.mean(nonshared_to_shared_hit_rate, axis=(1,2)), 'ro')\n",
    "\n",
    "ax.errorbar(np.nanmean(null_shared_to_shared), np.nanmean(null_nonshared_to_shared), \n",
    "                        xerr=np.nanstd(null_shared_to_shared), yerr=np.nanstd(null_nonshared_to_shared), color='k')\n",
    "ax.errorbar(np.nanmean(shared_to_shared_hit_rate), np.nanmean(nonshared_to_shared_hit_rate), \n",
    "                        xerr=np.nanstd(shared_to_shared_hit_rate), yerr=np.nanstd(nonshared_to_shared_hit_rate), color='r')\n",
    "\n",
    "ax.set_aspect('equal')\n",
    "ax.plot([0,1], [0,1], 'k--')\n",
    "ax.set_xlabel('From shared image')\n",
    "ax.set_ylabel('From non-shared image')\n",
    "\n",
    "scipy.stats.ranksums(null_shared_to_shared - null_nonshared_to_shared, np.mean(shared_to_shared_hit_rate, axis=1) - np.mean(nonshared_to_shared_hit_rate, axis=(1,2)), nan_policy='omit')\n"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "6ec94a09",
   "metadata": {},
   "source": [
    "## Hit rates for non-omissions, omissions and post-omissions"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "4ae06f1d",
   "metadata": {},
   "source": [
    "A few notes about the stim_table columns:\n",
    "- a lick bout is defined as any lick that follows the last lick by > 0.5 seconds\n",
    "- `lick_time` is really 'lick bout time'. Not all licks are registered in this column, just the first licks in lick bouts. If you want to see all the licks in a trial, `lick_times` stores them in a (string) list\n",
    "- `lick_for_flash` indicates whether a lick bout initiated after the start time of a flash and before the next flash. In the next cell we will add `lick_for_flash_during_response_window`, which will indicate whether a lick bout started in the response window after a flash (100-750 ms after flash start)\n",
    "- `first_lick_in_trial` should be named 'first lick bout in trial'. There are some edge cases when consummatory licks from the last trial continue on into the next trial, but this trial doesn't abort immediately. In these cases, you can have a lick bout that starts after these consummatory licks, and `first_lick_in_trial` will be True, but `before_first_trial_lick` will be false (since consummatory licks happened earlier in the trial).\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "72e3f94c",
   "metadata": {},
   "outputs": [],
   "source": [
    "response_rate_summary = {s:{col: np.nan for col in ['G_Familiar_private_nonchange', 'G_Familiar_shared_nonchange', \n",
    "                                                    'G_Novel_private_nonchange', 'G_Novel_shared_nonchange',\n",
    "                                                    'H_Familiar_private_nonchange', 'H_Familiar_shared_nonchange', \n",
    "                                                    'H_Novel_private_nonchange', 'H_Novel_shared_nonchange',\n",
    "                                                    'G_Familiar_private_hit', 'G_Familiar_shared_hit', \n",
    "                                                    'G_Novel_private_hit', 'G_Novel_shared_hit',\n",
    "                                                    'H_Familiar_private_hit', 'H_Familiar_shared_hit', \n",
    "                                                    'H_Novel_private_hit', 'H_Novel_shared_hit',\n",
    "                                                    'G_Familiar_private_fa', 'G_Familiar_shared_fa', \n",
    "                                                    'G_Novel_private_fa', 'G_Novel_shared_fa',\n",
    "                                                    'H_Familiar_private_fa', 'H_Familiar_shared_fa', \n",
    "                                                    'H_Novel_private_fa', 'H_Novel_shared_fa',\n",
    "                                                    'G_Familiar_omission', 'G_Novel_omission', \n",
    "                                                    'H_Familiar_omission', 'H_Novel_omission', \n",
    "                                                    'G_Familiar_postomission', 'G_Novel_postomission', \n",
    "                                                    'H_Familiar_postomission', 'H_Novel_postomission']} for s in good_behavior_sessions['ecephys_session_id'].values}\n",
    "\n",
    "for isess, session in good_behavior_sessions.iterrows():\n",
    "    \n",
    "    experience_level = session['experience_level']\n",
    "    image_set = session['image_set']\n",
    "    session_id = session['ecephys_session_id']\n",
    "    session_stim_table = stim_table[stim_table['session_id']==session_id]\n",
    "    session_trials = session_stim_table.groupby('behavior_trial_id').head(1)\n",
    "\n",
    "    response_rate_summary[session_id][image_set + '_' + experience_level + '_' + 'private_nonchange'] = get_private_nonchange_response_rate(session_stim_table)\n",
    "    response_rate_summary[session_id][image_set + '_' + experience_level + '_' + 'private_hit'] = get_private_hit_rate(session_trials)\n",
    "    response_rate_summary[session_id][image_set + '_' + experience_level + '_' + 'private_fa'] = get_private_fa_rate(session_trials)\n",
    "    response_rate_summary[session_id][image_set + '_' + experience_level + '_' + 'shared_nonchange'] = get_shared_nonchange_response_rate(session_stim_table)\n",
    "    response_rate_summary[session_id][image_set + '_' + experience_level + '_' + 'shared_hit'] = get_shared_hit_rate(session_trials)\n",
    "    response_rate_summary[session_id][image_set + '_' + experience_level + '_' + 'shared_fa'] = get_shared_fa_rate(session_trials)\n",
    "    response_rate_summary[session_id][image_set + '_' + experience_level + '_' + 'omission'] = get_omission_response_rate(session_stim_table)\n",
    "    response_rate_summary[session_id][image_set + '_' + experience_level + '_' + 'postomission'] = get_post_omission_response_rate(session_stim_table)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "1f407cfe",
   "metadata": {},
   "outputs": [],
   "source": [
    "response_rate_df = pd.DataFrame.from_dict(response_rate_summary, orient='index')"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "108c0766",
   "metadata": {},
   "outputs": [],
   "source": [
    "aggregated_labels = ['Familiar_shared', 'Familiar_private', 'Novel_shared', 'Novel_private', 'Novel', 'Familiar']\n",
    "response_categories = ['nonchange', 'hit',]# 'fa', 'omission', 'postomission']\n",
    "\n",
    "cols_to_aggregate = []\n",
    "for resp in response_categories:\n",
    "    for ag_label in aggregated_labels:\n",
    "\n",
    "        cols = [c for c in response_rate_df.columns if ag_label + '_' + resp in c]\n",
    "        if len(cols)>0:\n",
    "            cols_to_aggregate.append(cols)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "37edc4fc",
   "metadata": {},
   "outputs": [],
   "source": [
    "labels = []\n",
    "figure_labels = []\n",
    "values = []\n",
    "means = []\n",
    "sems = []\n",
    "for cols in cols_to_aggregate:\n",
    "    figure_labels.append(cols[0][2:].replace('shared', 'session \\n holdover').replace('private', '').replace('_', ' ').replace('hit', 'change'))\n",
    "\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)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "13536446",
   "metadata": {},
   "outputs": [],
   "source": [
    "plt.rcParams['font.size'] = 16\n",
    "fig, ax = plt.subplots()\n",
    "for ivals, vals in enumerate(values):\n",
    "    color = 'b' if 'Familiar' in labels[ivals] else 'r'\n",
    "    ax.boxplot(vals, positions=[ivals+1], widths=0.5, showfliers=False, whis=(10, 90), notch=True, boxprops={'color': color, 'linewidth': 2}, medianprops={'color': color, 'linewidth': 2}, whiskerprops={'color': color, 'linewidth': 2}, capprops={'color': color, 'linewidth': 2})\n",
    "\n",
    "ax.set_xticklabels(figure_labels, rotation=90)\n",
    "vbn_utils.formatFigure(fig, ax, yLabel='Response rate')\n",
    "fig.savefig(\"/Volumes/programs/mindscope/workgroups/np-behavior/VBN Manuscript/behavioral_response_rate_boxplot_test.png\")\n"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "70c4d21f",
   "metadata": {},
   "source": [
    "### Statistical comparisons of response rates"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "90c5ce4f",
   "metadata": {},
   "outputs": [],
   "source": [
    "fig, ax = plt.subplots()\n",
    "fig.set_size_inches(7,7)\n",
    "vbn_utils.plot_comparison_matrix(*values, colorbar=True, ax=ax, cmap='PiYG', binarize=True)\n",
    "ax.set_xticklabels(figure_labels, rotation=90)\n",
    "ax.set_yticklabels(figure_labels,)"
   ]
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "vbn_manuscript",
   "language": "python",
   "name": "python3"
  },
  "language_info": {
   "name": "python",
   "version": "3.8.20"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 5
}
