Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1,115 changes: 948 additions & 167 deletions dj_pipeline/notebooks/NP_figures/dual_occlusion.ipynb

Large diffs are not rendered by default.

373 changes: 291 additions & 82 deletions dj_pipeline/notebooks/NP_figures/dual_regression.ipynb

Large diffs are not rendered by default.

1,046 changes: 916 additions & 130 deletions dj_pipeline/notebooks/NP_figures/multi_occlusion.ipynb

Large diffs are not rendered by default.

6 changes: 2 additions & 4 deletions dj_pipeline/notebooks/NP_figures/multi_regression.ipynb
Original file line number Diff line number Diff line change
Expand Up @@ -4,9 +4,7 @@
"cell_type": "markdown",
"metadata": {},
"source": [
"# [Decision Variable Prediction] Figure 5\n",
"\n",
"## The amount of information available to the mouse correlates inversely with infotaxis behavior."
"# The amount of information available to the mouse correlates inversely with infotaxis behavior."
]
},
{
Expand Down Expand Up @@ -67,7 +65,7 @@
"\n",
"style()\n",
"\n",
"save_fig_path = \"notebooks/Paper_figures/Figure_output/\""
"save_fig_path = \"notebooks/NP_figures/NP_figures/NP_\""
]
},
{
Expand Down
87 changes: 44 additions & 43 deletions dj_pipeline/notebooks/Paper_figures/Figure_4.ipynb
Original file line number Diff line number Diff line change
Expand Up @@ -3805,7 +3805,7 @@
"# interpolated_df_all[\"aperture\"] = interpolated_df_all[\"aperture\"].astype(float)\n",
"\n",
"# df_model_all, coef = regression.predict_decision(\n",
"# df=interpolated_df_all, label=model_labels, per_mouse=True\n",
"# df=interpolated_df_all, label=model_labels, per_session=True\n",
"# )"
]
},
Expand Down Expand Up @@ -3845,7 +3845,7 @@
" window_pred, window_coef, window_scalers = regression.predict_decision(\n",
" df=window_df,\n",
" label=model_labels,\n",
" per_mouse=True,\n",
" per_session=True,\n",
" )\n",
"\n",
" window_pred[\"model_idx\"] = window_id\n",
Expand All @@ -3865,17 +3865,20 @@
},
{
"cell_type": "code",
"execution_count": 89,
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"decision_points_all = regression.find_decision_point(df_model_all, \n",
" threshold_uncertainty=0.2)"
"# interpolated_df_all[\"aperture\"] = interpolated_df_all[\"aperture\"].astype(float)\n",
"\n",
"# df_model_all, coef = regression.predict_decision(\n",
"# df=interpolated_df_all, label=model_labels, per_session=True\n",
"# )"
]
},
{
"cell_type": "code",
"execution_count": 90,
"execution_count": null,
"metadata": {},
"outputs": [
{
Expand All @@ -3898,48 +3901,46 @@
}
],
"source": [
"fig, ax = plt.subplots(1, 1, figsize=(3, 5))\n",
"WINDOW_COUNT = 10\n",
"\n",
"for aperture, color, ypos in zip(decision_points_all.sort_values(\"aperture\").aperture.unique(), \n",
" plotting.colors_aperture,\n",
" [0.10, 0.05]):\n",
" x = abs(decision_points_all[decision_points_all.aperture==aperture].groupby(\n",
" \"mouse_name\")[\"y\"].mean().values - 27)\n",
" y = trial_df_all[(trial_df_all.aperture==aperture)].groupby(\n",
" \"mouse_name\")[\"trial_rewarded\"].mean().values\n",
"\n",
"\n",
" # scatter\n",
" ax.scatter(\n",
" x, y,\n",
" color=color,\n",
" alpha=0.7,\n",
" s=50\n",
" )\n",
"# Build trial-progress windows like PredictionModel10Windows\n",
"df_multi = interpolated_df_all.copy()\n",
"df_multi[\"aperture\"] = df_multi[\"aperture\"].astype(float)\n",
"df_multi = df_multi.sort_values([\"dataset\", \"trial\", \"trial_length\"]).reset_index(drop=True)\n",
"\n",
" # regression line\n",
" slope, intercept, r_value, p_value, std_err = stats.linregress(x, y)\n",
" x_fit = np.linspace(x.min(), x.max(), 100)\n",
" y_fit = slope * x_fit + intercept\n",
" ax.plot(x_fit, y_fit, color=color, linestyle=\"--\")\n",
" \n",
" print(f\"Aperture {aperture}: r={r_value}, p={p_value}\")\n",
"\n",
" # annotate correlation\n",
" ax.text(\n",
" 0.05, ypos,\n",
" f\"r={r_value:.2f}, p={p_value:.3g}\",\n",
" transform=ax.transAxes,\n",
" va=\"top\", ha=\"left\",\n",
" color=color\n",
"trial_index = df_multi.groupby([\"dataset\", \"trial\"]).cumcount()\n",
"trial_size = df_multi.groupby([\"dataset\", \"trial\"])[\"trial\"].transform(\"size\")\n",
"df_multi[\"trial_window\"] = np.clip(\n",
" (trial_index * WINDOW_COUNT / trial_size).astype(int),\n",
" 0,\n",
" WINDOW_COUNT - 1,\n",
")\n",
"\n",
"# Train one model per window (LOGO across datasets)\n",
"df_model_parts = []\n",
"coef = {}\n",
"scalers_by_window = {}\n",
"\n",
"for window_id in range(WINDOW_COUNT):\n",
" window_df = df_multi[df_multi[\"trial_window\"] == window_id].copy()\n",
" if window_df.empty:\n",
" coef[window_id] = np.array([[np.nan]])\n",
" scalers_by_window[window_id] = []\n",
" continue\n",
"\n",
" window_pred, window_coef, window_scalers = regression.predict_decision(\n",
" df=window_df,\n",
" label=model_labels,\n",
" per_session=True,\n",
" )\n",
"\n",
" ax.set_xlabel(\"Distance to screen (cm)\")\n",
" ax.set_ylabel(\"Success rate\")\n",
" ax.set_ylim(0, 1.0)\n",
" sns.despine(offset=10, ax=ax)\n",
" window_pred[\"model_idx\"] = window_id\n",
" df_model_parts.append(window_pred)\n",
"\n",
" coef[window_id] = window_coef\n",
" scalers_by_window[window_id] = window_scalers\n",
"\n",
"plt.savefig(save_fig_path + \"figure4_decision_point_vs_performance.svg\", transparent=True)"
"df_model_all = pd.concat(df_model_parts, ignore_index=True)"
]
},
{
Expand Down
4 changes: 2 additions & 2 deletions dj_pipeline/notebooks/shape-analysis.ipynb
Original file line number Diff line number Diff line change
Expand Up @@ -1106,7 +1106,7 @@
"df_model, coef, _ = regression.predict_decision(\n",
" df=interpolated_j_shaped,\n",
" label=model_labels,\n",
" per_mouse=True,\n",
" per_session=True,\n",
" scale_data=True,\n",
" max_iter=500,\n",
")"
Expand Down Expand Up @@ -2062,7 +2062,7 @@
"df_model, coef, _ = regression.predict_decision(\n",
" df=interpolated_j_shaped,\n",
" label=model_labels,\n",
" per_mouse=True,\n",
" per_session=True,\n",
" scale_data=True,\n",
" max_iter=500,\n",
")"
Expand Down
20 changes: 16 additions & 4 deletions dj_pipeline/vr4mice/analysis/regression.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import copy
import warnings
from typing import List, Optional, Tuple, Union

import matplotlib as mpl
Expand Down Expand Up @@ -49,10 +50,11 @@ def predict_decision(
df,
label: List[str],
n_splits: int = 10,
per_mouse: bool = True,
per_session: bool = True,
max_iter: int = 100,
scale_data: bool = True,
random_state: Optional[int] = None,
per_mouse: Optional[bool] = None,
) -> Tuple[pd.DataFrame, npt.NDArray, List[Optional[dict]]]:

"""Predict the animal's decision based on the `label` data, through a logistic regression.
Expand All @@ -61,11 +63,12 @@ def predict_decision(
df: The dataframe.
label: A list of column names in the `df` dataframe.
n_splits: The number of splits fo the cross validation.
per_mouse: If `True` split the data per session, else split
randomly across all sessions. If per_mouse, we train on
per_session: If `True` split the data per session, else split
randomly across all sessions. If per_session, we train on
all sessions but one, and test on the left out session.
scale_data: If `True`, standardize the data before fitting the model.
random_state: Random state for reproducibility.
per_mouse: Deprecated alias for ``per_session`` kept for backward compatibility.

Returns:
A tuple of the input dataframe with added ``accuracy`` and ``proba_left`` columns,
Expand All @@ -86,6 +89,15 @@ def predict_decision(
```
"""

if per_mouse is not None:
warnings.warn(
"'per_mouse' is deprecated and will be removed in a future release; "
"use 'per_session' instead.",
DeprecationWarning,
stacklevel=2,
)
per_session = per_mouse

data = np.asarray(df[label].values)
labels = df.trial_left_choice.values

Expand All @@ -98,7 +110,7 @@ def predict_decision(
max_iter=max_iter, random_state=random_state
)

if per_mouse:
if per_session:
sessions = df.dataset.values
coefs = np.empty((len(np.unique(sessions)), n_features + 1))
scalers = []
Expand Down
8 changes: 4 additions & 4 deletions dj_pipeline/vr4mice/schema/decision.py
Original file line number Diff line number Diff line change
Expand Up @@ -423,7 +423,7 @@ class ModelParams(dj.Lookup):

@schema
class PredictionModel(dj.Computed):
"""Train logistic regression model per mouse using LOGO cross-validation."""
"""Train logistic regression model with leave-one-session-out cross-validation."""

definition = """
-> LabelSet
Expand All @@ -432,7 +432,7 @@ class PredictionModel(dj.Computed):
-> ExperimentStage
-> vr4mice.Batch
---
coefficients : <blob> # coefficients per session (per_mouse=True)
coefficients : <blob> # coefficients per session (per_session=True)
n_sessions : int32 # number of sessions included
sessions : <blob> # list of session dataset names
random_state : int32 # random state used for reproducibility
Expand Down Expand Up @@ -553,7 +553,7 @@ def make(self, key):
df_model, coef, _ = regression.predict_decision(
df=interpolated_df,
label=label_set,
per_mouse=True,
per_session=True,
max_iter=params["max_iter"],
scale_data=params["scale_data"],
random_state=random_state,
Expand Down Expand Up @@ -873,7 +873,7 @@ def make(self, key):
df_model, coef, scalers = regression.predict_decision(
df=window_df,
label=label_set,
per_mouse=True,
per_session=True,
max_iter=params["max_iter"],
scale_data=params["scale_data"],
random_state=random_state,
Expand Down
Loading