diff --git a/dj_pipeline/vr4mice/analysis/barcodes.py b/dj_pipeline/vr4mice/analysis/barcodes.py index d979845c..bc9b250f 100644 --- a/dj_pipeline/vr4mice/analysis/barcodes.py +++ b/dj_pipeline/vr4mice/analysis/barcodes.py @@ -95,8 +95,8 @@ def decode_teensy_barcodes( raise ValueError( "teensy_time, ttl_read, and photodiode_time must have the same shape" ) - if times.size and np.any(np.diff(times) <= 0): - raise ValueError("teensy_time must be strictly increasing") + if times.size and np.any(np.diff(times) < 0): + raise ValueError("teensy_time must be non-decreasing") if not np.isfinite(continuous_times).all(): raise ValueError("photodiode_time must contain only finite timestamps") diff --git a/dj_pipeline/vr4mice/analysis/dlc_helpers.py b/dj_pipeline/vr4mice/analysis/dlc_helpers.py index cc8ffea2..7f863f3a 100644 --- a/dj_pipeline/vr4mice/analysis/dlc_helpers.py +++ b/dj_pipeline/vr4mice/analysis/dlc_helpers.py @@ -367,8 +367,22 @@ def align_timestamps_to_step_time( if np.any(np.diff(step_time) < 0): raise ValueError("step_time must be sorted in ascending order") - indices = find_closest_indices(step_time.tolist(), timestamps.tolist()) - return step_time[np.asarray(indices, dtype=int)] + aligned = np.full(timestamps.shape, np.nan, dtype=np.float64) + in_range = (timestamps >= step_time[0]) & (timestamps <= step_time[-1]) + if not np.any(in_range): + return aligned + + ts = timestamps[in_range] + idx = np.searchsorted(step_time, ts, side="left") + left_idx = np.clip(idx - 1, 0, step_time.size - 1) + right_idx = np.clip(idx, 0, step_time.size - 1) + + left = step_time[left_idx] + right = step_time[right_idx] + choose_right = (idx > 0) & (idx < step_time.size) & ((ts - left) > (right - ts)) + nearest_idx = np.where(choose_right, right_idx, left_idx) + aligned[in_range] = step_time[nearest_idx] + return aligned def compute_circular_angular_velocity( diff --git a/dj_pipeline/vr4mice/analysis/np_sync.py b/dj_pipeline/vr4mice/analysis/np_sync.py index ee783a9d..41299696 100644 --- a/dj_pipeline/vr4mice/analysis/np_sync.py +++ b/dj_pipeline/vr4mice/analysis/np_sync.py @@ -8,24 +8,108 @@ import scipy.interpolate import scipy.stats -# Number of leading (chronologically earliest) VR barcode events excluded from the -# regression fit by default. DLC-live starts receiving data slightly after the -# Unity/game stream does, which is why downstream analysis already drops -# `trial == 1` as a DLC-live initialization trial (see vr4mice/analysis/analysis.py). -# Barcodes in that same early window can carry an unreliable onset_time_unity, -# biasing the fit if used as tie points, so we drop the earliest few by default. -DEFAULT_SKIP_FIRST_N_BARCODES = 10 +MIN_TIE_POINTS = 3 +OUTLIER_SIGMA = 5.0 +OUTLIER_FLOOR_MS = 30.0 +OUTLIER_MAX_FRACTION = 0.05 + + +def _boundary_repetitive_run_lengths(vr_times: np.ndarray) -> tuple[int, int]: + """Return leading/trailing consecutive equal-time run lengths. + + Direct `!=` is intentional here: onset_time_unity values at repeated + boundaries are copied from the same step_time array element and are + therefore expected to be bit-identical. + """ + if vr_times.size == 0: + return 0, 0 + + change_points = np.flatnonzero(vr_times[1:] != vr_times[:-1]) + 1 + run_starts = np.concatenate(([0], change_points)) + run_ends = np.concatenate((change_points, [vr_times.size])) + run_lengths = run_ends - run_starts + return int(run_lengths[0]), int(run_lengths[-1]) + + +def _trim_repetitive_boundary_timebins( + vr_times: np.ndarray, vr_values: np.ndarray +) -> tuple[np.ndarray, np.ndarray, int, int]: + """Drop consecutive repetitive onset_time_unity runs at start and end. + + When the first/last VR barcode events all map to the same Unity timebin, + those boundary runs are unreliable for cross-clock fitting and are removed. + """ + leading_run, trailing_run = _boundary_repetitive_run_lengths(vr_times) + n_trimmed_leading = leading_run if leading_run >= 2 else 0 + if leading_run >= 2: + vr_times = vr_times[leading_run:] + vr_values = vr_values[leading_run:] + + _, trailing_run = _boundary_repetitive_run_lengths(vr_times) + n_trimmed_trailing = trailing_run if trailing_run >= 2 else 0 + if trailing_run >= 2: + vr_times = vr_times[:-trailing_run] + vr_values = vr_values[:-trailing_run] + + return vr_times, vr_values, n_trimmed_leading, n_trimmed_trailing @dataclass(frozen=True) class BarcodeAlignmentFit: - """Linear fit + interpolator mapping VR time to NP time.""" + """Linear fit + interpolator mapping VR time to NP time, with diagnostics.""" slope: float intercept: float r2: float + rmse_ms: float + max_abs_residual_ms: float interpol_func: scipy.interpolate.interp1d shared_barcodes: np.ndarray + n_trimmed_leading: int + n_trimmed_trailing: int + n_rejected_outliers: int + + +def _inlier_mask( + vr_shared_times: np.ndarray, + np_shared_times: np.ndarray, + *, + sigma: float = OUTLIER_SIGMA, + floor_ms: float = OUTLIER_FLOOR_MS, + max_fraction: float = OUTLIER_MAX_FRACTION, + max_iterations: int = 10, +) -> np.ndarray: + """Iteratively reject tie points whose residual is a robust outlier. + + Residuals are centered on their median before scoring. A least-squares line + through a contaminated set can sit between clean and displaced points; an + uncentered rule then risks dropping all points instead of the offenders. + """ + keep = np.ones(vr_shared_times.shape, dtype=bool) + for _ in range(max_iterations): + fit = scipy.stats.linregress(vr_shared_times[keep], np_shared_times[keep]) + residual_ms = ( + np_shared_times - (fit.slope * vr_shared_times + fit.intercept) + ) * 1000.0 + center = float(np.median(residual_ms[keep])) + scale = 1.4826 * float(np.median(np.abs(residual_ms[keep] - center))) + threshold = max(float(floor_ms), float(sigma) * scale) + updated = np.abs(residual_ms - center) <= threshold + n_dropped = int((~updated).sum()) + if ( + n_dropped > max_fraction * vr_shared_times.size + or int(updated.sum()) < MIN_TIE_POINTS + ): + raise ValueError( + f"{n_dropped} of {vr_shared_times.size} barcode tie points are residual " + f"outliers (more than {max_fraction:.0%}); the VR/NP relation is not " + "simply linear for this session and no fit through it is trustworthy" + ) + if np.array_equal(updated, keep): + break + keep = updated + + return keep def align_barcodes( @@ -33,7 +117,8 @@ def align_barcodes( vr_values: np.ndarray, np_times: np.ndarray, np_values: np.ndarray, - skip_first_n_barcodes: int = 0, + *, + reject_outliers: bool = True, ) -> BarcodeAlignmentFit: """Fit VR time -> NP time from barcode values shared between both streams. @@ -42,17 +127,66 @@ def align_barcodes( https://github.com/AdaptiveMotorControlLab/auxPipelines-DataJoint_Mathis, adapted for this repo's VR (vr4mice) / NP (np_pipeline) schemas. + Before fitting, the function validates 1D and non-empty inputs, drops + barcode events with non-finite times on either stream, trims repetitive + boundary Unity timebins on the VR stream, intersects shared barcode values, + and requires at least ``MIN_TIE_POINTS`` shared tie points. + Args: vr_times: VR-side barcode onset times, ordered by event index (chronological). vr_values: VR-side barcode integer payloads, same order as `vr_times`. np_times: NP-side barcode onset times, ordered by event index. np_values: NP-side barcode integer payloads, same order as `np_times`. - skip_first_n_barcodes: number of leading (earliest) VR events to exclude - before matching, to avoid the DLC-live startup-lag window. + reject_outliers: When True, iteratively remove robust residual outliers + before fitting, using a median-centered MAD criterion with a + floor-based threshold in milliseconds. + + Raises: + ValueError: If shapes differ, inputs are not 1D, either stream is empty, + or fewer than ``MIN_TIE_POINTS`` shared tie points remain after + preprocessing. """ - if skip_first_n_barcodes: - vr_times = vr_times[skip_first_n_barcodes:] - vr_values = vr_values[skip_first_n_barcodes:] + vr_times = np.asarray(vr_times) + vr_values = np.asarray(vr_values) + np_times = np.asarray(np_times) + np_values = np.asarray(np_values) + + if vr_times.shape != vr_values.shape: + raise ValueError("vr_times and vr_values must have matching shapes") + if np_times.shape != np_values.shape: + raise ValueError("np_times and np_values must have matching shapes") + if vr_times.ndim != 1 or np_times.ndim != 1: + raise ValueError("barcode streams must be one-dimensional") + if vr_times.size == 0 or np_times.size == 0: + raise ValueError("both barcode streams must be non-empty") + + vr_finite = np.isfinite(vr_times) + if not np.all(vr_finite): + vr_times = vr_times[vr_finite] + vr_values = vr_values[vr_finite] + + np_finite = np.isfinite(np_times) + if not np.all(np_finite): + np_times = np_times[np_finite] + np_values = np_values[np_finite] + + if vr_times.size == 0 or np_times.size == 0: + raise ValueError( + "both barcode streams must contain at least one finite timepoint" + ) + + if np.unique(vr_times).size == 1: + raise ValueError( + "All barcode events have the same onset_time_unity; " + "cannot fit VR-to-NP alignment" + ) + + ( + vr_times, + vr_values, + n_trimmed_leading, + n_trimmed_trailing, + ) = _trim_repetitive_boundary_timebins(vr_times, vr_values) shared_barcodes, vr_index, np_index = np.intersect1d( vr_values, np_values, return_indices=True @@ -61,7 +195,25 @@ def align_barcodes( vr_shared_times = np.asarray(vr_times)[vr_index] np_shared_times = np.asarray(np_times)[np_index] + if vr_shared_times.size < MIN_TIE_POINTS: + raise ValueError( + f"Need at least {MIN_TIE_POINTS} shared barcode events after boundary trimming " + "to fit VR-to-NP alignment" + ) + + n_rejected_outliers = 0 + if reject_outliers: + keep = _inlier_mask(vr_shared_times, np_shared_times) + n_rejected_outliers = int((~keep).sum()) + vr_shared_times = vr_shared_times[keep] + np_shared_times = np_shared_times[keep] + shared_barcodes = shared_barcodes[keep] + linreg = scipy.stats.linregress(vr_shared_times, np_shared_times) + predicted_np = linreg.slope * vr_shared_times + linreg.intercept + residual_ms = (np_shared_times - predicted_np) * 1000.0 + rmse_ms = float(np.sqrt(np.mean(np.square(residual_ms)))) + max_abs_residual_ms = float(np.max(np.abs(residual_ms))) interpol_func = scipy.interpolate.interp1d( vr_shared_times, np_shared_times, @@ -73,6 +225,11 @@ def align_barcodes( slope=linreg.slope, intercept=linreg.intercept, r2=linreg.rvalue**2, + rmse_ms=rmse_ms, + max_abs_residual_ms=max_abs_residual_ms, interpol_func=interpol_func, shared_barcodes=shared_barcodes, + n_trimmed_leading=n_trimmed_leading, + n_trimmed_trailing=n_trimmed_trailing, + n_rejected_outliers=n_rejected_outliers, ) diff --git a/dj_pipeline/vr4mice/schema/barcodes.py b/dj_pipeline/vr4mice/schema/barcodes.py index ca2bf627..091fd3a6 100644 --- a/dj_pipeline/vr4mice/schema/barcodes.py +++ b/dj_pipeline/vr4mice/schema/barcodes.py @@ -140,7 +140,7 @@ class Event(dj.Part): barcode_value: int64 # Integer payload encoded by the barcode onset_sample: int64 # Teensy millisecond timestamp at event onset onset_time: float64 # Corresponding photodiode_time acquisition timestamp - onset_time_unity: float64 # Nearest base_analysis.DataFrame step_time + onset_time_unity=NULL: float64 # Nearest base_analysis.DataFrame step_time (NULL when outside step_time range) """ def make(self, key): @@ -191,7 +191,11 @@ def make(self, key): "barcode_value": event.value, "onset_sample": event.onset_sample, "onset_time": event.onset_time, - "onset_time_unity": onset_step_time, + "onset_time_unity": ( + float(onset_step_time) + if np.isfinite(onset_step_time) + else None + ), } for event, onset_step_time in zip( result.events, event_step_times, strict=True diff --git a/dj_pipeline/vr4mice/schema/np_sync.py b/dj_pipeline/vr4mice/schema/np_sync.py index 62ddcecf..834fadba 100644 --- a/dj_pipeline/vr4mice/schema/np_sync.py +++ b/dj_pipeline/vr4mice/schema/np_sync.py @@ -29,7 +29,7 @@ sys.path.insert(0, "/app/np_pipeline") from np_pipeline.schemas import barcodes as np_barcodes, session_link -from vr4mice.analysis.np_sync import DEFAULT_SKIP_FIRST_N_BARCODES, align_barcodes +from vr4mice.analysis.np_sync import align_barcodes from vr4mice.schema import base from vr4mice.schema import barcodes as vr_barcodes from vr4mice.schema import vr4mice @@ -75,10 +75,14 @@ class BarcodeSync(dj.Computed): -> vr_barcodes.TeensyBarcodes -> np_barcodes.ProbeBarcodeExtraction --- - skip_first_n_barcodes: int32 # leading VR barcode events excluded from the fit slope: float64 # Slope of the linear fit mapping VR time to NP time intercept: float64 # Intercept of the linear fit mapping VR time to NP time - r2: float64 # R-squared value of the linear fit + r2: float64 # R-squared of the fit. NOT a quality signal -- gate on rmse_ms + rmse_ms: float64 # RMS fit residual in milliseconds; the quality gate + max_abs_residual_ms: float64 # Largest single tie-point residual, milliseconds + n_shared_barcodes: int32 # Tie points the fit actually used + n_unmapped_vr_events: int32 # VR events dropped because onset_time_unity is NULL/NaN/inf + n_rejected_outliers: int32 # Tie points dropped as residual outliers interpol_func: # pickled scipy.interpolate.interp1d, VR time -> NP time barcode_overlap: float64 # Fraction of NP barcodes also found on the VR side """ @@ -93,9 +97,12 @@ def key_source(self): ) return source - skip_first_n_barcodes = DEFAULT_SKIP_FIRST_N_BARCODES min_shared_barcodes = 20 min_barcode_overlap = 0.90 + # The healthy cohort spans 5.98-8.10 ms against a Unity quantization floor of + # 18.5/sqrt(12) ~ 5.3 ms. 15 ms is ~2x the worst good fit and still an order of + # magnitude below the failures it catches. + max_rmse_ms = 15.0 def make(self, key): """Fit and insert one VR-time-to-NP-time alignment. @@ -145,6 +152,32 @@ def make(self, key): np_barcodes.ProbeBarcodeExtraction.Event & key ).to_arrays("barcode_value", "onset_time", order_by="barcode_index") + vr_values = np.asarray(vr_values) + vr_times = np.asarray(vr_times, dtype=np.float64) + np_values = np.asarray(np_values) + np_times = np.asarray(np_times, dtype=np.float64) + + vr_finite = np.isfinite(vr_times) + n_vr_dropped = int((~vr_finite).sum()) + if n_vr_dropped: + vr_values = vr_values[vr_finite] + vr_times = vr_times[vr_finite] + + np_finite = np.isfinite(np_times) + n_np_dropped = int((~np_finite).sum()) + if n_np_dropped: + np_values = np_values[np_finite] + np_times = np_times[np_finite] + + if n_vr_dropped or n_np_dropped: + logger.info( + "%s dropped %d non-finite VR and %d non-finite NP barcode events for %s", + self.__class__.__name__, + n_vr_dropped, + n_np_dropped, + key, + ) + if len(np_values) == 0: reason = ( "No NP barcode events found for key at populate time " @@ -166,11 +199,10 @@ def make(self, key): vr_values, np_times, np_values, - skip_first_n_barcodes=self.skip_first_n_barcodes, ) - # Measure overlap on full streams (independent of fit-time skipping) - # so DEFAULT_SKIP_FIRST_N_BARCODES does not bias this quality metric. + # Measure overlap on full streams (independent of fit-time trimming) + # so fit preprocessing does not bias this quality metric. full_shared_barcodes = np.intersect1d(vr_values, np_values) np_unique_barcodes = np.unique(np_values) barcode_overlap = len(full_shared_barcodes) / len(np_unique_barcodes) @@ -209,22 +241,49 @@ def make(self, key): ) return + if fit.rmse_ms > self.max_rmse_ms: + reason = ( + "Barcode alignment residuals too large for reliable NP-VR " + f"alignment (rmse_ms={fit.rmse_ms:.2f}, " + f"max_allowed={self.max_rmse_ms:.2f}, " + f"max_abs_residual_ms={fit.max_abs_residual_ms:.1f}, " + f"n_shared={len(fit.shared_barcodes)})" + ) + vr4mice.FailedSession().add_entry( + f"{key['dataset']}", f"{self.__class__.__name__}", reason + ) + logger.warning( + "%s %s for dataset %s", + self.__class__.__name__, + reason, + key["dataset"], + ) + return + self.insert1( { **key, - "skip_first_n_barcodes": self.skip_first_n_barcodes, "slope": fit.slope, "intercept": fit.intercept, "r2": fit.r2, + "rmse_ms": fit.rmse_ms, + "max_abs_residual_ms": fit.max_abs_residual_ms, + "n_shared_barcodes": len(fit.shared_barcodes), + "n_unmapped_vr_events": n_vr_dropped, + "n_rejected_outliers": fit.n_rejected_outliers, "interpol_func": pickle.dumps(fit.interpol_func), "barcode_overlap": barcode_overlap, } ) logger.info( - "%s aligned %d shared barcodes for %s", + "%s aligned %d shared barcodes for %s (rmse=%.2f ms, " + "dropped %d unmapped VR events, rejected %d outliers)", self.__class__.__name__, len(fit.shared_barcodes), key, + fit.rmse_ms, + n_vr_dropped, + fit.n_rejected_outliers, ) except Exception as err: diff --git a/tests/unit/test_barcodes.py b/tests/unit/test_barcodes.py index c876b68c..bac3fe3a 100644 --- a/tests/unit/test_barcodes.py +++ b/tests/unit/test_barcodes.py @@ -1,5 +1,7 @@ """Unit tests for Teensy barcode decoding.""" +from pathlib import Path + import numpy as np import pytest @@ -11,6 +13,14 @@ ) from dlc_helpers import align_timestamps_to_step_time +SCHEMA_BARCODES_PY = ( + Path(__file__).parent.parent.parent + / "dj_pipeline" + / "vr4mice" + / "schema" + / "barcodes.py" +) + def _barcode_edges(value: int, *, start_ms: int): config = BarcodeDecoderConfig() @@ -111,7 +121,7 @@ def test_decode_teensy_barcodes_returns_no_events_for_constant_signal(): [ ([0, 1], [0], [10, 11], "same shape"), ([0, 1], [0, 1], [10], "same shape"), - ([0, 0], [0, 1], [10, 11], "strictly increasing"), + ([1, 0], [0, 1], [10, 11], "non-decreasing"), ], ) def test_decode_teensy_barcodes_validates_sample_arrays( @@ -121,6 +131,17 @@ def test_decode_teensy_barcodes_validates_sample_arrays( decode_teensy_barcodes(times, states, photodiode_times) +def test_decode_teensy_barcodes_allows_duplicate_times_with_same_ttl_state(): + result = decode_teensy_barcodes( + [0, 0, 1, 2], + [0, 0, 0, 0], + [10.0, 10.0, 11.0, 12.0], + ) + + assert result.events == () + assert result.quality["edge_count"] == 0 + + def test_align_timestamps_to_data_frame_step_time(): step_time = np.asarray([0.10, 0.20, 0.30, 0.40]) @@ -130,3 +151,29 @@ def test_align_timestamps_to_data_frame_step_time(): ) assert result.tolist() == pytest.approx([0.10, 0.30, 0.40]) + + +def test_align_timestamps_to_step_time_returns_nan_outside_step_range(): + step_time = np.asarray([0.10, 0.20, 0.30, 0.40]) + + result = align_timestamps_to_step_time( + np.asarray([0.05, 0.14, 0.45]), + step_time, + ) + + assert np.isnan(result[0]) + assert result[1] == pytest.approx(0.10) + assert np.isnan(result[2]) + + +def test_teensy_barcodes_schema_allows_null_onset_time_unity(): + text = SCHEMA_BARCODES_PY.read_text() + + assert "onset_time_unity=NULL: float64" in text + + +def test_teensy_barcodes_make_converts_non_finite_unity_times_to_null(): + text = SCHEMA_BARCODES_PY.read_text() + + assert "if np.isfinite(onset_step_time)" in text + assert "else None" in text diff --git a/tests/unit/test_np_sync.py b/tests/unit/test_np_sync.py index 991c7f29..1255a58a 100644 --- a/tests/unit/test_np_sync.py +++ b/tests/unit/test_np_sync.py @@ -44,7 +44,26 @@ def test_align_barcodes_recovers_known_linear_fit(): assert fit.slope == pytest.approx(2.0) assert fit.intercept == pytest.approx(100.0) assert fit.r2 == pytest.approx(1.0) + assert fit.rmse_ms == pytest.approx(0.0) + assert fit.max_abs_residual_ms == pytest.approx(0.0) assert len(fit.shared_barcodes) == 20 + assert fit.n_trimmed_leading == 0 + assert fit.n_trimmed_trailing == 0 + assert fit.n_rejected_outliers == 0 + + +def test_align_barcodes_is_a_no_op_on_a_clean_stream(): + """Clean sessions should be left untouched by guards and robust rejection.""" + vr_times, vr_values, np_times, np_values = _linear_barcode_streams(n=40) + + guarded = align_barcodes(vr_times, vr_values, np_times, np_values) + raw = align_barcodes( + vr_times, vr_values, np_times, np_values, reject_outliers=False + ) + + assert (guarded.slope, guarded.intercept) == (raw.slope, raw.intercept) + assert guarded.n_rejected_outliers == 0 + assert (guarded.n_trimmed_leading, guarded.n_trimmed_trailing) == (0, 0) def test_align_barcodes_uses_only_shared_barcode_values(): @@ -61,21 +80,33 @@ def test_align_barcodes_uses_only_shared_barcode_values(): assert fit.slope == pytest.approx(2.0) -def test_align_barcodes_skip_first_n_excludes_leading_vr_events(): - vr_times, vr_values, np_times, np_values = _linear_barcode_streams(n=20) - # Corrupt the first two VR onset times so an unskipped fit would be biased. +def test_align_barcodes_trims_repetitive_boundary_timebins(): + vr_values = np.arange(8) + np_values = np.arange(8) + + # First two and last two events collapse onto one Unity timebin each. + vr_times = np.array([0.0, 0.0, 2.0, 3.0, 4.0, 5.0, 6.0, 6.0], dtype=np.float64) + np_times = 2.0 * np.arange(8, dtype=np.float64) + 100.0 + + fit = align_barcodes(vr_times, vr_values, np_times, np_values) + assert fit.slope == pytest.approx(2.0) + assert fit.intercept == pytest.approx(100.0) + assert len(fit.shared_barcodes) == 4 + assert fit.n_trimmed_leading == 2 + assert fit.n_trimmed_trailing == 2 + + +def test_align_barcodes_reports_real_repetitive_boundary_prefix_length(): + n = 60 + vr_times, vr_values, np_times, np_values = _linear_barcode_streams(n=n) vr_times = vr_times.copy() - vr_times[:2] += 1000.0 + vr_times[:17] = vr_times[16] - biased_fit = align_barcodes(vr_times, vr_values, np_times, np_values) - assert biased_fit.slope != pytest.approx(2.0) + fit = align_barcodes(vr_times, vr_values, np_times, np_values) - corrected_fit = align_barcodes( - vr_times, vr_values, np_times, np_values, skip_first_n_barcodes=2 - ) - assert corrected_fit.slope == pytest.approx(2.0) - assert corrected_fit.intercept == pytest.approx(100.0) - assert len(corrected_fit.shared_barcodes) == 18 + assert fit.n_trimmed_leading == 17 + assert len(fit.shared_barcodes) == n - 17 + assert fit.slope == pytest.approx(2.0) def test_align_barcodes_interpol_func_maps_vr_time_to_np_time(): @@ -88,6 +119,95 @@ def test_align_barcodes_interpol_func_maps_vr_time_to_np_time(): assert fit.interpol_func(5.0) == pytest.approx(110.0) +def test_align_barcodes_drops_non_finite_tie_points(): + vr_times, vr_values, np_times, np_values = _linear_barcode_streams(n=40) + vr_times = vr_times.copy() + vr_times[4] = np.nan + np_times = np_times.copy() + np_times[7] = np.inf + + fit = align_barcodes(vr_times, vr_values, np_times, np_values) + + assert fit.slope == pytest.approx(2.0) + assert fit.intercept == pytest.approx(100.0) + assert len(fit.shared_barcodes) == 38 + + +def test_align_barcodes_requires_finite_tie_points_after_filtering(): + vr_times = np.array([np.nan, np.nan, np.nan]) + vr_values = np.array([1, 2, 3]) + np_times = np.array([100.0, 101.0, 102.0]) + np_values = np.array([1, 2, 3]) + + with pytest.raises(ValueError, match="finite timepoint"): + align_barcodes(vr_times, vr_values, np_times, np_values) + + +def test_align_barcodes_requires_three_tie_points(): + values = np.arange(2) + + with pytest.raises(ValueError, match="at least 3"): + align_barcodes( + np.array([0.0, 1.0]), values, np.array([100.0, 102.0]), values + ) + + +def test_align_barcodes_accepts_reject_outliers_keyword(): + vr_times, vr_values, np_times, np_values = _linear_barcode_streams( + slope=2.0, intercept=100.0 + ) + + fit = align_barcodes( + vr_times, + vr_values, + np_times, + np_values, + reject_outliers=False, + ) + assert fit.slope == pytest.approx(2.0) + + +def test_align_barcodes_rejects_a_single_displaced_tie_point(): + vr_times, vr_values, np_times, np_values = _linear_barcode_streams(n=40) + np_times = np_times.copy() + np_times[17] += 0.725 + + fit = align_barcodes(vr_times, vr_values, np_times, np_values) + + assert fit.n_rejected_outliers == 1 + assert fit.slope == pytest.approx(2.0) + assert fit.rmse_ms == pytest.approx(0.0, abs=1e-6) + + +def test_align_barcodes_rejection_floor_spares_unity_quantization(): + for displacement_ms, expected in ((29, 0), (31, 1)): + vr_times, vr_values, np_times, np_values = _linear_barcode_streams(n=200) + np_times = np_times.copy() + np_times[100] += displacement_ms / 1000.0 + + fit = align_barcodes(vr_times, vr_values, np_times, np_values) + + assert fit.n_rejected_outliers == expected, displacement_ms + + +def test_align_barcodes_raises_when_disagreement_is_not_isolated(): + vr_times, vr_values, np_times, np_values = _linear_barcode_streams(n=40) + np_times = np_times.copy() + np_times[::4] += 0.5 + + with pytest.raises(ValueError, match="not simply linear"): + align_barcodes(vr_times, vr_values, np_times, np_values) + + +def test_align_barcodes_raises_on_degenerate_onset_time_unity(): + vr_values = np.arange(6) + + with pytest.raises(ValueError, match="same onset_time_unity"): + align_barcodes( + np.zeros(6), vr_values, np.arange(6, dtype=np.float64), vr_values + ) + + def test_analysis_np_sync_has_no_datajoint_or_np_pipeline_dependency(monkeypatch): """The pure alignment math must import and run with np_pipeline/datajoint absent.""" monkeypatch.setitem(sys.modules, "datajoint", None) @@ -183,6 +303,15 @@ def test_schema_make_handles_empty_np_events_with_clear_reason(): assert "No NP barcode events found for key at populate time" in make_src +def test_schema_make_filters_non_finite_tie_point_times_before_alignment(): + text = SCHEMA_NP_SYNC_PY.read_text() + make_src = _function_source(text, " def make(self, key):") + + assert "vr_finite = np.isfinite(vr_times)" in make_src + assert "np_finite = np.isfinite(np_times)" in make_src + assert "dropped %d non-finite VR and %d non-finite NP barcode events" in make_src + + def test_schema_make_uses_key_only_without_identity_parsing(): text = SCHEMA_NP_SYNC_PY.read_text() make_src = _function_source(text, " def make(self, key):") @@ -207,8 +336,21 @@ def test_schema_make_has_quality_gate_for_min_shared_and_overlap(): assert "min_shared_barcodes = 20" in text assert "min_barcode_overlap = 0.90" in text + assert "max_rmse_ms = 15.0" in text assert "Insufficient shared barcodes for reliable NP-VR alignment" in make_src assert "Insufficient NP-VR barcode overlap for reliable alignment" in make_src + assert "Barcode alignment residuals too large for reliable NP-VR" in make_src + assert "max_allowed={self.max_rmse_ms:.2f}" in make_src + + +def test_schema_definition_stores_alignment_diagnostics(): + text = SCHEMA_NP_SYNC_PY.read_text() + + assert "rmse_ms: float64" in text + assert "max_abs_residual_ms: float64" in text + assert "n_shared_barcodes: int32" in text + assert "n_unmapped_vr_events: int32" in text + assert "n_rejected_outliers: int32" in text def test_schema_module_docstring_states_vr_only_sessions_excluded_from_key_source():