ic-placelearning 0.0.1__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,11 @@
1
+ """IntelliCage Place Learning Toolkit.
2
+
3
+ The package provides reusable utilities to read IntelliCage exports, harmonize
4
+ mouse metadata with visits and nose-pokes, compute behavior metrics, and create
5
+ publication-oriented plots for place-learning experiments.
6
+ """
7
+
8
+ from .loader import CohortData, load_cohort_data
9
+
10
+ __all__ = ["CohortData", "load_cohort_data"]
11
+ __version__ = "0.0.0"
@@ -0,0 +1,477 @@
1
+ """Data loading utilities for IntelliCage place-learning experiments.
2
+
3
+ This module turns the raw IntelliCage exports into harmonized pandas
4
+ DataFrames. It reads:
5
+
6
+ 1. `Mice.txt` for mouse metadata and assigned corners.
7
+ 2. `Visits.txt` for visit-level behavior summaries.
8
+ 3. `Nosepokes.txt` for visit-linked nose-poke events.
9
+
10
+ The returned visit table contains both the old MATLAB-compatible metric
11
+ (`PlaceError == 0`) and the stricter analysis-oriented metric that requires a
12
+ correct place visit with at least one associated nose-poke and at least one
13
+ lick.
14
+
15
+ author: Fabrizio Musacchio
16
+ date: May 2026
17
+ """
18
+ # %% IMPORTS
19
+ from __future__ import annotations
20
+
21
+ from dataclasses import dataclass
22
+ from pathlib import Path
23
+
24
+ import numpy as np
25
+ import pandas as pd
26
+ # %% CONSTANTS
27
+ PHASE_NAMES: tuple[str, ...] = ("Phase1", "Phase2", "Phase3", "Phase4")
28
+ DEFAULT_PHASE_NAME_MAP: dict[str, int] = {
29
+ "Phase1": 1,
30
+ "Phase2": 2,
31
+ "Phase3": 3,
32
+ "Phase4": 4}
33
+ DEFAULT_GROUP_NAMES: tuple[str, ...] = tuple(f"Group {index}" for index in range(1, 11))
34
+ PHASE_DISPLAY_LABELS: dict[int, str] = {
35
+ 1: "Phase1",
36
+ 2: "Phase2",
37
+ 3: "Phase3",
38
+ 4: "Phase4"}
39
+ # %% DATA CLASSES
40
+ @dataclass(frozen=True)
41
+ class CohortData:
42
+ """Container that keeps the loaded cohort tables together."""
43
+
44
+ visits: pd.DataFrame
45
+ metadata: pd.DataFrame
46
+ nosepokes: pd.DataFrame
47
+ phase_manifest: pd.DataFrame
48
+
49
+ # %% HELPER FUNCTIONS
50
+ def _ordered_raw_group_names(groups: pd.Series) -> list[str]:
51
+ """Return raw group labels in first-seen order."""
52
+
53
+ ordered: list[str] = []
54
+ for group_name in groups.dropna().astype(str):
55
+ if group_name not in ordered:
56
+ ordered.append(group_name)
57
+ return ordered
58
+
59
+ def _complete_group_names(group_count: int, group_names: list[str] | tuple[str, ...] | None) -> list[str]:
60
+ """Fill a user-supplied group-name list with generic defaults."""
61
+
62
+ selected: list[str] = []
63
+ for group_name in group_names or []:
64
+ display_name = str(group_name)
65
+ if display_name and display_name not in selected:
66
+ selected.append(display_name)
67
+ next_index = len(selected) + 1
68
+ for default_name in DEFAULT_GROUP_NAMES[len(selected):]:
69
+ if len(selected) >= group_count:
70
+ break
71
+ if default_name not in selected:
72
+ selected.append(default_name)
73
+ next_index += 1
74
+ while len(selected) < group_count:
75
+ candidate = f"Group {next_index}"
76
+ if candidate not in selected:
77
+ selected.append(candidate)
78
+ next_index += 1
79
+ return selected[:group_count]
80
+
81
+ def _resolve_group_name_mapping(
82
+ groups: pd.Series,
83
+ group_names: list[str] | tuple[str, ...] | None) -> tuple[dict[str, str], list[str]]:
84
+ """Map raw dataset group labels to public display labels."""
85
+
86
+ raw_groups = _ordered_raw_group_names(groups)
87
+ display_names = _complete_group_names(len(raw_groups), group_names)
88
+ return dict(zip(raw_groups, display_names)), display_names
89
+
90
+ def read_mice_metadata(mice_path: Path, run_group: str) -> pd.DataFrame:
91
+ """Read one `Mice.txt` file and normalize its metadata columns."""
92
+
93
+ metadata = pd.read_csv(mice_path, sep="\t")
94
+ if "Corner Phase 3" not in metadata.columns and "Corner Phase 1" in metadata.columns:
95
+ metadata = metadata.rename(columns={"Corner Phase 1": "Corner Phase 3"})
96
+ if "Corner Phase 4" not in metadata.columns and "Corner Phase 2" in metadata.columns:
97
+ metadata = metadata.rename(columns={"Corner Phase 2": "Corner Phase 4"})
98
+ metadata = metadata.rename(
99
+ columns={
100
+ "VIRUS": "Group",
101
+ "Corner Phase 3": "CornerPhase3",
102
+ "Corner Phase 4": "CornerPhase4",
103
+ })
104
+ metadata["RunGroup"] = run_group
105
+ metadata["RFID"] = pd.to_numeric(metadata["RFID"], errors="raise").astype("Int64")
106
+ metadata["ET"] = metadata["ET"].astype("string").str.strip()
107
+ metadata["ETLabel"] = np.where(
108
+ metadata["ET"].str.match(r"^(ET|Lo)", case=False, na=False),
109
+ metadata["ET"],
110
+ "ET" + metadata["ET"])
111
+ metadata["DOB"] = pd.to_datetime(metadata["DOB"], format="%d.%m.%y", dayfirst=True, errors="coerce")
112
+ if "CornerPhase3" not in metadata.columns:
113
+ metadata["CornerPhase3"] = pd.Series(pd.NA, index=metadata.index, dtype="Int64")
114
+ else:
115
+ metadata["CornerPhase3"] = pd.to_numeric(metadata["CornerPhase3"], errors="coerce").astype("Int64")
116
+ if "CornerPhase4" not in metadata.columns:
117
+ metadata["CornerPhase4"] = pd.Series(pd.NA, index=metadata.index, dtype="Int64")
118
+ else:
119
+ metadata["CornerPhase4"] = pd.to_numeric(metadata["CornerPhase4"], errors="coerce").astype("Int64")
120
+ metadata["SEX"] = metadata["SEX"].astype("string")
121
+ metadata["Group"] = metadata["Group"].astype("string")
122
+ return metadata
123
+
124
+ def read_visits_file(visits_path: Path, run_group: str, phase_name: str, phase_number: int) -> pd.DataFrame:
125
+ """Read one IntelliCage `Visits.txt` file into a typed DataFrame."""
126
+
127
+ visits = pd.read_csv(visits_path, sep="\t")
128
+ visits["RunGroup"] = run_group
129
+ visits["Phase"] = phase_name
130
+ visits["PhaseNumber"] = int(phase_number)
131
+ visits["Start"] = pd.to_datetime(visits["Start"], errors="raise")
132
+ visits["End"] = pd.to_datetime(visits["End"], errors="raise")
133
+ visits["AnimalTag"] = pd.to_numeric(visits["AnimalTag"], errors="raise").astype("Int64")
134
+ visits["VisitID"] = pd.to_numeric(visits["VisitID"], errors="raise").astype("Int64")
135
+ visits["VisitDurationSeconds"] = (visits["End"] - visits["Start"]).dt.total_seconds()
136
+ visits["visit_has_lick"] = visits["LickNumber"].fillna(0).gt(0) | visits["LickDuration"].fillna(0).gt(0)
137
+ return visits
138
+
139
+ def read_nosepokes_file(
140
+ nosepokes_path: Path,
141
+ run_group: str,
142
+ phase_name: str,
143
+ phase_number: int,
144
+ ) -> pd.DataFrame:
145
+ """Read one IntelliCage `Nosepokes.txt` file into a typed DataFrame."""
146
+
147
+ nosepokes = pd.read_csv(nosepokes_path, sep="\t")
148
+ nosepokes["RunGroup"] = run_group
149
+ nosepokes["Phase"] = phase_name
150
+ nosepokes["PhaseNumber"] = int(phase_number)
151
+ nosepokes["Start"] = pd.to_datetime(nosepokes["Start"], errors="raise")
152
+ nosepokes["End"] = pd.to_datetime(nosepokes["End"], errors="raise")
153
+ nosepokes["VisitID"] = pd.to_numeric(nosepokes["VisitID"], errors="raise").astype("Int64")
154
+ if "LickStartTime" in nosepokes.columns:
155
+ nosepokes["LickStartTime"] = pd.to_datetime(nosepokes["LickStartTime"], errors="coerce")
156
+ return nosepokes
157
+
158
+ def summarize_nosepokes_by_visit(nosepokes: pd.DataFrame) -> pd.DataFrame:
159
+ """Aggregate raw nose-poke rows to one row per visit."""
160
+
161
+ summary = (
162
+ nosepokes.groupby(["RunGroup", "Phase", "PhaseNumber", "VisitID"], observed=True)
163
+ .agg(
164
+ nosepoke_event_count=("VisitID", "size"),
165
+ nosepoke_side_count=("Side", "nunique"),
166
+ nosepoke_lick_count=("LickNumber", "sum"),
167
+ nosepoke_lick_duration=("LickDuration", "sum"),
168
+ nosepoke_condition_error_count=("ConditionError", "sum"),
169
+ ).reset_index())
170
+ summary["has_nosepoke"] = summary["nosepoke_event_count"].gt(0)
171
+ summary["has_nosepoke_lick"] = summary["nosepoke_lick_count"].gt(0) | summary["nosepoke_lick_duration"].gt(0)
172
+ return summary
173
+
174
+ def _build_phase_manifest(visits: pd.DataFrame) -> pd.DataFrame:
175
+ """Create a manifest with the temporal extent of every phase."""
176
+
177
+ return (
178
+ visits.groupby(["RunGroup", "Phase", "PhaseNumber"], observed=True)
179
+ .agg(
180
+ PhaseStart=("Start", "min"),
181
+ PhaseEnd=("End", "max"),
182
+ VisitCount=("VisitID", "size"),
183
+ MouseCount=("AnimalTag", "nunique"))
184
+ .reset_index()
185
+ .sort_values(["RunGroup", "PhaseNumber"])
186
+ .reset_index(drop=True))
187
+
188
+ def _attach_time_reference_columns(visits: pd.DataFrame, phase_manifest: pd.DataFrame) -> pd.DataFrame:
189
+ """Add experiment-relative and phase-relative timing columns."""
190
+
191
+ experiment_starts = (
192
+ phase_manifest.loc[phase_manifest["PhaseNumber"].eq(1), ["RunGroup", "PhaseStart"]]
193
+ .rename(columns={"PhaseStart": "ExperimentStart"})
194
+ .copy())
195
+ phase_starts = phase_manifest.loc[:, ["RunGroup", "Phase", "PhaseNumber", "PhaseStart"]].copy()
196
+
197
+ enriched = visits.merge(experiment_starts, on="RunGroup", how="left", validate="many_to_one")
198
+ enriched = enriched.merge(
199
+ phase_starts,
200
+ on=["RunGroup", "Phase", "PhaseNumber"],
201
+ how="left",
202
+ validate="many_to_one")
203
+ enriched["experiment_elapsed_hours"] = (enriched["Start"] - enriched["ExperimentStart"]).dt.total_seconds() / 3600.0
204
+ enriched["phase_elapsed_hours"] = (enriched["Start"] - enriched["PhaseStart"]).dt.total_seconds() / 3600.0
205
+ enriched["experiment_day"] = np.floor(enriched["experiment_elapsed_hours"] / 24.0).astype(int)
206
+ enriched["phase_day"] = np.floor(enriched["phase_elapsed_hours"] / 24.0).astype(int) + 1
207
+ return enriched
208
+
209
+ def attach_analysis_time_columns(
210
+ visits: pd.DataFrame,
211
+ phase_manifest: pd.DataFrame,
212
+ *,
213
+ scheduled_phase_start_hours: dict[int, float],
214
+ mouse_day_start_hour: float,
215
+ experiment_day0_start_hour: float | None = None,
216
+ schedule_anchor_phase_number: int | None = None,
217
+ ) -> pd.DataFrame:
218
+ """Attach globally aligned analysis-time columns.
219
+
220
+ The raw IntelliCage exports store visits in phase-specific files and the
221
+ actual file boundaries can differ by a small amount from the intended
222
+ protocol timing. For cross-group comparisons we therefore
223
+ create a second time axis:
224
+
225
+ - experimental time starts at the mouse-day onset of day 0
226
+ - protocol phase windows are assigned from a global schedule in elapsed
227
+ hours rather than from the raw file boundary
228
+ - phase-relative elapsed time is then computed from this scheduled phase
229
+ start
230
+
231
+ Parameters
232
+ ----------
233
+ visits:
234
+ Visit-level table returned by :func:`load_cohort_data`.
235
+ phase_manifest:
236
+ Manifest with the observed temporal range of each raw phase file.
237
+ scheduled_phase_start_hours:
238
+ Mapping from phase number to global experiment-relative start hour.
239
+ An additional trailing marker, for example ``5=266``, can be provided
240
+ to define the exclusive end of phase 4.
241
+ mouse_day_start_hour:
242
+ Clock time that defines the beginning of the mouse day on day 0.
243
+ experiment_day0_start_hour:
244
+ Optional independent wall-clock hour that defines experiment elapsed
245
+ time zero on day 0. When omitted, the experiment timeline continues to
246
+ start at ``mouse_day_start_hour`` for backward compatibility. Set this
247
+ separately when day counting should start before the awake phase, for
248
+ example at midnight while the mouse day still begins at 07:00.
249
+ schedule_anchor_phase_number:
250
+ Optional raw phase number whose observed start should be aligned to the
251
+ configured scheduled start hour for every run group. This is useful
252
+ when early free-hab durations vary between runs, but all later phases
253
+ should be synchronized to the protocol transition point, for example
254
+ the observed start of NPA.
255
+ """
256
+
257
+ enriched = visits.copy()
258
+ phase1_starts = (
259
+ phase_manifest.loc[phase_manifest["PhaseNumber"].eq(1), ["RunGroup", "PhaseStart"]]
260
+ .rename(columns={"PhaseStart": "Phase1ObservedStart"})
261
+ .copy())
262
+ enriched = enriched.merge(phase1_starts, on="RunGroup", how="left", validate="many_to_one")
263
+ if enriched["Phase1ObservedStart"].isna().any():
264
+ raise ValueError("Could not determine the observed phase-1 start for every run group.")
265
+
266
+ analysis_origin_hour = (
267
+ float(mouse_day_start_hour)
268
+ if experiment_day0_start_hour is None
269
+ else float(experiment_day0_start_hour))
270
+ phase1_floor_day = enriched["Phase1ObservedStart"].dt.floor("D")
271
+ tentative_start = phase1_floor_day + pd.to_timedelta(analysis_origin_hour, unit="h")
272
+ starts_before_day_anchor = enriched["Phase1ObservedStart"] < tentative_start
273
+ enriched["AnalysisExperimentStart"] = tentative_start.where(
274
+ ~starts_before_day_anchor,
275
+ tentative_start - pd.to_timedelta(1, unit="D"))
276
+ enriched["analysis_experiment_elapsed_hours"] = (enriched["Start"] - enriched["AnalysisExperimentStart"]).dt.total_seconds() / 3600.0
277
+
278
+ sorted_phase_starts = sorted((int(key), float(value)) for key, value in scheduled_phase_start_hours.items())
279
+ if not sorted_phase_starts:
280
+ raise ValueError("`scheduled_phase_start_hours` must contain at least one phase start.")
281
+
282
+ if schedule_anchor_phase_number is not None:
283
+ anchor_phase_number = int(schedule_anchor_phase_number)
284
+ if anchor_phase_number not in scheduled_phase_start_hours:
285
+ raise ValueError("`schedule_anchor_phase_number` must be present in `scheduled_phase_start_hours`.")
286
+ anchor_rows = phase_manifest.loc[
287
+ phase_manifest["PhaseNumber"].eq(anchor_phase_number),
288
+ ["RunGroup", "PhaseStart"]].rename(columns={"PhaseStart": "AnchorPhaseObservedStart"})
289
+ if anchor_rows.empty:
290
+ raise ValueError("Could not find the requested anchor phase in the phase manifest.")
291
+ enriched = enriched.merge(anchor_rows, on="RunGroup", how="left", validate="many_to_one")
292
+ if enriched["AnchorPhaseObservedStart"].isna().any():
293
+ raise ValueError("Could not determine the observed anchor-phase start for every run group.")
294
+ enriched["anchor_phase_observed_hours"] = (enriched["AnchorPhaseObservedStart"] - enriched["AnalysisExperimentStart"]).dt.total_seconds() / 3600.0
295
+ enriched["schedule_alignment_offset_hours"] = (float(scheduled_phase_start_hours[anchor_phase_number]) - enriched["anchor_phase_observed_hours"])
296
+ enriched["analysis_experiment_elapsed_hours"] = (enriched["analysis_experiment_elapsed_hours"] + enriched["schedule_alignment_offset_hours"])
297
+ else:
298
+ enriched["schedule_alignment_offset_hours"] = 0.0
299
+
300
+ enriched["analysis_experiment_day"] = np.floor(enriched["analysis_experiment_elapsed_hours"] / 24.0).astype(int)
301
+
302
+ phase_rows: list[dict[str, float | int | str]] = []
303
+ for index, (phase_number, start_hour) in enumerate(sorted_phase_starts):
304
+ next_start = sorted_phase_starts[index + 1][1] if index + 1 < len(sorted_phase_starts) else np.inf
305
+ phase_rows.append(
306
+ {
307
+ "AnalysisPhaseNumber": phase_number,
308
+ "analysis_phase_start_hours": start_hour,
309
+ "analysis_phase_end_hours": next_start,
310
+ "AnalysisPhase": PHASE_DISPLAY_LABELS.get(phase_number, f"Phase{phase_number}"),
311
+ })
312
+
313
+ phase_table = pd.DataFrame(phase_rows)
314
+ valid_phase_table = phase_table.loc[phase_table["AnalysisPhaseNumber"].between(1, 4)].copy()
315
+ bins = [-np.inf, *valid_phase_table["analysis_phase_end_hours"].tolist()]
316
+ labels = valid_phase_table["AnalysisPhaseNumber"].tolist()
317
+ enriched["AnalysisPhaseNumber"] = pd.cut(
318
+ enriched["analysis_experiment_elapsed_hours"],
319
+ bins=bins,
320
+ labels=labels,
321
+ right=False).astype("Float64")
322
+ enriched["AnalysisPhaseNumber"] = enriched["AnalysisPhaseNumber"].astype("Int64")
323
+ phase_start_lookup = valid_phase_table.set_index("AnalysisPhaseNumber")["analysis_phase_start_hours"]
324
+ phase_name_lookup = valid_phase_table.set_index("AnalysisPhaseNumber")["AnalysisPhase"]
325
+ enriched["analysis_phase_start_hours"] = enriched["AnalysisPhaseNumber"].map(phase_start_lookup)
326
+ enriched["analysis_phase_elapsed_hours"] = (enriched["analysis_experiment_elapsed_hours"] - enriched["analysis_phase_start_hours"])
327
+ enriched["analysis_phase_day"] = np.floor(enriched["analysis_phase_elapsed_hours"] / 24.0).astype("Int64") + 1
328
+ enriched["AnalysisPhase"] = enriched["AnalysisPhaseNumber"].map(phase_name_lookup).astype("string")
329
+ enriched["AnalysisAssignedCorner"] = pd.Series(pd.NA, index=enriched.index, dtype="Int64")
330
+ enriched.loc[enriched["AnalysisPhaseNumber"].eq(3), "AnalysisAssignedCorner"] = enriched.loc[enriched["AnalysisPhaseNumber"].eq(3), "CornerPhase3"]
331
+ enriched.loc[enriched["AnalysisPhaseNumber"].eq(4), "AnalysisAssignedCorner"] = enriched.loc[enriched["AnalysisPhaseNumber"].eq(4), "CornerPhase4"]
332
+ enriched["correct_corner_visit"] = enriched["Corner"].eq(enriched["AnalysisAssignedCorner"])
333
+ enriched["correct_np_visit"] = enriched["correct_corner_visit"] & enriched["has_nosepoke"]
334
+ enriched["rewarded_correct_corner_visit"] = (enriched["correct_corner_visit"] & enriched["has_nosepoke"] & enriched["visit_has_lick"])
335
+ enriched["previous_correct_corner_visit"] = (enriched["AnalysisPhaseNumber"].eq(4) & enriched["Corner"].eq(enriched["CornerPhase3"]))
336
+ enriched["neutral_incorrect_corner_visit"] = (
337
+ enriched["AnalysisPhaseNumber"].eq(4)
338
+ & enriched["Corner"].notna()
339
+ & ~enriched["Corner"].eq(enriched["CornerPhase4"])
340
+ & ~enriched["Corner"].eq(enriched["CornerPhase3"]))
341
+ return enriched
342
+
343
+ def load_cohort_data(
344
+ dataset_root: Path | str,
345
+ *,
346
+ phase_name_map: dict[str, int] | None = None,
347
+ optional_phase_names: set[str] | list[str] | tuple[str, ...] | None = None,
348
+ drop_unmatched_visits: bool = False,
349
+ group_names: list[str] | tuple[str, ...] | None = None) -> CohortData:
350
+ """Load and harmonize one IntelliCage cohort directory.
351
+
352
+ Parameters
353
+ ----------
354
+ dataset_root:
355
+ Root directory of one IntelliCage cohort.
356
+ phase_name_map:
357
+ Mapping from subfolder names such as ``Phase1`` or ``SP2`` to the
358
+ raw phase numbers that should be assigned during loading.
359
+ optional_phase_names:
360
+ Folder names that may be absent for some run groups without causing
361
+ the loader to fail.
362
+ group_names:
363
+ Optional display names for the raw groups found in `Mice.txt`, in
364
+ first-seen order. Missing entries are filled with generic `Group N`
365
+ names so partially specified group-name lists remain valid.
366
+ """
367
+
368
+ dataset_root = Path(dataset_root)
369
+ selected_phase_map = DEFAULT_PHASE_NAME_MAP.copy()
370
+ if phase_name_map:
371
+ selected_phase_map = {str(key): int(value) for key, value in phase_name_map.items()}
372
+ optional_phase_name_set = {str(name) for name in (optional_phase_names or set())}
373
+ run_group_dirs = sorted(
374
+ path for path in dataset_root.iterdir() if path.is_dir() and path.name.startswith("Gruppe"))
375
+ if not run_group_dirs:
376
+ raise FileNotFoundError(f"No run-group directories found below {dataset_root}")
377
+
378
+ metadata_frames: list[pd.DataFrame] = []
379
+ visit_frames: list[pd.DataFrame] = []
380
+ nosepoke_frames: list[pd.DataFrame] = []
381
+
382
+ for run_group_dir in run_group_dirs:
383
+ mice_path = run_group_dir / "Mice.txt"
384
+ if not mice_path.exists():
385
+ raise FileNotFoundError(f"Missing metadata file: {mice_path}")
386
+ metadata_frames.append(read_mice_metadata(mice_path, run_group_dir.name))
387
+
388
+ for phase_name, phase_number in selected_phase_map.items():
389
+ visits_path = run_group_dir / phase_name / "IntelliCage" / "Visits.txt"
390
+ nosepokes_path = run_group_dir / phase_name / "IntelliCage" / "Nosepokes.txt"
391
+ if not visits_path.exists():
392
+ if phase_name in optional_phase_name_set:
393
+ continue
394
+ raise FileNotFoundError(f"Missing visit file: {visits_path}")
395
+ if not nosepokes_path.exists():
396
+ if phase_name in optional_phase_name_set:
397
+ continue
398
+ raise FileNotFoundError(f"Missing nose-poke file: {nosepokes_path}")
399
+ visit_frames.append(read_visits_file(visits_path, run_group_dir.name, phase_name, phase_number))
400
+ nosepoke_frames.append(read_nosepokes_file(nosepokes_path, run_group_dir.name, phase_name, phase_number))
401
+
402
+ metadata = pd.concat(metadata_frames, ignore_index=True)
403
+ visits = pd.concat(visit_frames, ignore_index=True)
404
+ nosepokes = pd.concat(nosepoke_frames, ignore_index=True)
405
+
406
+ nosepoke_summary = summarize_nosepokes_by_visit(nosepokes)
407
+ visits = visits.merge(
408
+ metadata,
409
+ left_on=["RunGroup", "AnimalTag"],
410
+ right_on=["RunGroup", "RFID"],
411
+ how="left",
412
+ validate="many_to_one")
413
+ if visits["ET"].isna().any():
414
+ missing_rows = visits.loc[visits["ET"].isna(), ["RunGroup", "AnimalTag"]].drop_duplicates()
415
+ if not drop_unmatched_visits:
416
+ raise ValueError(
417
+ "Some visits could not be matched to `Mice.txt` metadata. "
418
+ f"Missing pairs: {missing_rows.to_dict(orient='records')}")
419
+ visits = visits.loc[visits["ET"].notna()].copy()
420
+ valid_visit_keys = visits.loc[:, ["RunGroup", "Phase", "PhaseNumber", "VisitID"]].drop_duplicates()
421
+ nosepokes = nosepokes.merge(
422
+ valid_visit_keys,
423
+ on=["RunGroup", "Phase", "PhaseNumber", "VisitID"],
424
+ how="inner",
425
+ validate="many_to_one")
426
+
427
+ visits = visits.merge(
428
+ nosepoke_summary,
429
+ on=["RunGroup", "Phase", "PhaseNumber", "VisitID"],
430
+ how="left",
431
+ validate="many_to_one")
432
+ nosepoke_columns = [
433
+ "nosepoke_event_count",
434
+ "nosepoke_side_count",
435
+ "nosepoke_lick_count",
436
+ "nosepoke_lick_duration",
437
+ "nosepoke_condition_error_count"]
438
+ for column in nosepoke_columns:
439
+ visits[column] = visits[column].fillna(0)
440
+ for column in ["has_nosepoke", "has_nosepoke_lick"]:
441
+ visits[column] = visits[column].fillna(False)
442
+
443
+ visits["AssignedCorner"] = pd.Series(pd.NA, index=visits.index, dtype="Int64")
444
+ visits.loc[visits["PhaseNumber"].eq(3), "AssignedCorner"] = visits.loc[visits["PhaseNumber"].eq(3), "CornerPhase3"]
445
+ visits.loc[visits["PhaseNumber"].eq(4), "AssignedCorner"] = visits.loc[visits["PhaseNumber"].eq(4), "CornerPhase4"]
446
+ visits["assigned_corner_visit"] = visits["Corner"].eq(visits["AssignedCorner"])
447
+ visits["correct_place_visit"] = visits["PlaceError"].eq(0)
448
+ visits["correct_corner_visit"] = visits["assigned_corner_visit"]
449
+ visits["correct_np_visit"] = visits["assigned_corner_visit"] & visits["has_nosepoke"]
450
+ visits["rewarded_place_visit"] = (visits["correct_place_visit"] & visits["has_nosepoke"] & visits["visit_has_lick"])
451
+ visits["rewarded_correct_corner_visit"] = (visits["assigned_corner_visit"] & visits["has_nosepoke"] & visits["visit_has_lick"])
452
+ visits["previous_correct_corner_visit"] = visits["PhaseNumber"].eq(4) & visits["Corner"].eq(visits["CornerPhase3"])
453
+ visits["neutral_incorrect_corner_visit"] = (
454
+ visits["PhaseNumber"].eq(4)
455
+ & visits["Corner"].notna()
456
+ & ~visits["Corner"].eq(visits["CornerPhase4"])
457
+ & ~visits["Corner"].eq(visits["CornerPhase3"]))
458
+ visits["phase2_drinking_visit"] = (visits["PhaseNumber"].eq(2) & visits["has_nosepoke"] & visits["visit_has_lick"])
459
+
460
+ group_name_mapping, group_categories = _resolve_group_name_mapping(metadata["Group"], group_names)
461
+ visits["Group"] = visits["Group"].astype(str).map(group_name_mapping).fillna(visits["Group"].astype(str))
462
+ metadata["Group"] = metadata["Group"].astype(str).map(group_name_mapping).fillna(metadata["Group"].astype(str))
463
+ visits["Group"] = pd.Categorical(visits["Group"], categories=group_categories, ordered=True)
464
+ metadata["Group"] = pd.Categorical(metadata["Group"], categories=group_categories, ordered=True)
465
+
466
+ phase_manifest = _build_phase_manifest(visits)
467
+ visits = _attach_time_reference_columns(visits, phase_manifest)
468
+ visits = visits.sort_values(["RunGroup", "PhaseNumber", "Start", "VisitID"]).reset_index(drop=True)
469
+ metadata = metadata.sort_values(["Group", "ET"]).reset_index(drop=True)
470
+ nosepokes = nosepokes.sort_values(["RunGroup", "PhaseNumber", "Start", "VisitID"]).reset_index(drop=True)
471
+
472
+ return CohortData(
473
+ visits=visits,
474
+ metadata=metadata,
475
+ nosepokes=nosepokes,
476
+ phase_manifest=phase_manifest)
477
+ # %% END