diff --git a/docs/MIGRATION.md b/docs/MIGRATION.md index bd61491..8fdb0f2 100644 --- a/docs/MIGRATION.md +++ b/docs/MIGRATION.md @@ -151,8 +151,8 @@ Within a NIDAQ-timed session, mapped flashes and port edges must agree with the Bpod trial alignment within 100 ms. Trials with missing or misaligned hardware events are excluded rather than replaced with Bpod timestamps. -`PsychophysicalKernel.timing_source` records `nidaq`, `bpod`, or `mixed` for the -pooled result. Detailed session provenance remains derivable from +`PsychophysicalKernel.timing_source` is part of the primary key and records +`nidq` or `bpod` for the fit. Detailed session provenance remains derivable from `BehaviorAnalysisSet.TrialSet` and the ephys `EventMapping`; no duplicate per-session provenance table is needed. diff --git a/labdata_plugin/analysisschema.py b/labdata_plugin/analysisschema.py index 301c88b..61841cc 100644 --- a/labdata_plugin/analysisschema.py +++ b/labdata_plugin/analysisschema.py @@ -193,17 +193,17 @@ class PsychophysicalKernelFitConfig(dj.Lookup): @rojasbowe_schema class PsychophysicalKernel(dj.Computed): - """One pooled kernel per analysis set, subject, condition, and config.""" + """One pooled kernel per analysis set, subject, condition, config, and clock.""" definition = """ -> BehaviorAnalysisSet -> Subject trialset_description : varchar(54) -> PsychophysicalKernelFitConfig + timing_source : enum('nidq', 'bpod') --- fit_status : enum('fit', 'skipped') fit_message = NULL : varchar(256) - timing_source = 'bpod' : enum('nidaq', 'bpod', 'mixed') n_trials_fit : int n_bins_fit = NULL : int n_observed_per_bin = NULL : longblob @@ -227,7 +227,29 @@ def key_source(self): .aggr(BehaviorAnalysisSet.TrialSet(), n_trialsets="count(*)") .proj() ) - return subject_conditions * PsychophysicalKernelFitConfig() + base = subject_conditions * PsychophysicalKernelFitConfig() + key_fields = ( + "analysis_set_id", + "subject_name", + "trialset_description", + "kernel_fit_config_id", + ) + nidq_keys = [] + for row in base.fetch(as_dict=True): + trialset_keys = _selected_trialset_keys(row) + from behavior_analyses.kernel_timing import available_timing_sources + + if "nidq" in available_timing_sources(trialset_keys): + nidq_keys.append({field: row[field] for field in key_fields}) + + key_relation = dj.U(*key_fields, "timing_source") + bpod = key_relation & base.proj(*key_fields, timing_source="'bpod'") + if not nidq_keys: + return bpod + nidq = key_relation & (base & nidq_keys).proj( + *key_fields, timing_source="'nidq'" + ) + return bpod + nidq def make(self, key): config = (PsychophysicalKernelFitConfig() & key).fetch1() @@ -274,6 +296,7 @@ def _kernel_payload(key, config, trialset_keys): trialset_keys, key["trialset_description"], observation_window=str(config["observation_window"]), + timing_source=str(key["timing_source"]), ) residual, choices, n_observed, bin_centers, expected_counts = ( build_residual_rate_matrix( @@ -297,7 +320,6 @@ def _kernel_payload(key, config, trialset_keys): expected_counts=expected_counts, ) base = { - "timing_source": inputs["timing_source"], "n_trials_fit": int(result["n_trials_fit"]), "n_bins_fit": int(result["n_bins_fit"]), "n_observed_per_bin": n_observed, diff --git a/src/behavior_analyses/kernel_timing.py b/src/behavior_analyses/kernel_timing.py index a6a0aa4..f7220d4 100644 --- a/src/behavior_analyses/kernel_timing.py +++ b/src/behavior_analyses/kernel_timing.py @@ -5,6 +5,7 @@ import numpy as np TRIALSET_DATASET_KEY_FIELDS = ("subject_name", "session_name", "dataset_name") +TIMING_SOURCES = ("nidq", "bpod") REQUIRED_NIDAQ_EVENTS = ( "visual_stim", "trial_start", @@ -15,18 +16,45 @@ MAX_NIDAQ_ALIGNMENT_ERROR_S = 0.1 +def available_timing_sources(trialset_keys: list[dict[str, Any]]) -> list[str]: + """Return timing sources the selected trial sets can supply.""" + if not trialset_keys: + return [] + session_keys = {(key["subject_name"], key["session_name"]) for key in trialset_keys} + sources = ["bpod"] + if all( + has_nidq_timing( + _fetch_event_mapping_rows( + {"subject_name": subject, "session_name": session} + ) + ) + for subject, session in session_keys + ): + sources.insert(0, "nidq") + return sources + + def fetch_pooled_kernel_inputs( trialset_keys: list[dict[str, Any]], trialset_description: str, *, observation_window: str, + timing_source: str, ) -> dict[str, Any]: - """Fetch pooled trial inputs, preferring validated NIDAQ timing per session.""" + """Fetch pooled trial inputs from one requested timing source.""" if observation_window not in {"center_exit", "response"}: raise ValueError( "observation_window must be 'center_exit' or 'response', " f"got {observation_window!r}" ) + if timing_source not in TIMING_SOURCES: + raise ValueError( + f"timing_source must be one of {TIMING_SOURCES}, got {timing_source!r}" + ) + if timing_source not in available_timing_sources(trialset_keys): + raise ValueError( + f"Selected trial sets cannot supply timing_source={timing_source!r}" + ) session_inputs = [] seen_datasets = set() @@ -44,28 +72,32 @@ def fetch_pooled_kernel_inputs( field: dataset_key[field] for field in ("subject_name", "session_name") } mapping_rows = _fetch_event_mapping_rows(session_key) - if has_nidaq_visual_timing(mapping_rows): + if timing_source == "nidq": + if not has_nidq_timing(mapping_rows): + raise ValueError( + f"Session {session_key['subject_name']} " + f"{session_key['session_name']} cannot supply nidq timing: " + "incomplete EventMapping" + ) event_rows = _fetch_mapped_digital_event_rows(session_key, mapping_rows) - aligned_events = resolve_nidaq_event_arrays( + aligned_events = resolve_nidq_event_arrays( event_rows, mapping_rows, session_key["subject_name"], session_key["session_name"], ) - inputs = extract_nidaq_kernel_inputs( + inputs = extract_nidq_kernel_inputs( aligned_events, trial_rows, trialset_description, observation_window=observation_window, ) - inputs["timing_source"] = "nidaq" else: inputs = extract_bpod_kernel_inputs( trial_rows, trialset_description, observation_window=observation_window, ) - inputs["timing_source"] = "bpod" session_inputs.append(inputs) if not session_inputs: @@ -109,7 +141,7 @@ def extract_bpod_kernel_inputs( return result -def extract_nidaq_kernel_inputs( +def extract_nidq_kernel_inputs( aligned_events: dict[str, np.ndarray], trial_rows: list[dict[str, Any]], trialset_description: str, @@ -134,12 +166,12 @@ def extract_nidaq_kernel_inputs( "Insufficient finite Bpod/NIDAQ sync points to interpolate trial timing" ) bpod_sync = np.asarray([row["t_sync"] for row in sync_rows], dtype=float) - nidaq_sync = np.asarray( + nidq_sync = np.asarray( [trial_starts[int(row["trial_num"])] for row in sync_rows], dtype=float ) order = np.argsort(bpod_sync) bpod_sync = bpod_sync[order] - nidaq_sync = nidaq_sync[order] + nidq_sync = nidq_sync[order] stims = np.asarray(aligned_events["visual_stim"], dtype=float) center_exits = np.asarray(aligned_events["center_port_exit"], dtype=float) @@ -174,11 +206,11 @@ def extract_nidaq_kernel_inputs( np.interp( float(row["t_sync"]) + float(bpod_stims[0]), bpod_sync, - nidaq_sync, + nidq_sync, ) ) interpolated_exit = float( - np.interp(float(row["t_react"]), bpod_sync, nidaq_sync) + np.interp(float(row["t_react"]), bpod_sync, nidq_sync) ) trial_exits = center_exits[ (center_exits > trial_start) & (center_exits < trial_end) @@ -193,7 +225,7 @@ def extract_nidaq_kernel_inputs( if response_time is None or not np.isfinite(response_time): continue interpolated_response = float( - np.interp(float(response_time), bpod_sync, nidaq_sync) + np.interp(float(response_time), bpod_sync, nidq_sync) ) response_entries = ( right_entries if int(row["response"]) == 1 else left_entries @@ -222,12 +254,15 @@ def extract_nidaq_kernel_inputs( return result -def has_nidaq_visual_timing(mapping_rows: list[dict[str, Any]]) -> bool: - """Return whether a session declares a mapped NIDAQ/OneBox visual stream.""" - return any(row.get("event_name") == "visual_stim" for row in mapping_rows) +def has_nidq_timing(mapping_rows: list[dict[str, Any]]) -> bool: + """Return whether a session has one mapping for every required NIDAQ event.""" + mapped_names = [row.get("event_name") for row in mapping_rows] + return len(mapped_names) == len(set(mapped_names)) and all( + name in mapped_names for name in REQUIRED_NIDAQ_EVENTS + ) -def resolve_nidaq_event_arrays( +def resolve_nidq_event_arrays( event_rows: list[dict[str, Any]], mapping_rows: list[dict[str, Any]], subject: str, @@ -293,9 +328,7 @@ def resolve_nidaq_event_arrays( def combine_kernel_inputs(session_inputs: list[dict[str, Any]]) -> dict[str, Any]: - """Combine per-session inputs and summarize their timing provenance.""" - sources = {inputs["timing_source"] for inputs in session_inputs} - timing_source = next(iter(sources)) if len(sources) == 1 else "mixed" + """Combine per-session inputs from one timing source.""" return { "stim_times_per_trial": [ stims @@ -310,7 +343,6 @@ def combine_kernel_inputs(session_inputs: list[dict[str, Any]]) -> dict[str, Any ), "response_values": _concatenate(session_inputs, "response_values", dtype=int), "trial_rate_hz": _concatenate(session_inputs, "trial_rate_hz", dtype=float), - "timing_source": timing_source, } diff --git a/tests/test_kernel_timing.py b/tests/test_kernel_timing.py index 095b6f1..b753e7f 100644 --- a/tests/test_kernel_timing.py +++ b/tests/test_kernel_timing.py @@ -3,6 +3,7 @@ from pathlib import Path import sys import unittest +from unittest.mock import patch import numpy as np @@ -35,6 +36,14 @@ def setUp(self): "t_response": None, }, ] + self.trialset_keys = [ + { + "subject_name": "GRB006", + "session_name": "20240821_121447", + "dataset_name": "chipmunk", + "trialset_description": "visual", + } + ] def test_bpod_window_uses_reaction_or_response_time(self): from behavior_analyses.kernel_timing import extract_bpod_kernel_inputs @@ -53,11 +62,11 @@ def test_bpod_window_uses_reaction_or_response_time(self): np.testing.assert_allclose(center["observation_end_times"], [0.5]) np.testing.assert_allclose(response["observation_end_times"], [0.9]) - def test_nidaq_visual_mapping_requires_complete_valid_events(self): - from behavior_analyses.kernel_timing import resolve_nidaq_event_arrays + def test_nidq_visual_mapping_requires_complete_valid_events(self): + from behavior_analyses.kernel_timing import resolve_nidq_event_arrays with self.assertRaisesRegex(ValueError, "Missing EventMapping"): - resolve_nidaq_event_arrays( + resolve_nidq_event_arrays( [], [ { @@ -71,23 +80,23 @@ def test_nidaq_visual_mapping_requires_complete_valid_events(self): "session", ) - def test_nidaq_window_uses_mapped_flash_and_port_times(self): + def test_nidq_window_uses_mapped_flash_and_port_times(self): from behavior_analyses.kernel_timing import ( - extract_nidaq_kernel_inputs, - resolve_nidaq_event_arrays, + extract_nidq_kernel_inputs, + resolve_nidq_event_arrays, ) mapping_rows, event_rows = self._mapped_event_fixture() - aligned = resolve_nidaq_event_arrays( + aligned = resolve_nidq_event_arrays( event_rows, mapping_rows, "GRB006", "session" ) - center = extract_nidaq_kernel_inputs( + center = extract_nidq_kernel_inputs( aligned, self.trial_rows, "visual", observation_window="center_exit", ) - response = extract_nidaq_kernel_inputs( + response = extract_nidq_kernel_inputs( aligned, self.trial_rows, "visual", @@ -101,15 +110,15 @@ def test_nidaq_window_uses_mapped_flash_and_port_times(self): np.testing.assert_allclose(center["observation_end_times"], [0.5]) np.testing.assert_allclose(response["observation_end_times"], [0.9]) - def test_bpod_and_nidaq_match_for_aligned_trial(self): + def test_bpod_and_nidq_match_for_aligned_trial(self): from behavior_analyses.kernel_timing import ( extract_bpod_kernel_inputs, - extract_nidaq_kernel_inputs, - resolve_nidaq_event_arrays, + extract_nidq_kernel_inputs, + resolve_nidq_event_arrays, ) mapping_rows, event_rows = self._mapped_event_fixture() - aligned = resolve_nidaq_event_arrays( + aligned = resolve_nidq_event_arrays( event_rows, mapping_rows, "GRB006", "session" ) @@ -117,40 +126,129 @@ def test_bpod_and_nidaq_match_for_aligned_trial(self): bpod = extract_bpod_kernel_inputs( self.trial_rows, "visual", observation_window=window ) - nidaq = extract_nidaq_kernel_inputs( + nidq = extract_nidq_kernel_inputs( aligned, self.trial_rows, "visual", observation_window=window, ) - self.assertEqual(bpod["response_values"], nidaq["response_values"]) + self.assertEqual(bpod["response_values"], nidq["response_values"]) np.testing.assert_allclose( - bpod["observation_end_times"], nidaq["observation_end_times"] + bpod["observation_end_times"], nidq["observation_end_times"] ) np.testing.assert_allclose( - bpod["stim_times_per_trial"][0], nidaq["stim_times_per_trial"][0] + bpod["stim_times_per_trial"][0], nidq["stim_times_per_trial"][0] ) - def test_combined_provenance_is_mixed(self): - from behavior_analyses.kernel_timing import ( - combine_kernel_inputs, - extract_bpod_kernel_inputs, - ) + def test_available_timing_sources_include_nidq_when_mapped(self): + from behavior_analyses.kernel_timing import available_timing_sources - first = extract_bpod_kernel_inputs( - self.trial_rows, "visual", observation_window="center_exit" - ) - second = extract_bpod_kernel_inputs( - self.trial_rows, "visual", observation_window="center_exit" - ) - first["timing_source"] = "bpod" - second["timing_source"] = "nidaq" + mapping_rows, _ = self._mapped_event_fixture() + with patch( + "behavior_analyses.kernel_timing._fetch_event_mapping_rows", + return_value=mapping_rows, + ): + sources = available_timing_sources(self.trialset_keys) + + self.assertEqual(sources, ["nidq", "bpod"]) + + def test_available_timing_sources_are_bpod_only_without_mapping(self): + from behavior_analyses.kernel_timing import available_timing_sources + + with patch( + "behavior_analyses.kernel_timing._fetch_event_mapping_rows", + return_value=[], + ): + sources = available_timing_sources(self.trialset_keys) + + self.assertEqual(sources, ["bpod"]) + + def test_available_timing_sources_require_all_nidq_mappings(self): + from behavior_analyses.kernel_timing import available_timing_sources + + mapping_rows, _ = self._mapped_event_fixture() + incomplete = [row for row in mapping_rows if row["event_name"] != "right_port"] + with patch( + "behavior_analyses.kernel_timing._fetch_event_mapping_rows", + return_value=incomplete, + ): + sources = available_timing_sources(self.trialset_keys) + + self.assertEqual(sources, ["bpod"]) - combined = combine_kernel_inputs([first, second]) + def test_fetch_pooled_kernel_inputs_requires_requested_timing_source(self): + from behavior_analyses.kernel_timing import fetch_pooled_kernel_inputs - self.assertEqual(combined["timing_source"], "mixed") - self.assertEqual(len(combined["stim_times_per_trial"]), 2) + with patch( + "behavior_analyses.kernel_timing.available_timing_sources", + return_value=["bpod"], + ): + with self.assertRaisesRegex( + ValueError, "cannot supply timing_source='nidq'" + ): + fetch_pooled_kernel_inputs( + self.trialset_keys, + "visual", + observation_window="center_exit", + timing_source="nidq", + ) + + def test_same_session_yields_distinct_rows_per_timing_source(self): + from behavior_analyses.kernel_timing import fetch_pooled_kernel_inputs + + mapping_rows, event_rows = self._mapped_event_fixture() + + def fake_trial_rows(_dataset_key): + return self.trial_rows + + def fake_mapping(_session_key): + return mapping_rows + + def fake_events(_session_key, _mapping_rows): + return event_rows + + with ( + patch( + "behavior_analyses.kernel_timing._fetch_chipmunk_trial_rows", + side_effect=fake_trial_rows, + ), + patch( + "behavior_analyses.kernel_timing._fetch_event_mapping_rows", + side_effect=fake_mapping, + ), + patch( + "behavior_analyses.kernel_timing._fetch_mapped_digital_event_rows", + side_effect=fake_events, + ), + ): + bpod = fetch_pooled_kernel_inputs( + self.trialset_keys, + "visual", + observation_window="center_exit", + timing_source="bpod", + ) + nidq = fetch_pooled_kernel_inputs( + self.trialset_keys, + "visual", + observation_window="center_exit", + timing_source="nidq", + ) + + shared_key = { + "analysis_set_id": "test_set", + "subject_name": "GRB006", + "trialset_description": "visual", + "kernel_fit_config_id": 0, + } + bpod_key = {**shared_key, "timing_source": "bpod"} + nidq_key = {**shared_key, "timing_source": "nidq"} + self.assertNotEqual(bpod_key["timing_source"], nidq_key["timing_source"]) + self.assertEqual( + {field: bpod_key[field] for field in shared_key}, + {field: nidq_key[field] for field in shared_key}, + ) + self.assertEqual(bpod["response_values"], nidq["response_values"]) @staticmethod def _mapped_event_fixture(): diff --git a/tests/test_schema_imports.py b/tests/test_schema_imports.py index 339591e..d266955 100644 --- a/tests/test_schema_imports.py +++ b/tests/test_schema_imports.py @@ -5,7 +5,7 @@ import sys import types import unittest -from unittest.mock import patch +from unittest.mock import MagicMock, patch REPO_ROOT = Path(__file__).resolve().parents[1] @@ -67,6 +67,69 @@ def test_analysis_schema_imports_with_fake_labdata(self): module = importlib.import_module("labdata_plugin.analysisschema") module = importlib.reload(module) + source = MagicMock() + base = MagicMock() + bpod = MagicMock() + nidq = MagicMock() + union = MagicMock() + source.aggr.return_value.proj.return_value = base + base.__mul__.return_value = base + base.fetch.return_value = [ + { + "analysis_set_id": "set", + "subject_name": "GRB006", + "trialset_description": "visual", + "kernel_fit_config_id": 0, + } + ] + base.proj.return_value = bpod + base.__and__.return_value.proj.return_value = nidq + key_relation = MagicMock() + key_relation.__and__.side_effect = [bpod, nidq] + bpod.__add__.return_value = union + universal_set = MagicMock(side_effect=[source, key_relation]) + + with ( + patch.object(module.dj, "U", universal_set, create=True), + patch.object( + module, + "BehaviorAnalysisSet", + types.SimpleNamespace(TrialSet=MagicMock()), + ), + patch.object(module, "PsychophysicalKernelFitConfig", MagicMock()), + patch.object( + module, + "_selected_trialset_keys", + return_value=[ + { + "subject_name": "GRB006", + "session_name": "session", + "dataset_name": "chipmunk", + "trialset_description": "visual", + } + ], + ), + patch( + "behavior_analyses.kernel_timing.available_timing_sources", + return_value=["nidq", "bpod"], + ), + ): + descriptor = module.PsychophysicalKernel.key_source + key_source = descriptor.fget( + object.__new__(module.PsychophysicalKernel) + ) + + self.assertTrue( + all( + isinstance(arg, str) + for call in universal_set.call_args_list + for arg in call.args + ) + ) + self.assertIn("timing_source", universal_set.call_args_list[-1].args) + bpod.__add__.assert_called_once_with(nidq) + self.assertIs(key_source, union) + self.assertTrue(hasattr(module, "BehaviorAnalysisSet")) self.assertTrue(hasattr(module, "PsychometricFitConfig")) self.assertTrue(hasattr(module, "PsychophysicalKernelFitConfig")) @@ -91,6 +154,11 @@ def test_analysis_schema_imports_with_fake_labdata(self): [row[0] for row in module.PsychophysicalKernelFitConfig.contents], [0, 1], ) + self.assertIn( + "timing_source : enum('nidq', 'bpod')", + module.PsychophysicalKernel.definition.split("---")[0], + ) + self.assertNotIn("mixed", module.PsychophysicalKernel.definition) self.assertFalse(hasattr(module, "LearningSessionMetrics"))