diff --git a/README.md b/README.md index d8a7721..9a10012 100644 --- a/README.md +++ b/README.md @@ -12,7 +12,7 @@ and asks a sharper estimation-theoretic question: > evidence that the world model itself should change—and can that event signal > still be identified after realistic release and measurement dynamics? -Version 0.5 implements six computational layers: +Version 0.7 implements seven computational layers: 1. conjugate Dirichlet learning for one transition model; 2. exact HMM filtering over already learned transition contexts; @@ -25,7 +25,11 @@ Version 0.5 implements six computational layers: baselines, and held-out candidate comparison; 6. causal closed-loop triggering with explicit transport delay, randomized latency, local eligibility traces, yoked sham controls, and held-out recovery - of the causal stimulation window. + of the causal stimulation window; +7. a versioned, leakage-audited replay artifact contract and a prespecified + animal-level comparison of behavioral filtering-to-smoothing revision against + online-surprise, content, location, recency, prospective, and TD-error fields, + with recovery gates and explicit abstention. This repository is a computational hypothesis-testing project. It does **not** claim that any candidate has already been established as the biological ACh @@ -307,10 +311,13 @@ whether the discrete timescale posterior is concentrated and stable. - **Stage 5 — complete:** causal online triggering, independently calibrated delay, randomized timing, yoked active/sham perturbation, eligibility-family recovery, and falsification against null and latency-independent effects. -- **Stage 6 — next:** replay as smoothing-based revision rather than - unconstrained internally generated prediction error. +- **Stage 6 — implemented, claim-gated:** a leakage-audited replay artifact + contract, filtering-to-smoothing spatial-field comparison, animal-level + simultaneous contrasts, and post-decoder recovery gates. A strict real-data + freeze is still required before making any biological replay claim. -See [`docs/closed_loop.md`](docs/closed_loop.md) for causal triggering and +See [`docs/pf_replay_spatial_revision.md`](docs/pf_replay_spatial_revision.md) +for the replay contract, claim boundary, and fixed analysis; [`docs/closed_loop.md`](docs/closed_loop.md) for causal triggering and eligibility-window recovery, [`docs/measurement_model.md`](docs/measurement_model.md) for the ACh measurement derivation, [`docs/partial_observation.md`](docs/partial_observation.md) for multisensory diff --git a/docs/pf_replay_spatial_revision.md b/docs/pf_replay_spatial_revision.md new file mode 100644 index 0000000..f293393 --- /dev/null +++ b/docs/pf_replay_spatial_revision.md @@ -0,0 +1,306 @@ +# Pfeiffer/Foster spatial replay-revision contract + +## Claim boundary + +Pfeiffer/Foster replay is strongly prospective by task design. This analysis does +not relabel it as retrospective replay. It asks a narrower model-comparison +question: + +> Is the spatial content of an awake replay event better explained by a signed +> record of earlier behavior-state revisions than by prospective, recency, +> content, finite prediction-error, current-location, or null priors? + +A positive result would establish a geometry consistent with smoothing-based +revision under the specified behavioral model. It would not establish that +replay is generated by smoothing, that the animal implements this exact model, +or that acetylcholine encodes the revision. + +The currently frozen Pfeiffer/Foster posterior-bin CSV contains posterior means, +MAP positions, and entropies but not the full raw spatial log-emission tensor. +It is insufficient for this primary comparison. The raw likelihood artifact +must therefore be regenerated; MAP-path scoring is not an acceptable substitute. + +## Frozen latent variable and observation schedule + +Before a real-data run, version 2 fixes the behavioral construction as follows. + +- Latent state: the compact destination-well identity for a completed + historical RUN traversal (at most the session's well library, not the full + PF spatial grid). Early in a traversal this state is uncertain; the later + approach and well visit can revise the earlier belief. +- External observation: tracked x/y frames and the terminal well visit from + that RUN traversal, evaluated under behavior-only route-to-well templates. +- Transition: a first-order well transition fitted only from RUN behavior whose + timestamps precede the target replay event. Bayesian-ACh stores this compact + transition as row-stochastic P[source, destination]. A Hippo trace must export + its exact transpose T[destination, source]; the artifact records the + convention. +- Snippet: one movement segment from a completed well-to-well fill interval. + Filtering is prefix-only inside the segment; fixed-interval smoothing may use + later tracking frames in that same segment. Because segmentation, smoothing, + and the destination window consume the full source fill interval, the route + becomes available only at `interval_end_time_s`. A movement end before the + replay is not sufficient when its fill interval ends later. Smoothing never + crosses the replay event. +- Revision field: for historical time h, + `KL(p_smooth,h || p_filter,h) * (p_smooth,h - p_filter,h)`, projected from + compact wells onto the replay grid with snippet-specific route kernels learned + only from earlier behavior, then summed with a predeclared exponential age + weight. The sign is retained; negative revisions are not rectified. +- Neural observation: the replay spike-derived spatial log-emission tensor + during the event. It is used only as the held-out quantity to score and is + never fed back into the behavioral smoother. + +If the tracked behavior and transition model produce negligible +filtering-to-smoothing KL, the revision candidate is unidentifiable and the +analysis abstains. That outcome is not evidence against neural replay. + +## Exact revision field and interpretation + +For replay event e, the exported field is + +```text +R_e(x) = sum_r KL(q_r || f_r) + * exp(-(start_e - end_r) / 300 s) + * sum_k (q_r(k) - f_r(k)) M_r,k(x) +z_e(x) = R_e(x) / sum_x |R_e(x)| +``` + +Here r ranges over routes whose fill intervals end before the ripple, +`end_r` is that fill-interval end (the evidence-availability and age timestamp), +f is the destination-well filter at the fixed route prefix, q is the smoother at that +same historical point, and M maps each well state to the event grid. The +downstream score base-centers and base-standardizes z, then tunes only a +nonnegative temperature on training animals. Thus the primary comparison tests +the signed spatial pattern, not the total amount of historical revision. + +| Modeling choice | Frozen value | Consequence | +|---|---:|---| +| latent state | destination-well identity | no within-route state switching | +| within-route transition | identity matrix | destination is static | +| transition pseudocount | 0.5 | regularizes sparse origin-to-well counts | +| route resampling | 21 arc points | fixes correspondence across paths | +| filtering prefix | 50% (11 of 21 points) | defines the historical revision time | +| path observation width | 15 cm | fixes template likelihood resolution | +| terminal-label error | 0.02 | strongly anchors the completed destination | +| spatial route-kernel width | 10 cm | maps well revisions to decoder bins | +| revision age constant | 300 s | exponentially discounts older routes | +| revision magnitude | KL(q || f) | asymmetric information-gain weight | +| field normalization | divide by sum absolute mass | preserves sign, removes scale | +| score normalization | nuisance-base centering/SD | further removes eventwise scale | +| temperature grid | 0, .25, .5, 1, 2, 4, 8 | polarity fixed; strength tuned LO-rat | + +A route template uses earlier same-origin/same-destination routes when present +and otherwise falls back to earlier routes with the same destination. That +fallback, the hyperparameters above, the all-route aggregation, the active +decoder grid, and complete-case selection are modeling choices, not identified +biological quantities. + +The replay candidate prior is constant across event time bins and the score is +a mean of per-bin spatial marginal log likelihoods. Therefore a positive +contrast would identify better time-marginal spatial alignment with this field, +not replay sequence order, forward/reverse direction, or ordered reinstatement. +It would not identify a smoothing generator, the animal's algorithm, causality, +or acetylcholine coding. + +## Cutoff semantics + +For event e with start time s_e and end time f_e: + +```text +decoder_training_cutoff_s[e] <= s_e +history_cutoff_s[e] <= s_e +field_available_s[e, candidate] <= s_e +replay log emissions use only [s_e, f_e] +place-field and decoder training observations end at decoder_training_cutoff_s[e] +later outcome time > f_e +``` + +The event cohort must also be independent of decoded spatial content. The +claim-bearing producer reselects RUN ripples by raw LFP peak power under a +predeclared top-N-per-session schedule. Ranking is performed over the full +session, so this is an offline conditional-content sample—not an online causal +claim about replay-event incidence. The earlier 160-event table selected +downstream of full-session decoder evidence is not an admissible primary cohort. +Every raw dataset path, size, and SHA-256 is checked against the canonical lock; +missing, changed, and unlocked extra files fail before export. Historical route +tables must be regenerated from a clean commit and smoothed independently +within each completed fill interval. Event eligibility, behavioral history, +recency age, route-specific template training, and historical well locations all +gate on the fill-interval end, not the earlier trimmed movement end. + +The decoder must be refitted at each event cutoff (or use a demonstrably +equivalent prefix cache). The existing producer's one-time, full-session RUN +place-field fit violates this rule and cannot produce a claim-bearing artifact. +A timestamp copied into the manifest is not sufficient: the fitted observations +must actually end at that timestamp. + +Decoder point-spread calibration uses a fixed 120-second interval ending +strictly before each replay, 100-ms bins, moving RUN samples outside ripples, a +minimum of 20 valid bins, and the 68th-percentile position error. This is a +documented design amendment, not a preregistered choice: the initial 60-second +rule left two of the 160 fixed LFP-selected events with only 17 and 10 valid +bins. An outcome-blind, score-blind counts preflight over 60, 120, and 180 +seconds found 120 seconds to be the shortest tested global window supporting +all 160 events (minimum 42 bins). The global window was changed before replay +scoring; no event was dropped or backfilled. Only valid bins are evaluated, in +deterministic chunks of at most 32 time bins. The consumer requires decoder +configuration digest +`a79fa8a1f55a964c4367853cc120efc9b742ec4e327c277ede78cfd6a277f20b`; +tracked positions are truncated at the strict event cutoff before speed and +actual-position interpolation, preventing centered-gradient boundary leakage. +this binds the fixed window, bin width, support threshold, q68 statistic, +chunk size, and the unchanged encoder/emission settings. + +The compact well state keeps Bayesian-ACh pair arrays small; the PF spatial +grid is never passed through its dense pair-marginal implementation. Spatial +traces, if needed, use the sparse Hippo trace. The prospective field uses the +behavior-only transition model available at s_e; it cannot use the route the +animal actually takes later. Recency, +current-location, posterior-content, online-surprise, and finite TD-error +fields likewise use only observations available by s_e. + +Later valid outcomes, including the next visited well, are stored in a +physically separate hash-bound artifact. They may be used for a secondary +behavioral association only after the predictor artifact has been frozen. They +are not inputs to candidate construction, temperature selection, recovery, or +the primary neural likelihood comparison. + +## Candidate priors and raw-emission score + +Every candidate shares the same event-specific nuisance base over active +spatial bins. The base may encode decoder support and a predeclared occupancy +floor, but it cannot use candidate-specific evidence. For candidate field +z_c,e(x) and a temperature selected on training animals only, + +```text +q_c,e(x) proportional to base_e(x) * exp(temperature_c * z_c,e(x)). +``` + +The null has temperature zero. The primary event score marginalizes every valid +raw neural log-emission row: + +```text +score_c,e = mean_t log sum_x exp(log_emission_e,t(x)) q_c,e(x). +``` + +Per-row log-emission offsets are retained so absolute evidence can be +reconstructed; candidate differences are invariant to those offsets. This +score deliberately avoids MAP paths. A shared-dynamics HMM sensitivity may be +added later, but no candidate may receive different transition flexibility. + +The frozen candidates are: + +1. `smoothing_revision`: the signed KL-weighted historical revision field; +2. `online_surprise`: spatially binned behavior-filter predictive surprise; +3. `posterior_content`: historical posterior occupancy without differencing; +4. `current_location`: a kernel around the pre-replay tracked location; +5. `recency`: exponentially aged pre-replay behavior occupancy; +6. `prospective`: the behavior-only next-state distribution available online; +7. `td_error`: finite, clipped reward/value prediction errors observed before + replay and mapped to their locations; and +8. `null`: the shared nuisance base alone. + +## Grouping and abstention + +Temperatures are selected without the held-out rat. Event scores are averaged +within session and then equally across rats. The confirmatory target is fixed in +advance as `smoothing_revision`; a candidate selected after viewing all +held-out rats is reported only as a descriptive winner. + +For inference, the software forms paired rat-level contrasts between +`smoothing_revision` and every finite alternative, including prospective, +recency, posterior content, TD error, and null. A joint animal bootstrap uses +the maximum shortfall over all contrasts to produce simultaneous one-sided +lower bounds. Each rat-level contrast also receives an exact one-sided sign-flip +test. Smoothing is identified only if every simultaneous lower bound is +strictly positive and every exact p value is at most the prespecified alpha. +Consequently four rats can never support a 95% animal-level identification +claim (the smallest attainable one-sided p value is 1/16). The descriptive winner-minus-runner interval never controls +the claim. + +Recovery is executable rather than asserted. It first subsets to the exact +all-candidate complete-case cohort, retaining original event IDs and reporting +every exclusion. The recovery runner draws latent +spatial bins from each frozen candidate prior, emits Gaussian raw spatial +log-likelihood rows around those bins, and sends those rows through the full +temperature-selection and scoring code. Coordinates and widths are in +centimeters. Gaussian width is not an arbitrary grid constant: it is the +event-specific held-out RUN decoder point-spread/error multiplied by the +prespecified stress range 0.5, 1.0, and 2.0. + +This is post-decoder scoring recovery at empirical decoder resolution; it is +not end-to-end spike/place-field decoder recovery. Because injection uses the +same frozen candidate fields that the scorer receives, successful recovery +checks implementation and finite decoder resolution, not robustness to a +misspecified smoothing field or to alternative behavioral hyperparameters. A nuisance-base/null +injection is a required negative control: no non-null candidate may be decisive. +Every pure generator must be recovered with fixed-generator simultaneous +contrasts under both leave-one-rat-out and leave-one-session-out calibration at +every width. Registered 50/50 mixtures pair `smoothing_revision` with TD error, +prospective, recency, and posterior-content candidates. A mixture passes only +when no pure candidate (including null) is decisive; a decisive alternative win +is a recovery failure, not mixture abstention. The persisted gate consists of +the per-generator, per-split, per-width records; it has no manually set pass +booleans. + +The software computes scores even when it must abstain. It reports +`status="identified"` only if all of the following hold: + +- the common complete-case cohort spans at least five rats and has adequate + exact animal sign-flip resolution; +- computed pure-candidate injection recovery succeeds under both split schemes; +- the registered 50/50 mixture remains ambiguous; +- candidate-field collinearity remains below the frozen threshold; and +- every prespecified smoothing contrast has a positive simultaneous lower bound. + +Failure of any gate yields `abstain`. A zero-event complete-case cohort still +freezes a technical-abstention artifact with every excluded event ID, empty +header-only score/recovery tables, failed gates, and hashes; it does not crash +or silently disappear. An abstention must not be rewritten as a biological null +or as evidence for unconstrained prediction error. + +## Artifact files + +Schema `bayesian-ach.replay-spatial.v2` writes: + +- `replay_spatial_predictors.npz`: raw shifted log emissions, offsets, masks, + event-specific centimeter spatial coordinates, held-out RUN decoder + point-spread/error, nuisance base, all pre-replay candidate fields, + availability times, optional posterior-derived well masses, and + rat/session/event identifiers; +- `replay_spatial_manifest.json`: clean producer commit, locked-dataset + digest and deterministic full-tree verification report, clean route-producer + commit/config/table hashes, ordered cohort and event-audit hashes, + transition/trace convention, offline raw-LFP event-selection schedule, + pre-event temporal-holdout point-spread schedule, selection/behavior/decoder + parameter digests, the named event-audit and copied route-provenance files, + and the predictor SHA-256 digest. + +Schema `bayesian-ach.replay-later-outcome.v1` writes separate outcome NPZ and +manifest files bound to the already frozen predictor digest. Loading uses +`allow_pickle=False`. It opens and hash-checks the actual full-tree verifier +report, copied clean route manifest, event audit, and predictor. Verifier report +content, route producer/config/table hashes, audit identifiers, and the ordered +cohort digest must all agree with the predictor manifest before analysis. + + +## Frozen real-data runner + +Run from a clean Bayesian-ACh worktree and write outside that worktree: + +```bash +env PYTHONPATH=src python scripts/run_pf_replay_spatial_analysis.py \ + --predictor-dir /path/to/pf-replay-spatial-contract \ + --output-dir /path/to/pf-replay-spatial-analysis +``` + +The claim-bearing runner fixes temperatures to +`(0, 0.25, 0.5, 1, 2, 4, 8)`, five minimum independent rats, 5,000 animal +bootstrap replicates, 95% simultaneous coverage, correlation threshold 0.98, +and seed 7. Recovery fixes injection temperature 4, empirical point-spread +multipliers `(0.5, 1, 2)`, 0.02-nat emission noise, seed 701, LOAO/LOSO pure +generators, the four registered smoothing mixtures, and a nuisance-base null +negative control. The analysis manifest binds the clean consumer commit, +predictor/producer/dataset/route hashes, frozen configs, exclusion IDs, all +compact output hashes, recovery status, and final abstention reasons. diff --git a/scripts/check_markdown_math.py b/scripts/check_markdown_math.py index 24979d1..88c9642 100644 --- a/scripts/check_markdown_math.py +++ b/scripts/check_markdown_math.py @@ -9,7 +9,6 @@ from pathlib import Path - ROOT = Path(__file__).resolve().parents[1] SKIP_PARTS = {".git", ".venv", "build", "dist"} diff --git a/scripts/run_pf_replay_spatial_analysis.py b/scripts/run_pf_replay_spatial_analysis.py new file mode 100644 index 0000000..9b59942 --- /dev/null +++ b/scripts/run_pf_replay_spatial_analysis.py @@ -0,0 +1,368 @@ +#!/usr/bin/env python3 +"""Freeze the prespecified PF replay spatial-revision comparison.""" + +from __future__ import annotations + +import argparse +import csv +import hashlib +import json +import re +import subprocess +from collections.abc import Sequence +from dataclasses import asdict +from datetime import datetime, timezone +from pathlib import Path +from typing import Any + +import numpy as np + +from bayesian_ach.replay_artifact import load_spatial_predictor_artifact +from bayesian_ach.replay_recovery import ( + SpatialInjectionRecoveryConfig, + run_spatial_recovery_checks, +) +from bayesian_ach.replay_spatial import ( + SpatialComparisonConfig, + compare_spatial_replay_candidates, +) + +ANALYSIS_SCHEMA = "bayesian-ach.pf-replay-spatial-analysis.v1" +FOLD_OUTPUT = "pf_replay_spatial_candidate_folds.csv" +CONTRAST_OUTPUT = "pf_replay_spatial_target_contrasts.csv" +RAT_SCORE_OUTPUT = "pf_replay_spatial_rat_scores.csv" +RECOVERY_OUTPUT = "pf_replay_spatial_recovery.csv" +EXCLUSION_OUTPUT = "pf_replay_spatial_recovery_exclusions.csv" +GATE_OUTPUT = "pf_replay_spatial_gates.csv" +REPORT_OUTPUT = "pf_replay_spatial_report.md" +MANIFEST_OUTPUT = "pf_replay_spatial_analysis_manifest.json" +_COMMIT_PATTERN = re.compile(r"^[0-9a-f]{40}$") +ROOT = Path(__file__).resolve().parents[1] + + +def _sha256(path: Path) -> str: + digest = hashlib.sha256() + with path.open("rb") as handle: + for chunk in iter(lambda: handle.read(1024 * 1024), b""): + digest.update(chunk) + return digest.hexdigest() + + +def _clean_commit() -> str: + commit = subprocess.run( + ["git", "rev-parse", "HEAD"], + cwd=ROOT, + check=True, + capture_output=True, + text=True, + ).stdout.strip() + dirty = subprocess.run( + ["git", "status", "--porcelain", "--untracked-files=normal"], + cwd=ROOT, + check=True, + capture_output=True, + text=True, + ).stdout + if _COMMIT_PATTERN.fullmatch(commit) is None: + raise ValueError("analysis must run from a committed Git checkout") + if dirty.strip(): + raise ValueError("analysis must run from a clean committed worktree") + return commit + + +def _write_rows( + path: Path, + rows: list[dict[str, Any]], + fieldnames: Sequence[str], +) -> None: + columns = list(fieldnames) + if not columns: + raise ValueError(f"fieldnames for {path.name} must not be empty") + if any(list(row) != columns for row in rows): + raise ValueError(f"rows for {path.name} have inconsistent columns") + with path.open("w", encoding="utf-8", newline="") as handle: + writer = csv.DictWriter(handle, fieldnames=columns, lineterminator="\n") + writer.writeheader() + writer.writerows(rows) + + +def run_analysis( + predictor_directory: str | Path, + output_directory: str | Path, + *, + comparison_config: SpatialComparisonConfig | None = None, + injection_config: SpatialInjectionRecoveryConfig | None = None, +) -> dict[str, Path]: + """Run recovery and real scoring, then freeze compact hash-bound evidence.""" + + consumer_commit = _clean_commit() + comparison_config = ( + SpatialComparisonConfig( + minimum_rats=5, + bootstrap_replicates=5000, + simultaneous_confidence_level=0.95, + seed=7, + ) + if comparison_config is None + else comparison_config + ) + injection_config = ( + SpatialInjectionRecoveryConfig( + injection_temperature=4.0, + spatial_sigma_multipliers=(0.5, 1.0, 2.0), + coordinate_units="cm", + emission_noise_sd_nats=0.02, + seed=701, + ) + if injection_config is None + else injection_config + ) + comparison_config.validate() + injection_config.validate() + + predictor_source = Path(predictor_directory) + frozen = load_spatial_predictor_artifact(predictor_source) + recovery = run_spatial_recovery_checks( + frozen.dataset, + comparison_config, + injection_config, + ) + comparison = compare_spatial_replay_candidates( + frozen.dataset, + comparison_config, + recovery_gate=recovery, + ) + + output = Path(output_directory) + output.mkdir(parents=True, exist_ok=True) + fold_rows = [asdict(record) for record in comparison.folds] + contrast_rows = [asdict(record) for record in comparison.target_contrasts] + rat_score_rows = [ + { + "rat": rat, + "candidate": candidate, + "mean_log_score_per_bin": float(comparison.rat_scores[rat_index, candidate_index]), + } + for rat_index, rat in enumerate(comparison.rat_ids) + for candidate_index, candidate in enumerate(comparison.candidate_names) + ] + recovery_rows = [ + {"recovery_kind": kind, **asdict(record)} + for kind, records in ( + ("pure", recovery.pure_records), + ("mixture", recovery.mixture_records), + ("null_negative_control", recovery.null_records), + ) + for record in records + ] + exclusion_rows = [ + {"event_id": event_id, "reason": "not_in_all_candidate_complete_case_cohort"} + for event_id in recovery.excluded_event_ids + ] + if not exclusion_rows: + exclusion_rows = [{"event_id": "", "reason": "none"}] + + target_lower = np.asarray( + [record.simultaneous_lower_bound for record in comparison.target_contrasts], + dtype=float, + ) + target_p = np.asarray( + [record.exact_one_sided_sign_flip_p for record in comparison.target_contrasts], + dtype=float, + ) + alpha = 1.0 - comparison_config.simultaneous_confidence_level + gate_rows = [ + { + "gate": "minimum_independent_rats", + "passed": len(comparison.rat_ids) >= comparison_config.minimum_rats, + "value": len(comparison.rat_ids), + "required": comparison_config.minimum_rats, + }, + { + "gate": "exact_animal_sign_flip_resolution", + "passed": bool(np.all(np.isfinite(target_p)) and np.all(target_p <= alpha + 1e-15)), + "value": float(np.max(target_p)) if target_p.size else float("nan"), + "required": f"all <= {alpha}", + }, + { + "gate": "simultaneous_smoothing_contrasts", + "passed": bool(np.all(np.isfinite(target_lower)) and np.all(target_lower > 0.0)), + "value": float(np.min(target_lower)) if target_lower.size else float("nan"), + "required": "all > 0", + }, + { + "gate": "post_decoder_recovery", + "passed": recovery.passed, + "value": ( + len(recovery.pure_records) + + len(recovery.mixture_records) + + len(recovery.null_records) + ), + "required": "all pure pass; mixtures and null abstain", + }, + { + "gate": "candidate_field_collinearity", + "passed": comparison.maximum_field_correlation + <= comparison_config.maximum_field_correlation, + "value": comparison.maximum_field_correlation, + "required": comparison_config.maximum_field_correlation, + }, + { + "gate": "overall_identification", + "passed": comparison.status == "identified", + "value": comparison.status, + "required": "identified", + }, + ] + + paths = { + FOLD_OUTPUT: output / FOLD_OUTPUT, + CONTRAST_OUTPUT: output / CONTRAST_OUTPUT, + RAT_SCORE_OUTPUT: output / RAT_SCORE_OUTPUT, + RECOVERY_OUTPUT: output / RECOVERY_OUTPUT, + EXCLUSION_OUTPUT: output / EXCLUSION_OUTPUT, + GATE_OUTPUT: output / GATE_OUTPUT, + } + table_rows = { + FOLD_OUTPUT: ( + fold_rows, + ( + "candidate", + "held_out_rat", + "temperature", + "mean_log_score_per_bin", + "n_events", + "n_sessions", + ), + ), + CONTRAST_OUTPUT: ( + contrast_rows, + ( + "alternative", + "mean_margin", + "simultaneous_lower_bound", + "exact_one_sided_sign_flip_p", + ), + ), + RAT_SCORE_OUTPUT: ( + rat_score_rows, + ("rat", "candidate", "mean_log_score_per_bin"), + ), + RECOVERY_OUTPUT: ( + recovery_rows, + ( + "recovery_kind", + "generator", + "split_unit", + "selected_candidate", + "selected_margin", + "selected_margin_lower", + "decisive", + "n_held_out_groups", + "spatial_sigma_multiplier", + ), + ), + EXCLUSION_OUTPUT: ( + exclusion_rows, + ("event_id", "reason"), + ), + GATE_OUTPUT: ( + gate_rows, + ("gate", "passed", "value", "required"), + ), + } + for name, (rows, fieldnames) in table_rows.items(): + _write_rows(paths[name], rows, fieldnames) + + report = [ + "# PF replay spatial revision analysis", + "", + f"Status: **{comparison.status}**.", + ( + "This is a conditional replay-content comparison, not a causal " + "event-incidence analysis." + ), + ( + "Recovery is post-decoder Gaussian emission-score recovery at the " + "empirical RUN point spread, not end-to-end spike-decoder recovery." + ), + "", + f"- Predictor events: {frozen.dataset.n_events}", + f"- Complete-case recovery events: {recovery.common_event_count}", + f"- Excluded from recovery: {len(recovery.excluded_event_ids)}", + f"- Independent rats scored: {len(comparison.rat_ids)}", + f"- Descriptive winner: {comparison.winner}", + f"- Recovery gate: {'PASS' if recovery.passed else 'FAIL'}", + f"- Maximum field correlation: {comparison.maximum_field_correlation:.6g}", + ( + "- Abstention reasons: " + + (", ".join(comparison.abstention_reasons) or "none") + ), + "", + "Prespecified smoothing-revision contrasts:", + ] + report.extend( + ( + f"- versus {record.alternative}: margin {record.mean_margin:+.6g}, " + f"simultaneous lower {record.simultaneous_lower_bound:+.6g}, " + f"exact one-sided sign-flip p={record.exact_one_sided_sign_flip_p:.6g}" + ) + for record in comparison.target_contrasts + ) + report_path = output / REPORT_OUTPUT + report_path.write_text("\n".join(report) + "\n", encoding="utf-8") + paths[REPORT_OUTPUT] = report_path + + output_sha256 = {name: _sha256(path) for name, path in paths.items()} + input_manifest_path = predictor_source / "replay_spatial_manifest.json" + analysis_manifest = { + "schema_version": ANALYSIS_SCHEMA, + "created_at_utc": datetime.now(timezone.utc).isoformat(), + "consumer_repository": "IPS-Stuttgart/Bayesian-ACh", + "consumer_commit": consumer_commit, + "consumer_clean_worktree": True, + "predictor_sha256": frozen.predictor_sha256, + "predictor_manifest_sha256": _sha256(input_manifest_path), + "predictor_producer_commit": frozen.manifest.producer_commit, + "dataset_sha256": frozen.manifest.dataset_sha256, + "dataset_verifier_report_sha256": ( + frozen.manifest.dataset_verifier_report_sha256 + ), + "route_manifest_file_sha256": frozen.manifest.route_manifest_file_sha256, + "source_event_count": frozen.dataset.n_events, + "common_event_count": recovery.common_event_count, + "excluded_event_ids": list(recovery.excluded_event_ids), + "comparison_config": asdict(comparison_config), + "injection_config": asdict(injection_config), + "recovery_scope": "post_decoder_gaussian_raw_emission_scoring", + "event_selection_scope": frozen.manifest.event_selection_time_scope, + "claim_scope": "conditional_replay_content_not_event_incidence", + "recovery_passed": recovery.passed, + "status": comparison.status, + "abstention_reasons": list(comparison.abstention_reasons), + "outputs_sha256": output_sha256, + } + manifest_path = output / MANIFEST_OUTPUT + manifest_path.write_text( + json.dumps(analysis_manifest, indent=2, sort_keys=True) + "\n", + encoding="utf-8", + ) + paths[MANIFEST_OUTPUT] = manifest_path + return paths + + +def parse_args(argv: Sequence[str] | None = None) -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--predictor-dir", required=True) + parser.add_argument("--output-dir", required=True) + return parser.parse_args(argv) + + +def main(argv: Sequence[str] | None = None) -> int: + args = parse_args(argv) + run_analysis(args.predictor_dir, args.output_dir) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/src/bayesian_ach/replay_artifact.py b/src/bayesian_ach/replay_artifact.py new file mode 100644 index 0000000..65bba36 --- /dev/null +++ b/src/bayesian_ach/replay_artifact.py @@ -0,0 +1,607 @@ +"""Versioned predictor/outcome artifact contract for real replay analyses.""" + +from __future__ import annotations + +import csv +import hashlib +import json +import re +from dataclasses import asdict, dataclass +from pathlib import Path +from typing import Final + +import numpy as np +from numpy.typing import NDArray + +from bayesian_ach.replay_spatial import ( + REPLAY_SPATIAL_SCHEMA_VERSION, + SpatialReplayDataset, +) + +LATER_OUTCOME_SCHEMA_VERSION: Final[str] = "bayesian-ach.replay-later-outcome.v1" +HIPPO_TRACE_SCHEMA_VERSION: Final[str] = ( + "hipporeplayimm.first-order-smoothing-trace.v1" +) +HIPPO_TRANSITION_CONVENTION: Final[str] = ( + "column-stochastic: transition[destination, source] = " + "P(x_t=destination | x_(t-1)=source)" +) +_SHA256_PATTERN = re.compile(r"^[0-9a-f]{64}$") +_COMMIT_PATTERN = re.compile(r"^[0-9a-f]{40}$") +PF_FROZEN_DECODER_PARAMETERS_SHA256: Final[str] = ( + "a79fa8a1f55a964c4367853cc120efc9b742ec4e327c277ede78cfd6a277f20b" +) + + +def _sha256(path: Path) -> str: + digest = hashlib.sha256() + with path.open("rb") as handle: + for chunk in iter(lambda: handle.read(1024 * 1024), b""): + digest.update(chunk) + return digest.hexdigest() + + +@dataclass(frozen=True, slots=True) +class ReplaySpatialManifest: + """Provenance that can be frozen before any later outcome is joined.""" + + producer_repository: str + producer_commit: str + dataset_id: str + dataset_sha256: str + dataset_manifest_file_sha256: str + dataset_verifier_report_file: str + dataset_verifier_report_sha256: str + dataset_verified_file_count: int + dataset_verified_total_bytes: int + dataset_verified_session_count: int + dataset_verified_file_records_sha256: str + route_manifest_file: str + route_manifest_file_sha256: str + route_producer_commit: str + route_producer_clean_worktree: bool + route_parameters_sha256: str + route_segments_sha256: str + route_points_sha256: str + cohort_sha256: str + event_audit_file: str + event_audit_sha256: str + event_selection_parameters_sha256: str + behavior_field_parameters_sha256: str + decoder_parameters_sha256: str + trace_schema_version: str = HIPPO_TRACE_SCHEMA_VERSION + transition_convention: str = HIPPO_TRANSITION_CONVENTION + candidate_evidence_cutoff: str = "strict_pre_replay" + likelihood_domain: str = "max_shifted_log_emission_plus_offset" + decoder_training_schedule: str = "event_specific_prefix_refit" + decoder_point_spread_schedule: str = "pre_event_temporal_holdout_run_68pct" + event_selection_schedule: str = "lfp_raw_peak_power_top_n_per_session" + event_selection_time_scope: str = "full_session_offline_rank" + dataset_verification_schedule: str = ( + "locked_full_tree_path_size_sha256_no_extra_files" + ) + route_smoothing_scope: str = "within_completed_fill_interval" + spatial_coordinate_units: str = "cm" + well_mass_source: str = "raw_log_emission_posterior" + behavior_latent_state: str = "compact_destination_well" + behavior_observation_schedule: str = ( + "tracked_position_and_completed_fill_intervals_pre_replay" + ) + state_to_spatial_mapping: str = "pre_replay_route_kernel" + replay_feedback_used: bool = False + outcomes_in_predictor: bool = False + producer_clean_worktree: bool = True + schema_version: str = REPLAY_SPATIAL_SCHEMA_VERSION + + def validate(self) -> None: + if self.schema_version != REPLAY_SPATIAL_SCHEMA_VERSION: + raise ValueError("unsupported replay spatial schema version") + if self.trace_schema_version != HIPPO_TRACE_SCHEMA_VERSION: + raise ValueError("unsupported HippoReplayDynamics trace schema") + if self.transition_convention != HIPPO_TRANSITION_CONVENTION: + raise ValueError("transition convention does not match the trace contract") + if not self.producer_repository or not self.dataset_id: + raise ValueError("producer_repository and dataset_id are required") + if _COMMIT_PATTERN.fullmatch(self.producer_commit) is None: + raise ValueError("producer_commit must be a lowercase 40-character commit SHA") + if _SHA256_PATTERN.fullmatch(self.dataset_sha256) is None: + raise ValueError("dataset_sha256 must be a lowercase SHA-256 digest") + for value, name in ( + ( + self.dataset_manifest_file_sha256, + "dataset_manifest_file_sha256", + ), + ( + self.dataset_verifier_report_sha256, + "dataset_verifier_report_sha256", + ), + ( + self.dataset_verified_file_records_sha256, + "dataset_verified_file_records_sha256", + ), + ( + self.route_manifest_file_sha256, + "route_manifest_file_sha256", + ), + (self.route_parameters_sha256, "route_parameters_sha256"), + (self.route_segments_sha256, "route_segments_sha256"), + (self.route_points_sha256, "route_points_sha256"), + (self.cohort_sha256, "cohort_sha256"), + (self.event_audit_sha256, "event_audit_sha256"), + ( + self.event_selection_parameters_sha256, + "event_selection_parameters_sha256", + ), + ( + self.behavior_field_parameters_sha256, + "behavior_field_parameters_sha256", + ), + (self.decoder_parameters_sha256, "decoder_parameters_sha256"), + ): + if _SHA256_PATTERN.fullmatch(value) is None: + raise ValueError(f"{name} must be a lowercase SHA-256 digest") + if ( + self.decoder_parameters_sha256 + != PF_FROZEN_DECODER_PARAMETERS_SHA256 + ): + raise ValueError( + "decoder_parameters_sha256 does not match the frozen " + "120-second point-spread configuration" + ) + if _COMMIT_PATTERN.fullmatch(self.route_producer_commit) is None: + raise ValueError( + "route_producer_commit must be a lowercase 40-character commit SHA" + ) + if not self.route_producer_clean_worktree: + raise ValueError("route producer must run from a clean committed worktree") + if self.dataset_verifier_report_file != ( + "replay_spatial_dataset_verification.json" + ): + raise ValueError("dataset verifier report file is not the frozen name") + if self.route_manifest_file != "replay_spatial_route_manifest.json": + raise ValueError("route manifest file is not the frozen name") + if self.event_audit_file != "replay_spatial_event_audit.csv": + raise ValueError("event audit file is not the frozen name") + if ( + self.dataset_verified_file_count < 1 + or self.dataset_verified_total_bytes < 1 + or self.dataset_verified_session_count < 1 + ): + raise ValueError("dataset verification counts must be positive") + report = { + "schema_version": "hipporeplayimm.pf-dataset-verification.v1", + "status": "pass", + "dataset_sha256": self.dataset_sha256, + "dataset_manifest_file_sha256": self.dataset_manifest_file_sha256, + "verified_file_count": self.dataset_verified_file_count, + "verified_total_bytes": self.dataset_verified_total_bytes, + "verified_session_count": self.dataset_verified_session_count, + "verified_file_records_sha256": ( + self.dataset_verified_file_records_sha256 + ), + "missing_files": [], + "extra_files": [], + } + report_sha256 = hashlib.sha256( + (json.dumps(report, indent=2, sort_keys=True) + "\n").encode("utf-8") + ).hexdigest() + if report_sha256 != self.dataset_verifier_report_sha256: + raise ValueError("dataset verifier report digest does not match its content") + if self.candidate_evidence_cutoff != "strict_pre_replay": + raise ValueError("candidate evidence must be frozen strictly before replay") + if self.likelihood_domain != "max_shifted_log_emission_plus_offset": + raise ValueError("raw replay scores require shifted log emissions and offsets") + if self.decoder_training_schedule != "event_specific_prefix_refit": + raise ValueError("decoder training must use an event-specific prefix refit") + if ( + self.decoder_point_spread_schedule + != "pre_event_temporal_holdout_run_68pct" + ): + raise ValueError("decoder point spread must use the frozen prefix holdout") + if self.event_selection_schedule != "lfp_raw_peak_power_top_n_per_session": + raise ValueError("event selection must use raw LFP power only") + if self.event_selection_time_scope != "full_session_offline_rank": + raise ValueError("event selection time scope must be the frozen offline rank") + if self.dataset_verification_schedule != ( + "locked_full_tree_path_size_sha256_no_extra_files" + ): + raise ValueError("dataset verification must check the entire locked tree") + if self.route_smoothing_scope != "within_completed_fill_interval": + raise ValueError("route smoothing must not cross completed-route boundaries") + if self.spatial_coordinate_units != "cm": + raise ValueError("spatial coordinates and point spread must use cm") + if self.well_mass_source != "raw_log_emission_posterior": + raise ValueError("well masses must be derived from the raw posterior") + if self.behavior_latent_state != "compact_destination_well": + raise ValueError("behavioral smoothing must use the compact well state") + if ( + self.behavior_observation_schedule + != "tracked_position_and_completed_fill_intervals_pre_replay" + ): + raise ValueError("behavior observation schedule is not the frozen schedule") + if self.state_to_spatial_mapping != "pre_replay_route_kernel": + raise ValueError("state-to-spatial mapping is not the frozen mapping") + if self.replay_feedback_used: + raise ValueError("decoded replay must not be fed back as a new observation") + if self.outcomes_in_predictor: + raise ValueError("later outcomes must not be present in the predictor artifact") + if not self.producer_clean_worktree: + raise ValueError("producer must run from a clean committed worktree") + + +def _expected_dataset_verifier_report( + manifest: ReplaySpatialManifest, +) -> dict[str, object]: + return { + "schema_version": "hipporeplayimm.pf-dataset-verification.v1", + "status": "pass", + "dataset_sha256": manifest.dataset_sha256, + "dataset_manifest_file_sha256": manifest.dataset_manifest_file_sha256, + "verified_file_count": manifest.dataset_verified_file_count, + "verified_total_bytes": manifest.dataset_verified_total_bytes, + "verified_session_count": manifest.dataset_verified_session_count, + "verified_file_records_sha256": ( + manifest.dataset_verified_file_records_sha256 + ), + "missing_files": [], + "extra_files": [], + } + + +def _hash_checked_json( + path: Path, + expected_sha256: str, + *, + label: str, +) -> dict[str, object]: + if not path.is_file() or _sha256(path) != expected_sha256: + raise ValueError(f"{label} SHA-256 does not match its manifest") + payload = json.loads(path.read_text(encoding="utf-8")) + if not isinstance(payload, dict): + raise ValueError(f"{label} must contain a JSON object") + return payload + + +def _validate_provenance_sidecars( + source: Path, + dataset: SpatialReplayDataset, + manifest: ReplaySpatialManifest, +) -> None: + report = _hash_checked_json( + source / manifest.dataset_verifier_report_file, + manifest.dataset_verifier_report_sha256, + label="dataset verifier report", + ) + if report != _expected_dataset_verifier_report(manifest): + raise ValueError("dataset verifier report content disagrees with manifest") + + route = _hash_checked_json( + source / manifest.route_manifest_file, + manifest.route_manifest_file_sha256, + label="route provenance manifest", + ) + if ( + route.get("analysis") != "replay_behavior_route_primitives" + or route.get("producer_commit") != manifest.route_producer_commit + or route.get("producer_clean_worktree") is not True + or route.get("route_smoothing_scope") != manifest.route_smoothing_scope + or route.get("parameters_sha256") != manifest.route_parameters_sha256 + ): + raise ValueError("route provenance content disagrees with manifest") + output_sha256 = route.get("output_sha256") + if not isinstance(output_sha256, dict): + raise ValueError("route provenance output hashes are missing") + if ( + output_sha256.get("replay_behavior_route_segments.csv") + != manifest.route_segments_sha256 + or output_sha256.get("replay_behavior_route_segment_points.csv") + != manifest.route_points_sha256 + ): + raise ValueError("route table hashes disagree with route provenance") + + audit_path = source / manifest.event_audit_file + if not audit_path.is_file() or _sha256(audit_path) != manifest.event_audit_sha256: + raise ValueError("event audit SHA-256 does not match its manifest") + with audit_path.open("r", encoding="utf-8", newline="") as handle: + rows = list(csv.DictReader(handle)) + if len(rows) != dataset.n_events: + raise ValueError("event audit row count does not match predictor events") + event_ids = tuple(str(row.get("event_id", "")) for row in rows) + sessions = tuple(str(row.get("session", "")) for row in rows) + rats = tuple(str(row.get("rat", "")) for row in rows) + if ( + event_ids != dataset.event_ids + or sessions != tuple(np.asarray(dataset.session_ids, dtype=str)) + or rats != tuple(np.asarray(dataset.rat_ids, dtype=str)) + ): + raise ValueError("event audit identifiers do not match predictor arrays") + try: + cohort = [ + { + "event_id": event_id, + "session": session, + "event_index": int(row["event_index"]), + } + for event_id, session, row in zip( + event_ids, + sessions, + rows, + strict=True, + ) + ] + except (KeyError, TypeError, ValueError) as error: + raise ValueError("event audit event_index values are invalid") from error + cohort_sha256 = hashlib.sha256( + ( + json.dumps(cohort, sort_keys=True, separators=(",", ":")) + + "\n" + ).encode("utf-8") + ).hexdigest() + if cohort_sha256 != manifest.cohort_sha256: + raise ValueError("event audit cohort digest does not match its manifest") + + +@dataclass(frozen=True, slots=True) +class FrozenPredictorArtifact: + dataset: SpatialReplayDataset + manifest: ReplaySpatialManifest + predictor_sha256: str + + +def write_spatial_predictor_artifact( + directory: Path | str, + dataset: SpatialReplayDataset, + manifest: ReplaySpatialManifest, +) -> FrozenPredictorArtifact: + """Write a predictor-only NPZ and a hash-binding JSON manifest.""" + + dataset.validate() + manifest.validate() + output = Path(directory) + output.mkdir(parents=True, exist_ok=True) + _validate_provenance_sidecars(output, dataset, manifest) + predictor_path = output / "replay_spatial_predictors.npz" + well_masses = ( + np.empty((dataset.n_events, 0), dtype=float) + if dataset.well_masses is None + else np.asarray(dataset.well_masses, dtype=float) + ) + np.savez_compressed( + predictor_path, + event_ids=np.asarray(dataset.event_ids, dtype=str), + rat_ids=np.asarray(dataset.rat_ids, dtype=str), + session_ids=np.asarray(dataset.session_ids, dtype=str), + event_start_s=np.asarray(dataset.event_start_s, dtype=float), + event_end_s=np.asarray(dataset.event_end_s, dtype=float), + history_cutoff_s=np.asarray(dataset.history_cutoff_s, dtype=float), + decoder_training_cutoff_s=np.asarray( + dataset.decoder_training_cutoff_s, + dtype=float, + ), + field_available_s=np.asarray(dataset.field_available_s, dtype=float), + log_emissions=np.asarray(dataset.log_emissions, dtype=float), + log_emission_offsets=np.asarray(dataset.log_emission_offsets, dtype=float), + time_mask=np.asarray(dataset.time_mask, dtype=bool), + active_spatial_mask=np.asarray(dataset.active_spatial_mask, dtype=bool), + spatial_coordinates=np.asarray(dataset.spatial_coordinates, dtype=float), + decoder_point_spread_cm=np.asarray( + dataset.decoder_point_spread_cm, + dtype=float, + ), + nuisance_base=np.asarray(dataset.nuisance_base, dtype=float), + candidate_fields=np.asarray(dataset.candidate_fields, dtype=float), + candidate_available=np.asarray(dataset.candidate_available, dtype=bool), + candidate_names=np.asarray(dataset.candidate_names, dtype=str), + well_masses=well_masses, + well_ids=np.asarray(dataset.well_ids, dtype=str), + ) + predictor_sha256 = _sha256(predictor_path) + manifest_payload = { + **asdict(manifest), + "predictor_file": predictor_path.name, + "predictor_sha256": predictor_sha256, + } + manifest_path = output / "replay_spatial_manifest.json" + manifest_path.write_text( + json.dumps(manifest_payload, indent=2, sort_keys=True) + "\n", + encoding="utf-8", + ) + return FrozenPredictorArtifact(dataset, manifest, predictor_sha256) + + +def load_spatial_predictor_artifact( + directory: Path | str, +) -> FrozenPredictorArtifact: + """Load, hash-check, and validate a predictor-only artifact.""" + + source = Path(directory) + manifest_payload = json.loads( + (source / "replay_spatial_manifest.json").read_text(encoding="utf-8") + ) + predictor_name = str(manifest_payload.pop("predictor_file")) + expected_sha256 = str(manifest_payload.pop("predictor_sha256")) + if _SHA256_PATTERN.fullmatch(expected_sha256) is None: + raise ValueError("predictor_sha256 is not a valid SHA-256 digest") + predictor_path = source / predictor_name + observed_sha256 = _sha256(predictor_path) + if observed_sha256 != expected_sha256: + raise ValueError("predictor artifact SHA-256 does not match its manifest") + manifest = ReplaySpatialManifest(**manifest_payload) + manifest.validate() + + with np.load(predictor_path, allow_pickle=False) as arrays: + well_mass = np.asarray(arrays["well_masses"], dtype=float) + dataset = SpatialReplayDataset( + event_ids=tuple(str(value) for value in arrays["event_ids"]), + rat_ids=np.asarray(arrays["rat_ids"], dtype=str), + session_ids=np.asarray(arrays["session_ids"], dtype=str), + event_start_s=np.asarray(arrays["event_start_s"], dtype=float), + event_end_s=np.asarray(arrays["event_end_s"], dtype=float), + history_cutoff_s=np.asarray(arrays["history_cutoff_s"], dtype=float), + decoder_training_cutoff_s=np.asarray( + arrays["decoder_training_cutoff_s"], + dtype=float, + ), + field_available_s=np.asarray(arrays["field_available_s"], dtype=float), + log_emissions=np.asarray(arrays["log_emissions"], dtype=float), + log_emission_offsets=np.asarray( + arrays["log_emission_offsets"], + dtype=float, + ), + time_mask=np.asarray(arrays["time_mask"], dtype=bool), + active_spatial_mask=np.asarray( + arrays["active_spatial_mask"], + dtype=bool, + ), + spatial_coordinates=np.asarray( + arrays["spatial_coordinates"], + dtype=float, + ), + decoder_point_spread_cm=np.asarray( + arrays["decoder_point_spread_cm"], + dtype=float, + ), + nuisance_base=np.asarray(arrays["nuisance_base"], dtype=float), + candidate_fields=np.asarray(arrays["candidate_fields"], dtype=float), + candidate_available=np.asarray( + arrays["candidate_available"], + dtype=bool, + ), + candidate_names=tuple(str(value) for value in arrays["candidate_names"]), + well_masses=None if well_mass.shape[1] == 0 else well_mass, + well_ids=tuple(str(value) for value in arrays["well_ids"]), + ) + dataset.validate() + _validate_provenance_sidecars(source, dataset, manifest) + return FrozenPredictorArtifact(dataset, manifest, observed_sha256) + + +@dataclass(frozen=True, slots=True) +class LaterOutcomeTable: + """Behavior observed after replay and stored outside the predictor artifact.""" + + event_ids: tuple[str, ...] + outcome_time_s: NDArray[np.float64] + next_well_ids: tuple[str, ...] + valid: NDArray[np.bool_] + + def validate(self, predictors: SpatialReplayDataset) -> None: + n_rows = len(self.event_ids) + if n_rows < 1 or len(set(self.event_ids)) != n_rows: + raise ValueError("outcome event_ids must be nonempty and unique") + times = np.asarray(self.outcome_time_s, dtype=float) + valid = np.asarray(self.valid) + if times.shape != (n_rows,) or not np.all(np.isfinite(times)): + raise ValueError("outcome_time_s must contain one finite value per row") + if valid.dtype != np.bool_ or valid.shape != (n_rows,): + raise ValueError("outcome valid flags must be boolean") + if len(self.next_well_ids) != n_rows: + raise ValueError("next_well_ids must contain one value per outcome") + event_lookup = { + event_id: index for index, event_id in enumerate(predictors.event_ids) + } + for row, event_id in enumerate(self.event_ids): + if event_id not in event_lookup: + raise ValueError(f"outcome references unknown event {event_id!r}") + predictor_index = event_lookup[event_id] + if times[row] <= predictors.event_end_s[predictor_index]: + raise ValueError("later outcomes must occur strictly after replay ends") + if valid[row]: + if not self.next_well_ids[row]: + raise ValueError("valid outcomes require a next well") + if ( + predictors.well_ids + and self.next_well_ids[row] not in predictors.well_ids + ): + raise ValueError("valid outcome well is absent from predictor support") + + +@dataclass(frozen=True, slots=True) +class FrozenOutcomeArtifact: + outcomes: LaterOutcomeTable + predictor_sha256: str + outcome_sha256: str + + +def write_later_outcome_artifact( + directory: Path | str, + predictors: FrozenPredictorArtifact, + outcomes: LaterOutcomeTable, +) -> FrozenOutcomeArtifact: + """Write outcomes separately and bind them to an already frozen predictor.""" + + outcomes.validate(predictors.dataset) + output = Path(directory) + output.mkdir(parents=True, exist_ok=True) + outcome_path = output / "replay_later_outcomes.npz" + np.savez_compressed( + outcome_path, + event_ids=np.asarray(outcomes.event_ids, dtype=str), + outcome_time_s=np.asarray(outcomes.outcome_time_s, dtype=float), + next_well_ids=np.asarray(outcomes.next_well_ids, dtype=str), + valid=np.asarray(outcomes.valid, dtype=bool), + ) + outcome_sha256 = _sha256(outcome_path) + payload = { + "schema_version": LATER_OUTCOME_SCHEMA_VERSION, + "outcome_file": outcome_path.name, + "outcome_sha256": outcome_sha256, + "predictor_sha256": predictors.predictor_sha256, + } + (output / "replay_later_outcome_manifest.json").write_text( + json.dumps(payload, indent=2, sort_keys=True) + "\n", + encoding="utf-8", + ) + return FrozenOutcomeArtifact( + outcomes, + predictors.predictor_sha256, + outcome_sha256, + ) + + +def load_later_outcome_artifact( + directory: Path | str, + predictors: FrozenPredictorArtifact, +) -> FrozenOutcomeArtifact: + """Load outcomes only after verifying their frozen predictor binding.""" + + source = Path(directory) + payload = json.loads( + (source / "replay_later_outcome_manifest.json").read_text(encoding="utf-8") + ) + if payload.get("schema_version") != LATER_OUTCOME_SCHEMA_VERSION: + raise ValueError("unsupported later-outcome schema version") + if payload.get("predictor_sha256") != predictors.predictor_sha256: + raise ValueError("outcomes are not bound to this predictor artifact") + outcome_path = source / str(payload["outcome_file"]) + observed_sha256 = _sha256(outcome_path) + if observed_sha256 != payload.get("outcome_sha256"): + raise ValueError("outcome artifact SHA-256 does not match its manifest") + with np.load(outcome_path, allow_pickle=False) as arrays: + outcomes = LaterOutcomeTable( + event_ids=tuple(str(value) for value in arrays["event_ids"]), + outcome_time_s=np.asarray(arrays["outcome_time_s"], dtype=float), + next_well_ids=tuple(str(value) for value in arrays["next_well_ids"]), + valid=np.asarray(arrays["valid"], dtype=bool), + ) + outcomes.validate(predictors.dataset) + return FrozenOutcomeArtifact( + outcomes, + predictors.predictor_sha256, + observed_sha256, + ) + + +__all__ = [ + "FrozenOutcomeArtifact", + "FrozenPredictorArtifact", + "HIPPO_TRACE_SCHEMA_VERSION", + "HIPPO_TRANSITION_CONVENTION", + "LATER_OUTCOME_SCHEMA_VERSION", + "LaterOutcomeTable", + "ReplaySpatialManifest", + "load_later_outcome_artifact", + "load_spatial_predictor_artifact", + "write_later_outcome_artifact", + "write_spatial_predictor_artifact", +] diff --git a/src/bayesian_ach/replay_recovery.py b/src/bayesian_ach/replay_recovery.py new file mode 100644 index 0000000..2b73f06 --- /dev/null +++ b/src/bayesian_ach/replay_recovery.py @@ -0,0 +1,467 @@ +"""Emission-level injection recovery for spatial replay model comparison.""" + +from __future__ import annotations + +from dataclasses import dataclass, replace + +import numpy as np +from numpy.typing import NDArray +from scipy.special import logsumexp + +from bayesian_ach.replay_spatial import ( + SPATIAL_CANDIDATE_NAMES, + SpatialComparisonConfig, + SpatialRecoveryGate, + SpatialRecoveryRecord, + SpatialReplayComparison, + SpatialReplayDataset, + compare_spatial_replay_candidates, + spatial_common_candidate_mask, + subset_spatial_replay_dataset, +) + + +@dataclass(frozen=True, slots=True) +class SpatialInjectionRecoveryConfig: + """Frozen settings for post-decoder raw-emission scoring recovery.""" + + injection_temperature: float = 4.0 + spatial_sigma_multipliers: tuple[float, ...] = (0.5, 1.0, 2.0) + coordinate_units: str = "cm" + emission_noise_sd_nats: float = 0.02 + mixtures: tuple[tuple[str, str], ...] = ( + ("smoothing_revision", "td_error"), + ("smoothing_revision", "prospective"), + ("smoothing_revision", "recency"), + ("smoothing_revision", "posterior_content"), + ) + seed: int = 701 + + def validate(self) -> None: + if not np.isfinite(self.injection_temperature) or self.injection_temperature <= 0: + raise ValueError("injection_temperature must be finite and positive") + multipliers = np.asarray(self.spatial_sigma_multipliers, dtype=float) + if ( + multipliers.ndim != 1 + or multipliers.size < 1 + or not np.all(np.isfinite(multipliers)) + or np.any(multipliers <= 0.0) + or len(set(float(value) for value in multipliers)) != multipliers.size + ): + raise ValueError( + "spatial_sigma_multipliers must be unique finite positive values" + ) + if self.coordinate_units != "cm": + raise ValueError("spatial coordinates and decoder point spread must use cm") + if ( + not np.isfinite(self.emission_noise_sd_nats) + or self.emission_noise_sd_nats < 0 + ): + raise ValueError("emission_noise_sd_nats must be finite and nonnegative") + for first, second in self.mixtures: + if ( + first not in SPATIAL_CANDIDATE_NAMES + or second not in SPATIAL_CANDIDATE_NAMES + or first == second + ): + raise ValueError("mixtures must contain two distinct registered candidates") + + +def _standardized_fields( + dataset: SpatialReplayDataset, +) -> tuple[NDArray[np.float64], NDArray[np.float64]]: + base = np.asarray(dataset.nuisance_base, dtype=float).copy() + base /= base.sum(axis=1, keepdims=True) + fields = np.asarray(dataset.candidate_fields, dtype=float) + mean = np.sum(base[:, None, :] * fields, axis=2, keepdims=True) + centered = fields - mean + variance = np.sum(base[:, None, :] * centered**2, axis=2) + if np.any(variance <= 1e-14): + raise ValueError("recovery requires every candidate field to be nonconstant") + standardized = centered / np.sqrt(variance)[:, :, None] + support = np.broadcast_to(dataset.active_spatial_mask[:, None, :], fields.shape) + standardized[~support] = 0.0 + return base, standardized + + +def inject_spatial_replay_emissions( + dataset: SpatialReplayDataset, + generators: tuple[str, ...], + config: SpatialInjectionRecoveryConfig | None = None, + *, + spatial_sigma_multiplier: float = 1.0, + seed: int | None = None, +) -> SpatialReplayDataset: + """Replace replay emissions with draws from one candidate or a 50/50 mixture. + + Candidate priors generate latent grid bins. A Gaussian likelihood around the + sampled bin is then written in the same max-shifted raw-log-emission domain + as the real decoder. The candidate fields themselves are not altered. + """ + + dataset.validate() + config = SpatialInjectionRecoveryConfig() if config is None else config + config.validate() + if len(generators) not in (0, 1, 2): + raise ValueError( + "generators must be empty for the nuisance-base null, " + "one candidate, or a 50/50 pair" + ) + if any(name not in dataset.candidate_names for name in generators): + raise ValueError("all generators must be registered candidates") + multiplier = float(spatial_sigma_multiplier) + if not np.isfinite(multiplier) or multiplier <= 0.0: + raise ValueError("spatial_sigma_multiplier must be finite and positive") + generator_indices = [dataset.candidate_names.index(name) for name in generators] + for index in generator_indices: + if not np.all(dataset.candidate_available[:, index]): + raise ValueError("recovery generator is unavailable for at least one event") + + base, standardized = _standardized_fields(dataset) + rng = np.random.default_rng(config.seed if seed is None else seed) + emissions = np.full_like(dataset.log_emissions, -np.inf, dtype=float) + for event_index in range(dataset.n_events): + active = np.asarray(dataset.active_spatial_mask[event_index], dtype=bool) + active_indices = np.flatnonzero(active) + coordinates = dataset.spatial_coordinates[event_index, active] + for time_index in np.flatnonzero(dataset.time_mask[event_index]): + candidate_index: int | None + if len(generator_indices) == 0: + candidate_index = None + elif len(generator_indices) == 1: + candidate_index = generator_indices[0] + else: + generator_position = (event_index + int(time_index)) % 2 + candidate_index = generator_indices[generator_position] + log_prior = np.log(base[event_index, active]) + if candidate_index is not None: + log_prior += ( + config.injection_temperature + * standardized[event_index, candidate_index, active] + ) + probabilities = np.exp(log_prior - logsumexp(log_prior)) + sampled_position = int(rng.choice(len(active_indices), p=probabilities)) + squared_distance = np.sum( + (coordinates - coordinates[sampled_position]) ** 2, + axis=1, + ) + spatial_sigma_cm = ( + dataset.decoder_point_spread_cm[event_index] * multiplier + ) + row = -0.5 * squared_distance / spatial_sigma_cm**2 + if config.emission_noise_sd_nats > 0: + row += rng.normal( + 0.0, + config.emission_noise_sd_nats, + size=row.shape, + ) + row -= np.max(row) + emissions[event_index, time_index, active] = row + + return replace( + dataset, + log_emissions=np.asarray(emissions, dtype=np.float64), + log_emission_offsets=np.zeros_like(dataset.log_emission_offsets, dtype=float), + ) + + +def _comparison_for_split( + dataset: SpatialReplayDataset, + comparison_config: SpatialComparisonConfig, + split_unit: str, +) -> SpatialReplayComparison: + if split_unit == "leave_one_rat_out": + grouped = dataset + elif split_unit == "leave_one_session_out": + compound = np.asarray( + [ + f"{rat}|{session}" + for rat, session in zip( + np.asarray(dataset.rat_ids, dtype=str), + np.asarray(dataset.session_ids, dtype=str), + strict=True, + ) + ], + dtype=str, + ) + grouped = replace(dataset, rat_ids=compound, session_ids=compound) + else: + raise ValueError("split_unit must be leave_one_rat_out or leave_one_session_out") + return compare_spatial_replay_candidates(grouped, comparison_config) + + +def _fixed_candidate_contrast( + result: SpatialReplayComparison, + candidate: str, + comparison_config: SpatialComparisonConfig, + *, + seed: int, +) -> tuple[float, float, bool]: + candidate_index = result.candidate_names.index(candidate) + alternatives = [ + index + for index in range(len(result.candidate_names)) + if index != candidate_index + ] + paired = ( + result.rat_scores[:, [candidate_index]] + - result.rat_scores[:, alternatives] + ) + if paired.shape[0] < 2: + return float("nan"), float("nan"), False + observed = paired.mean(axis=0) + rng = np.random.default_rng(seed) + indices = rng.integers( + 0, + paired.shape[0], + size=(comparison_config.bootstrap_replicates, paired.shape[0]), + ) + bootstrap = paired[indices].mean(axis=1) + shortfall = np.max(observed[None, :] - bootstrap, axis=1) + critical = float( + np.quantile(shortfall, comparison_config.simultaneous_confidence_level) + ) + lower = observed - critical + return float(np.min(observed)), float(np.min(lower)), bool(np.all(lower > 0.0)) + + +def run_spatial_recovery_checks( + dataset: SpatialReplayDataset, + comparison_config: SpatialComparisonConfig | None = None, + injection_config: SpatialInjectionRecoveryConfig | None = None, +) -> SpatialRecoveryGate: + """Compute post-decoder LOAO/LOSO scoring recovery on the common cohort. + + This tests Gaussian raw-emission score discrimination at the empirical + decoder point spread. It is not end-to-end spike/place-field decoder + recovery. Original event IDs and exclusions are retained in the gate. + """ + + dataset.validate() + source_event_ids = tuple(dataset.event_ids) + common_mask = spatial_common_candidate_mask(dataset) + excluded_event_ids = tuple( + event_id + for event_id, keep in zip(source_event_ids, common_mask, strict=True) + if not bool(keep) + ) + if not np.any(common_mask): + comparison_config = ( + SpatialComparisonConfig() + if comparison_config is None + else comparison_config + ) + comparison_config.validate() + injection_config = ( + SpatialInjectionRecoveryConfig() + if injection_config is None + else injection_config + ) + injection_config.validate() + return SpatialRecoveryGate( + pure_records=(), + mixture_records=(), + null_records=(), + source_event_count=len(source_event_ids), + common_event_count=0, + excluded_event_ids=excluded_event_ids, + required_mixtures=tuple( + "+".join(mixture) for mixture in injection_config.mixtures + ), + required_sigma_multipliers=tuple( + float(value) + for value in injection_config.spatial_sigma_multipliers + ), + ) + dataset = subset_spatial_replay_dataset(dataset, common_mask) + comparison_config = ( + SpatialComparisonConfig() + if comparison_config is None + else comparison_config + ) + comparison_config.validate() + injection_config = ( + SpatialInjectionRecoveryConfig() + if injection_config is None + else injection_config + ) + injection_config.validate() + + pure_records: list[SpatialRecoveryRecord] = [] + mixture_records: list[SpatialRecoveryRecord] = [] + null_records: list[SpatialRecoveryRecord] = [] + split_units = ("leave_one_rat_out", "leave_one_session_out") + for generator_index, generator in enumerate(dataset.candidate_names): + for multiplier_index, multiplier in enumerate( + injection_config.spatial_sigma_multipliers + ): + injected = inject_spatial_replay_emissions( + dataset, + (generator,), + injection_config, + spatial_sigma_multiplier=float(multiplier), + seed=( + injection_config.seed + + 10_007 * generator_index + + 101 * multiplier_index + ), + ) + for split_index, split_unit in enumerate(split_units): + result = _comparison_for_split( + injected, + comparison_config, + split_unit, + ) + margin, lower, decisive = _fixed_candidate_contrast( + result, + generator, + comparison_config, + seed=( + injection_config.seed + + 10_007 * generator_index + + 101 * multiplier_index + + split_index + ), + ) + pure_records.append( + SpatialRecoveryRecord( + generator=generator, + split_unit=split_unit, + selected_candidate=result.winner, + selected_margin=margin, + selected_margin_lower=lower, + decisive=decisive, + n_held_out_groups=len(result.rat_ids), + spatial_sigma_multiplier=float(multiplier), + ) + ) + + mixture_names: list[str] = [] + for mixture_index, mixture in enumerate(injection_config.mixtures): + mixture_name = "+".join(mixture) + mixture_names.append(mixture_name) + for multiplier_index, multiplier in enumerate( + injection_config.spatial_sigma_multipliers + ): + injected = inject_spatial_replay_emissions( + dataset, + mixture, + injection_config, + spatial_sigma_multiplier=float(multiplier), + seed=( + injection_config.seed + + 100_003 + + 10_007 * mixture_index + + 101 * multiplier_index + ), + ) + for split_index, split_unit in enumerate(split_units): + result = _comparison_for_split( + injected, + comparison_config, + split_unit, + ) + summaries = [ + ( + candidate, + *_fixed_candidate_contrast( + result, + candidate, + comparison_config, + seed=( + injection_config.seed + + 100_003 + + 10_007 * mixture_index + + 101 * multiplier_index + + split_index + + 1_000_003 * candidate_index + ), + ), + ) + for candidate_index, candidate in enumerate( + result.candidate_names + ) + ] + selected = max(summaries, key=lambda values: values[2]) + mixture_records.append( + SpatialRecoveryRecord( + generator=mixture_name, + split_unit=split_unit, + selected_candidate=selected[0], + selected_margin=selected[1], + selected_margin_lower=selected[2], + decisive=any(values[3] for values in summaries), + n_held_out_groups=len(result.rat_ids), + spatial_sigma_multiplier=float(multiplier), + ) + ) + + for multiplier_index, multiplier in enumerate( + injection_config.spatial_sigma_multipliers + ): + injected = inject_spatial_replay_emissions( + dataset, + (), + injection_config, + spatial_sigma_multiplier=float(multiplier), + seed=injection_config.seed + 900_001 + 101 * multiplier_index, + ) + for split_index, split_unit in enumerate(split_units): + result = _comparison_for_split( + injected, + comparison_config, + split_unit, + ) + summaries = [ + ( + candidate, + *_fixed_candidate_contrast( + result, + candidate, + comparison_config, + seed=( + injection_config.seed + + 900_001 + + 101 * multiplier_index + + split_index + + 1_000_003 * candidate_index + ), + ), + ) + for candidate_index, candidate in enumerate( + result.candidate_names[:-1] + ) + ] + selected = max(summaries, key=lambda values: values[2]) + null_records.append( + SpatialRecoveryRecord( + generator="null", + split_unit=split_unit, + selected_candidate=selected[0], + selected_margin=selected[1], + selected_margin_lower=selected[2], + decisive=any(values[3] for values in summaries), + n_held_out_groups=len(result.rat_ids), + spatial_sigma_multiplier=float(multiplier), + ) + ) + + return SpatialRecoveryGate( + pure_records=tuple(pure_records), + mixture_records=tuple(mixture_records), + null_records=tuple(null_records), + source_event_count=len(source_event_ids), + common_event_count=dataset.n_events, + excluded_event_ids=excluded_event_ids, + required_mixtures=tuple(mixture_names), + required_sigma_multipliers=tuple( + float(value) for value in injection_config.spatial_sigma_multipliers + ), + ) + + +__all__ = [ + "SpatialInjectionRecoveryConfig", + "inject_spatial_replay_emissions", + "run_spatial_recovery_checks", +] diff --git a/src/bayesian_ach/replay_spatial.py b/src/bayesian_ach/replay_spatial.py new file mode 100644 index 0000000..59137d0 --- /dev/null +++ b/src/bayesian_ach/replay_spatial.py @@ -0,0 +1,947 @@ +"""Leakage-safe spatial test of replay as filtering-to-smoothing revision.""" + +from __future__ import annotations + +from dataclasses import dataclass, replace +from typing import Any, Final + +import numpy as np +from numpy.typing import ArrayLike, NDArray +from scipy.special import logsumexp + +REPLAY_SPATIAL_SCHEMA_VERSION: Final[str] = "bayesian-ach.replay-spatial.v2" +SPATIAL_CANDIDATE_NAMES: Final[tuple[str, ...]] = ( + "smoothing_revision", + "online_surprise", + "posterior_content", + "current_location", + "recency", + "prospective", + "td_error", +) +NULL_CANDIDATE_NAME: Final[str] = "null" +_FLOAT_TOL = 1e-10 + + +def _probability_rows(values: ArrayLike, *, name: str) -> NDArray[np.float64]: + array = np.asarray(values, dtype=float) + if array.ndim != 2 or array.shape[0] < 1 or array.shape[1] < 2: + raise ValueError(f"{name} must have shape (positive snippets, at least two states)") + if not np.all(np.isfinite(array)) or np.any(array < 0.0): + raise ValueError(f"{name} must contain finite nonnegative values") + totals = array.sum(axis=1) + if np.any(totals <= 0.0): + raise ValueError(f"every {name} row must contain positive mass") + return np.asarray(array / totals[:, None], dtype=np.float64) + + +def _categorical_kl_rows( + posterior: NDArray[np.float64], + prior: NDArray[np.float64], +) -> NDArray[np.float64]: + positive = posterior > 0.0 + if np.any(positive & (prior <= 0.0)): + raise ValueError("smoothing posterior has mass outside filtering support") + terms = np.zeros_like(posterior) + terms[positive] = posterior[positive] * np.log( + posterior[positive] / prior[positive] + ) + return np.asarray(terms.sum(axis=1), dtype=np.float64) + + +@dataclass(frozen=True, slots=True) +class SignedRevisionField: + """Pre-replay spatial field assembled from historical smoothing revisions.""" + + signed_field: NDArray[np.float64] + per_snippet_kl: NDArray[np.float64] + snippet_weights: NDArray[np.float64] + total_weight: float + identifiable: bool + + +def build_signed_revision_field( + filtered_probabilities: ArrayLike, + smoothed_probabilities: ArrayLike, + snippet_end_s: ArrayLike, + *, + event_start_s: float, + state_to_spatial: ArrayLike | None = None, + recency_tau_s: float = 30.0, + minimum_total_weight: float = 1e-10, +) -> SignedRevisionField: + """Build a signed spatial revision field from strictly historical snippets. + + Each state-wise difference (smoothed minus filtered) is projected onto the + spatial grid and weighted by KL(smoothed || filtered). An exponential age + weight is optional through recency_tau_s. Every snippet must end no later + than the replay event; replay emissions and later outcomes are forbidden. + """ + + filtered = _probability_rows(filtered_probabilities, name="filtered_probabilities") + smoothed = _probability_rows(smoothed_probabilities, name="smoothed_probabilities") + if smoothed.shape != filtered.shape: + raise ValueError("filtered and smoothed probabilities must have matching shapes") + ends = np.asarray(snippet_end_s, dtype=float) + if ends.shape != (filtered.shape[0],) or not np.all(np.isfinite(ends)): + raise ValueError("snippet_end_s must contain one finite time per snippet") + event_start = float(event_start_s) + if not np.isfinite(event_start): + raise ValueError("event_start_s must be finite") + if np.any(ends > event_start + _FLOAT_TOL): + raise ValueError("all smoothing snippets must end before replay starts") + tau = float(recency_tau_s) + if not np.isfinite(tau) or tau <= 0.0: + raise ValueError("recency_tau_s must be finite and positive") + threshold = float(minimum_total_weight) + if not np.isfinite(threshold) or threshold < 0.0: + raise ValueError("minimum_total_weight must be finite and nonnegative") + + n_snippets, n_states = filtered.shape + if state_to_spatial is None: + mapping = np.broadcast_to( + np.eye(n_states, dtype=float)[None, :, :], + (n_snippets, n_states, n_states), + ).copy() + else: + raw_mapping = np.asarray(state_to_spatial, dtype=float) + if raw_mapping.ndim == 2: + if raw_mapping.shape[0] != n_states or raw_mapping.shape[1] < 2: + raise ValueError( + "state_to_spatial must have shape (state, at least two spatial bins)" + ) + mapping = np.broadcast_to( + raw_mapping[None, :, :], + (n_snippets, *raw_mapping.shape), + ).copy() + elif raw_mapping.ndim == 3: + if ( + raw_mapping.shape[0] != n_snippets + or raw_mapping.shape[1] != n_states + or raw_mapping.shape[2] < 2 + ): + raise ValueError( + "time-varying state_to_spatial must have shape " + "(snippet, state, at least two spatial bins)" + ) + mapping = raw_mapping.copy() + else: + raise ValueError("state_to_spatial must be two- or three-dimensional") + if not np.all(np.isfinite(mapping)) or np.any(mapping < 0.0): + raise ValueError("state_to_spatial must contain finite nonnegative values") + row_mass = mapping.sum(axis=2) + if np.any(row_mass <= 0.0): + raise ValueError("every state_to_spatial row must contain positive mass") + mapping /= row_mass[:, :, None] + + kl = _categorical_kl_rows(smoothed, filtered) + age = np.maximum(event_start - ends, 0.0) + weights = kl * np.exp(-age / tau) + signed_snippets = np.einsum( + "hs,hsb->hb", + smoothed - filtered, + mapping, + ) + signed_field = np.sum(weights[:, None] * signed_snippets, axis=0) + total_weight = float(weights.sum()) + return SignedRevisionField( + signed_field=np.asarray(signed_field, dtype=np.float64), + per_snippet_kl=kl, + snippet_weights=np.asarray(weights, dtype=np.float64), + total_weight=total_weight, + identifiable=bool(total_weight > threshold), + ) + + +@dataclass(frozen=True, slots=True) +class SpatialReplayDataset: + """Frozen predictor-only artifact used for spatial replay comparison. + + Log emissions are stored after a per-event/time additive offset has been + removed. log_emission_offsets restores that offset. Candidate fields and + their source cutoffs are frozen before the replay event. Later behavioral + outcomes are intentionally absent from this object. + """ + + event_ids: tuple[str, ...] + rat_ids: NDArray[np.str_] + session_ids: NDArray[np.str_] + event_start_s: NDArray[np.float64] + event_end_s: NDArray[np.float64] + history_cutoff_s: NDArray[np.float64] + decoder_training_cutoff_s: NDArray[np.float64] + field_available_s: NDArray[np.float64] + log_emissions: NDArray[np.float64] + log_emission_offsets: NDArray[np.float64] + time_mask: NDArray[np.bool_] + active_spatial_mask: NDArray[np.bool_] + spatial_coordinates: NDArray[np.float64] + decoder_point_spread_cm: NDArray[np.float64] + nuisance_base: NDArray[np.float64] + candidate_fields: NDArray[np.float64] + candidate_available: NDArray[np.bool_] + candidate_names: tuple[str, ...] = SPATIAL_CANDIDATE_NAMES + well_masses: NDArray[np.float64] | None = None + well_ids: tuple[str, ...] = () + + @property + def n_events(self) -> int: + return len(self.event_ids) + + @property + def n_time(self) -> int: + return int(self.log_emissions.shape[1]) + + @property + def n_spatial_bins(self) -> int: + return int(self.log_emissions.shape[2]) + + def validate(self) -> None: + """Reject shape, provenance, support, and future-leakage violations.""" + + n_events = self.n_events + if n_events < 1 or len(set(self.event_ids)) != n_events: + raise ValueError("event_ids must be nonempty and unique") + if self.candidate_names != SPATIAL_CANDIDATE_NAMES: + raise ValueError( + "candidate_names must equal the frozen spatial candidate registry" + ) + for name in self.event_ids: + if not str(name): + raise ValueError("event_ids must not contain empty values") + + rats = np.asarray(self.rat_ids, dtype=str) + sessions = np.asarray(self.session_ids, dtype=str) + if rats.shape != (n_events,) or sessions.shape != (n_events,): + raise ValueError("rat_ids and session_ids must contain one value per event") + if np.any(rats == "") or np.any(sessions == ""): + raise ValueError("rat_ids and session_ids must not contain empty values") + + starts = np.asarray(self.event_start_s, dtype=float) + ends = np.asarray(self.event_end_s, dtype=float) + history = np.asarray(self.history_cutoff_s, dtype=float) + decoder = np.asarray(self.decoder_training_cutoff_s, dtype=float) + for values, name in ( + (starts, "event_start_s"), + (ends, "event_end_s"), + (history, "history_cutoff_s"), + (decoder, "decoder_training_cutoff_s"), + ): + if values.shape != (n_events,) or not np.all(np.isfinite(values)): + raise ValueError(f"{name} must contain one finite value per event") + if np.any(ends <= starts): + raise ValueError("each event must end after it starts") + if np.any(history > starts + _FLOAT_TOL): + raise ValueError("history_cutoff_s must not extend into replay") + if np.any(decoder > starts + _FLOAT_TOL): + raise ValueError("decoder training must not use observations after replay starts") + + emissions = np.asarray(self.log_emissions, dtype=float) + if emissions.ndim != 3 or emissions.shape[0] != n_events: + raise ValueError("log_emissions must have shape (event, time, spatial_bin)") + if emissions.shape[2] < 2: + raise ValueError("log_emissions must contain at least two spatial bins") + + fields = np.asarray(self.candidate_fields, dtype=float) + expected_fields = ( + n_events, + len(self.candidate_names), + self.n_spatial_bins, + ) + if fields.shape != expected_fields or not np.all(np.isfinite(fields)): + raise ValueError( + "candidate_fields must be finite with shape " + "(event, candidate, spatial_bin)" + ) + field_times = np.asarray(self.field_available_s, dtype=float) + if field_times.shape != (n_events, len(self.candidate_names)): + raise ValueError("field_available_s must have shape (event, candidate)") + if not np.all(np.isfinite(field_times)): + raise ValueError("field_available_s must be finite") + if np.any(field_times > starts[:, None] + _FLOAT_TOL): + raise ValueError("candidate fields must use only evidence available before replay") + available = np.asarray(self.candidate_available) + if available.dtype != np.bool_ or available.shape != field_times.shape: + raise ValueError("candidate_available must be boolean with shape (event, candidate)") + + mask = np.asarray(self.time_mask) + if mask.dtype != np.bool_ or mask.shape != emissions.shape[:2]: + raise ValueError("time_mask must be boolean with shape (event, time)") + if np.any(mask.sum(axis=1) < 1): + raise ValueError("every event must contain at least one valid emission row") + spatial = np.asarray(self.active_spatial_mask) + if spatial.dtype != np.bool_ or spatial.shape != ( + n_events, + emissions.shape[2], + ): + raise ValueError( + "active_spatial_mask must be boolean with shape (event, spatial_bin)" + ) + if np.any(spatial.sum(axis=1) < 2): + raise ValueError("every event must contain at least two active spatial bins") + coordinates = np.asarray(self.spatial_coordinates, dtype=float) + if coordinates.shape != (n_events, emissions.shape[2], 2): + raise ValueError( + "spatial_coordinates must have shape (event, spatial_bin, xy)" + ) + if not np.all(np.isfinite(coordinates[spatial])): + raise ValueError("active spatial coordinates must be finite") + if np.any(np.isfinite(coordinates[~spatial])): + raise ValueError("inactive spatial coordinates must be NaN") + point_spread = np.asarray(self.decoder_point_spread_cm, dtype=float) + if ( + point_spread.shape != (n_events,) + or not np.all(np.isfinite(point_spread)) + or np.any(point_spread <= 0.0) + ): + raise ValueError( + "decoder_point_spread_cm must contain one finite positive value per event" + ) + if np.any(np.isnan(emissions)) or np.any(emissions == np.inf): + raise ValueError("log_emissions must not contain NaN or positive infinity") + offsets = np.asarray(self.log_emission_offsets, dtype=float) + if offsets.shape != mask.shape or not np.all(np.isfinite(offsets)): + raise ValueError( + "log_emission_offsets must be finite with shape (event, time)" + ) + + for event_index in range(n_events): + active = spatial[event_index] + for time_index in range(emissions.shape[1]): + row = emissions[event_index, time_index] + if mask[event_index, time_index]: + finite = np.isfinite(row) & active + if not np.any(finite): + raise ValueError("every valid emission row needs finite active support") + if np.any(np.isfinite(row[~active])): + raise ValueError("emissions outside active spatial support must be -inf") + row_max = float(np.max(row[finite])) + if abs(row_max) > 1e-8: + raise ValueError( + "each valid log-emission row must be max-shifted to zero" + ) + elif np.any(np.isfinite(row)): + raise ValueError("padded emission rows must contain only -inf") + + base = np.asarray(self.nuisance_base, dtype=float) + if base.shape != spatial.shape or not np.all(np.isfinite(base)): + raise ValueError("nuisance_base must be finite with shape (event, spatial_bin)") + if np.any(base < 0.0) or np.any(base[~spatial] != 0.0): + raise ValueError("nuisance_base must be nonnegative and zero off active support") + if np.any(base.sum(axis=1) <= 0.0): + raise ValueError("every nuisance_base row must contain positive mass") + if np.any(fields[~np.broadcast_to(spatial[:, None, :], fields.shape)] != 0.0): + raise ValueError("candidate_fields must be zero off active spatial support") + + if self.well_masses is None: + if self.well_ids: + raise ValueError("well_ids require well_masses") + else: + well_mass = np.asarray(self.well_masses, dtype=float) + if ( + well_mass.ndim != 2 + or well_mass.shape[0] != n_events + or well_mass.shape[1] != len(self.well_ids) + or len(set(self.well_ids)) != len(self.well_ids) + ): + raise ValueError("well_masses and well_ids have incompatible shapes") + if not np.all(np.isfinite(well_mass)) or np.any(well_mass < 0.0): + raise ValueError("well_masses must contain finite nonnegative values") + if not np.allclose( + well_mass.sum(axis=1), + 1.0, + rtol=0.0, + atol=1e-8, + ): + raise ValueError("every well_masses row must sum to one") + + +@dataclass(frozen=True, slots=True) +class SpatialComparisonConfig: + """Predeclared grouped scoring and abstention thresholds.""" + + temperatures: tuple[float, ...] = (0.0, 0.25, 0.5, 1.0, 2.0, 4.0, 8.0) + minimum_rats: int = 5 + bootstrap_replicates: int = 5000 + maximum_field_correlation: float = 0.98 + simultaneous_confidence_level: float = 0.95 + seed: int = 7 + + def validate(self) -> None: + temperatures = np.asarray(self.temperatures, dtype=float) + if temperatures.ndim != 1 or temperatures.size < 2: + raise ValueError("temperatures must contain at least two values") + if not np.all(np.isfinite(temperatures)) or np.any(temperatures < 0.0): + raise ValueError("temperatures must be finite and nonnegative") + if 0.0 not in self.temperatures: + raise ValueError("temperatures must include the null temperature zero") + if self.minimum_rats < 2: + raise ValueError("minimum_rats must be at least two") + if self.bootstrap_replicates < 100: + raise ValueError("bootstrap_replicates must be at least 100") + if not 0.0 < self.maximum_field_correlation <= 1.0: + raise ValueError("maximum_field_correlation must lie in (0, 1]") + if not 0.5 < self.simultaneous_confidence_level < 1.0: + raise ValueError("simultaneous_confidence_level must lie in (0.5, 1)") + + +@dataclass(frozen=True, slots=True) +class SpatialRecoveryRecord: + """One computed emission-injection recovery result.""" + + generator: str + split_unit: str + selected_candidate: str + selected_margin: float + selected_margin_lower: float + decisive: bool + n_held_out_groups: int + spatial_sigma_multiplier: float + + +@dataclass(frozen=True, slots=True) +class SpatialRecoveryGate: + """Computed recovery evidence required before a biological claim is made. + + Records are produced by run_spatial_recovery_checks. There are no + user-asserted pass flags: every pure generator must be decisively recovered + under both held-out-animal and held-out-session calibration, while every + registered 50/50 mixture must trigger uncertainty abstention under both. + """ + + pure_records: tuple[SpatialRecoveryRecord, ...] + mixture_records: tuple[SpatialRecoveryRecord, ...] + null_records: tuple[SpatialRecoveryRecord, ...] + source_event_count: int + common_event_count: int + excluded_event_ids: tuple[str, ...] + required_mixtures: tuple[str, ...] = ( + "smoothing_revision+td_error", + "smoothing_revision+prospective", + "smoothing_revision+recency", + "smoothing_revision+posterior_content", + ) + required_sigma_multipliers: tuple[float, ...] = (0.5, 1.0, 2.0) + + @property + def passed(self) -> bool: + required_splits = ("leave_one_rat_out", "leave_one_session_out") + required_cells = { + (split_unit, float(multiplier)) + for split_unit in required_splits + for multiplier in self.required_sigma_multipliers + } + for generator in SPATIAL_CANDIDATE_NAMES: + matches = [ + record + for record in self.pure_records + if record.generator == generator + ] + observed_cells = { + (record.split_unit, float(record.spatial_sigma_multiplier)) + for record in matches + } + if ( + observed_cells != required_cells + or len(matches) != len(required_cells) + or any( + record.selected_candidate != generator or not record.decisive + for record in matches + ) + ): + return False + if ( + self.source_event_count < self.common_event_count + or self.common_event_count < 1 + or self.source_event_count - self.common_event_count + != len(self.excluded_event_ids) + ): + return False + for mixture in self.required_mixtures: + matches = [ + record + for record in self.mixture_records + if record.generator == mixture + ] + observed_cells = { + (record.split_unit, float(record.spatial_sigma_multiplier)) + for record in matches + } + if ( + observed_cells != required_cells + or len(matches) != len(required_cells) + or any(record.decisive for record in matches) + ): + return False + null_cells = { + (record.split_unit, float(record.spatial_sigma_multiplier)) + for record in self.null_records + if record.generator == NULL_CANDIDATE_NAME + } + if ( + null_cells != required_cells + or len(self.null_records) != len(required_cells) + or any(record.decisive for record in self.null_records) + ): + return False + return True + + +@dataclass(frozen=True, slots=True) +class SpatialCandidateFold: + candidate: str + held_out_rat: str + temperature: float + mean_log_score_per_bin: float + n_events: int + n_sessions: int + + +@dataclass(frozen=True, slots=True) +class SpatialTargetContrast: + """Prespecified animal-level smoothing contrast with simultaneous coverage.""" + + alternative: str + mean_margin: float + simultaneous_lower_bound: float + exact_one_sided_sign_flip_p: float + + +@dataclass(frozen=True, slots=True) +class SpatialReplayComparison: + """LORO comparison with confirmatory, prespecified smoothing contrasts. + + winner and winner_margin_ci are descriptive because their identities are + selected on the same held-out scores. The status decision uses only the + prespecified smoothing-revision contrasts against every alternative and + their joint bootstrap lower bounds. + """ + + candidate_names: tuple[str, ...] + folds: tuple[SpatialCandidateFold, ...] + rat_ids: tuple[str, ...] + rat_scores: NDArray[np.float64] + winner: str + runner_up: str + winner_margin: float + winner_margin_ci: tuple[float, float] + target_candidate: str + target_contrasts: tuple[SpatialTargetContrast, ...] + simultaneous_confidence_level: float + maximum_field_correlation: float + common_event_count: int + status: str + abstention_reasons: tuple[str, ...] + + +def subset_spatial_replay_dataset( + dataset: SpatialReplayDataset, + selected: ArrayLike, +) -> SpatialReplayDataset: + """Return an auditable event subset while preserving original identifiers.""" + + dataset.validate() + mask = np.asarray(selected) + if mask.dtype != np.bool_ or mask.shape != (dataset.n_events,): + raise ValueError("selected must be a boolean vector with one value per event") + if not np.any(mask): + raise ValueError("selected must retain at least one event") + + def take(values: ArrayLike) -> NDArray[Any]: + return np.asarray(values)[mask] + + subset = replace( + dataset, + event_ids=tuple( + event_id + for event_id, keep in zip(dataset.event_ids, mask, strict=True) + if bool(keep) + ), + rat_ids=take(dataset.rat_ids), + session_ids=take(dataset.session_ids), + event_start_s=take(dataset.event_start_s), + event_end_s=take(dataset.event_end_s), + history_cutoff_s=take(dataset.history_cutoff_s), + decoder_training_cutoff_s=take(dataset.decoder_training_cutoff_s), + field_available_s=take(dataset.field_available_s), + log_emissions=take(dataset.log_emissions), + log_emission_offsets=take(dataset.log_emission_offsets), + time_mask=take(dataset.time_mask), + active_spatial_mask=take(dataset.active_spatial_mask), + spatial_coordinates=take(dataset.spatial_coordinates), + decoder_point_spread_cm=take(dataset.decoder_point_spread_cm), + nuisance_base=take(dataset.nuisance_base), + candidate_fields=take(dataset.candidate_fields), + candidate_available=take(dataset.candidate_available), + well_masses=( + None if dataset.well_masses is None else take(dataset.well_masses) + ), + ) + subset.validate() + return subset + + +def _exact_one_sided_sign_flip_p(values: NDArray[np.float64]) -> float: + """Exact randomization p for a positive equal-animal mean contrast.""" + + differences = np.asarray(values, dtype=float) + if ( + differences.ndim != 1 + or differences.size < 1 + or not np.all(np.isfinite(differences)) + ): + return float("nan") + observed = float(np.mean(differences)) + n_animals = int(differences.size) + pattern_ids = np.arange(1 << n_animals, dtype=np.uint64)[:, None] + bit_ids = np.arange(n_animals, dtype=np.uint64)[None, :] + signs = np.where(((pattern_ids >> bit_ids) & 1) == 1, 1.0, -1.0) + null_statistics = np.mean(signs * differences[None, :], axis=1) + return float(np.mean(null_statistics >= observed - 1e-15)) + + +def _normalized_base(dataset: SpatialReplayDataset) -> NDArray[np.float64]: + base = np.asarray(dataset.nuisance_base, dtype=float).copy() + base /= base.sum(axis=1, keepdims=True) + return base + + +def _standardized_fields( + dataset: SpatialReplayDataset, + base: NDArray[np.float64], +) -> tuple[NDArray[np.float64], NDArray[np.bool_]]: + fields = np.asarray(dataset.candidate_fields, dtype=float) + mean = np.sum(base[:, None, :] * fields, axis=2, keepdims=True) + centered = fields - mean + variance = np.sum(base[:, None, :] * centered**2, axis=2) + usable = variance > 1e-14 + scale = np.sqrt(np.maximum(variance, 1e-14)) + standardized = centered / scale[:, :, None] + standardized[~np.broadcast_to(dataset.active_spatial_mask[:, None, :], fields.shape)] = 0.0 + return standardized, usable + + +def spatial_common_candidate_mask( + dataset: SpatialReplayDataset, +) -> NDArray[np.bool_]: + """Identify the predeclared all-candidate complete-case event cohort.""" + + dataset.validate() + base = _normalized_base(dataset) + _, field_usable = _standardized_fields(dataset, base) + common = np.all(np.asarray(dataset.candidate_available, dtype=bool), axis=1) + common &= np.all(field_usable, axis=1) + return np.asarray(common, dtype=np.bool_) + + +def _event_scores_for_temperature( + dataset: SpatialReplayDataset, + base: NDArray[np.float64], + standardized: NDArray[np.float64], + candidate_index: int | None, + temperature: float, +) -> NDArray[np.float64]: + scores = np.empty(dataset.n_events, dtype=float) + for event_index in range(dataset.n_events): + active = dataset.active_spatial_mask[event_index] + log_prior = np.full(dataset.n_spatial_bins, -np.inf, dtype=float) + active_base = base[event_index, active] + log_prior[active] = np.log(active_base) + if candidate_index is not None: + log_prior[active] += ( + float(temperature) * standardized[event_index, candidate_index, active] + ) + log_prior[active] -= logsumexp(log_prior[active]) + total = 0.0 + count = 0 + for time_index in np.flatnonzero(dataset.time_mask[event_index]): + total += float( + logsumexp( + dataset.log_emissions[event_index, time_index, active] + + log_prior[active] + ) + + dataset.log_emission_offsets[event_index, time_index] + ) + count += 1 + scores[event_index] = total / count + return scores + + +def _equal_rat_session_mean( + values: NDArray[np.float64], + selected: NDArray[np.bool_], + rats: NDArray[np.str_], + sessions: NDArray[np.str_], +) -> float: + rat_means: list[float] = [] + for rat in sorted(set(rats[selected])): + rat_mask = selected & (rats == rat) + session_means = [ + float(np.mean(values[rat_mask & (sessions == session)])) + for session in sorted(set(sessions[rat_mask])) + ] + if session_means: + rat_means.append(float(np.mean(session_means))) + return float(np.mean(rat_means)) if rat_means else -np.inf + + +def _field_correlation( + standardized: NDArray[np.float64], + common: NDArray[np.bool_], + active: NDArray[np.bool_], +) -> float: + vectors: list[NDArray[np.float64]] = [] + for candidate_index in range(standardized.shape[1]): + values = np.concatenate( + [ + standardized[event_index, candidate_index, active[event_index]] + for event_index in np.flatnonzero(common) + ] + ) + vectors.append(values) + matrix = np.asarray(vectors, dtype=float) + nonconstant = np.std(matrix, axis=1) > 1e-12 + if int(np.sum(nonconstant)) < 2: + return 1.0 + correlation = np.corrcoef(matrix[nonconstant]) + off_diagonal = np.abs( + correlation[~np.eye(correlation.shape[0], dtype=bool)] + ) + return float(np.max(off_diagonal)) if off_diagonal.size else 0.0 + + +def compare_spatial_replay_candidates( + dataset: SpatialReplayDataset, + config: SpatialComparisonConfig | None = None, + *, + recovery_gate: SpatialRecoveryGate | None = None, +) -> SpatialReplayComparison: + """Score raw replay emissions with LORO tuning and animal-level abstention.""" + + dataset.validate() + config = SpatialComparisonConfig() if config is None else config + config.validate() + rats = np.asarray(dataset.rat_ids, dtype=str) + sessions = np.asarray(dataset.session_ids, dtype=str) + base = _normalized_base(dataset) + standardized, field_usable = _standardized_fields(dataset, base) + common = np.all(np.asarray(dataset.candidate_available, dtype=bool), axis=1) + common &= np.all(field_usable, axis=1) + unique_rats = tuple(sorted(set(rats[common]))) + all_names = (*dataset.candidate_names, NULL_CANDIDATE_NAME) + + temperature_scores: dict[tuple[int, float], NDArray[np.float64]] = {} + for candidate_index in range(len(dataset.candidate_names)): + for temperature in config.temperatures: + temperature_scores[candidate_index, float(temperature)] = ( + _event_scores_for_temperature( + dataset, + base, + standardized, + candidate_index, + float(temperature), + ) + ) + null_scores = _event_scores_for_temperature( + dataset, + base, + standardized, + None, + 0.0, + ) + + folds: list[SpatialCandidateFold] = [] + rat_score_rows: list[list[float]] = [] + retained_rats: list[str] = [] + for held_out_rat in unique_rats: + train = common & (rats != held_out_rat) + test = common & (rats == held_out_rat) + if not np.any(train) or not np.any(test): + continue + row: list[float] = [] + for candidate_index, candidate in enumerate(dataset.candidate_names): + objectives = [ + _equal_rat_session_mean( + temperature_scores[candidate_index, float(temperature)], + train, + rats, + sessions, + ) + for temperature in config.temperatures + ] + best_index = int(np.argmax(objectives)) + temperature = float(config.temperatures[best_index]) + score = temperature_scores[candidate_index, temperature] + held_score = _equal_rat_session_mean(score, test, rats, sessions) + row.append(held_score) + folds.append( + SpatialCandidateFold( + candidate=candidate, + held_out_rat=held_out_rat, + temperature=temperature, + mean_log_score_per_bin=held_score, + n_events=int(np.sum(test)), + n_sessions=len(set(sessions[test])), + ) + ) + held_null = _equal_rat_session_mean(null_scores, test, rats, sessions) + row.append(held_null) + folds.append( + SpatialCandidateFold( + candidate=NULL_CANDIDATE_NAME, + held_out_rat=held_out_rat, + temperature=0.0, + mean_log_score_per_bin=held_null, + n_events=int(np.sum(test)), + n_sessions=len(set(sessions[test])), + ) + ) + retained_rats.append(held_out_rat) + rat_score_rows.append(row) + + target_candidate = "smoothing_revision" + target_index = all_names.index(target_candidate) + alternative_indices = [ + index for index in range(len(all_names)) if index != target_index + ] + ci: tuple[float, float] + if rat_score_rows: + rat_scores = np.asarray(rat_score_rows, dtype=float) + mean_scores = rat_scores.mean(axis=0) + order = np.argsort(mean_scores)[::-1] + winner_index = int(order[0]) + runner_index = int(order[1]) + winner = all_names[winner_index] + runner_up = all_names[runner_index] + rat_differences = rat_scores[:, winner_index] - rat_scores[:, runner_index] + winner_margin = float(np.mean(rat_differences)) + rng = np.random.default_rng(config.seed) + descriptive_draws = rng.choice( + rat_differences, + size=(config.bootstrap_replicates, len(rat_differences)), + replace=True, + ).mean(axis=1) + tail = (1.0 - config.simultaneous_confidence_level) / 2.0 + quantiles = np.quantile( + descriptive_draws, + [tail, 1.0 - tail], + ) + ci = (float(quantiles[0]), float(quantiles[1])) + + paired = ( + rat_scores[:, [target_index]] + - rat_scores[:, alternative_indices] + ) + observed_margins = paired.mean(axis=0) + bootstrap_indices = rng.integers( + 0, + paired.shape[0], + size=(config.bootstrap_replicates, paired.shape[0]), + ) + bootstrap_margins = paired[bootstrap_indices].mean(axis=1) + maximum_shortfall = np.max( + observed_margins[None, :] - bootstrap_margins, + axis=1, + ) + critical_value = float( + np.quantile( + maximum_shortfall, + config.simultaneous_confidence_level, + ) + ) + simultaneous_lower = observed_margins - critical_value + target_contrasts = tuple( + SpatialTargetContrast( + alternative=all_names[alternative_index], + mean_margin=float(observed_margins[position]), + simultaneous_lower_bound=float(simultaneous_lower[position]), + exact_one_sided_sign_flip_p=_exact_one_sided_sign_flip_p( + paired[:, position] + ), + ) + for position, alternative_index in enumerate(alternative_indices) + ) + else: + rat_scores = np.empty((0, len(all_names)), dtype=float) + winner = NULL_CANDIDATE_NAME + runner_up = dataset.candidate_names[0] + winner_margin = float("nan") + ci = (float("nan"), float("nan")) + target_contrasts = tuple( + SpatialTargetContrast( + alternative=all_names[index], + mean_margin=float("nan"), + simultaneous_lower_bound=float("nan"), + exact_one_sided_sign_flip_p=float("nan"), + ) + for index in alternative_indices + ) + + maximum_correlation = ( + _field_correlation(standardized, common, dataset.active_spatial_mask) + if np.any(common) + else 1.0 + ) + reasons: list[str] = [] + if len(retained_rats) < config.minimum_rats: + reasons.append("too_few_independent_rats") + if recovery_gate is None: + reasons.append("recovery_gate_missing") + elif not recovery_gate.passed: + reasons.append("recovery_gate_failed") + if maximum_correlation > config.maximum_field_correlation: + reasons.append("candidate_fields_collinear") + target_lower = np.asarray( + [contrast.simultaneous_lower_bound for contrast in target_contrasts], + dtype=float, + ) + if np.any(~np.isfinite(target_lower)) or np.any(target_lower <= 0.0): + reasons.append("smoothing_revision_contrast_uncertain") + target_sign_flip_p = np.asarray( + [contrast.exact_one_sided_sign_flip_p for contrast in target_contrasts], + dtype=float, + ) + alpha = 1.0 - config.simultaneous_confidence_level + if ( + np.any(~np.isfinite(target_sign_flip_p)) + or np.any(target_sign_flip_p > alpha + 1e-15) + ): + reasons.append("animal_sign_flip_resolution_insufficient") + status = "identified" if not reasons else "abstain" + + return SpatialReplayComparison( + candidate_names=all_names, + folds=tuple(folds), + rat_ids=tuple(retained_rats), + rat_scores=rat_scores, + winner=winner, + runner_up=runner_up, + winner_margin=winner_margin, + winner_margin_ci=ci, + target_candidate=target_candidate, + target_contrasts=target_contrasts, + simultaneous_confidence_level=config.simultaneous_confidence_level, + maximum_field_correlation=maximum_correlation, + common_event_count=int(np.sum(common)), + status=status, + abstention_reasons=tuple(reasons), + ) + + +__all__ = [ + "NULL_CANDIDATE_NAME", + "REPLAY_SPATIAL_SCHEMA_VERSION", + "SPATIAL_CANDIDATE_NAMES", + "SignedRevisionField", + "SpatialCandidateFold", + "SpatialComparisonConfig", + "SpatialRecoveryGate", + "SpatialRecoveryRecord", + "SpatialReplayComparison", + "SpatialTargetContrast", + "SpatialReplayDataset", + "build_signed_revision_field", + "compare_spatial_replay_candidates", + "spatial_common_candidate_mask", + "subset_spatial_replay_dataset", +] diff --git a/tests/test_replay_spatial.py b/tests/test_replay_spatial.py new file mode 100644 index 0000000..0b5984e --- /dev/null +++ b/tests/test_replay_spatial.py @@ -0,0 +1,685 @@ +from __future__ import annotations + +import csv +import hashlib +import io +import json +from dataclasses import replace +from pathlib import Path + +import numpy as np +import pytest + +from bayesian_ach.replay_artifact import ( + LaterOutcomeTable, + ReplaySpatialManifest, + load_later_outcome_artifact, + load_spatial_predictor_artifact, + write_later_outcome_artifact, + write_spatial_predictor_artifact, +) +from bayesian_ach.replay_recovery import ( + SpatialInjectionRecoveryConfig, + run_spatial_recovery_checks, +) +from bayesian_ach.replay_spatial import ( + SPATIAL_CANDIDATE_NAMES, + SpatialComparisonConfig, + SpatialRecoveryGate, + SpatialRecoveryRecord, + SpatialReplayDataset, + build_signed_revision_field, + compare_spatial_replay_candidates, + subset_spatial_replay_dataset, +) + + +def _softmax(values: np.ndarray) -> np.ndarray: + shifted = values - np.max(values, axis=-1, keepdims=True) + probabilities = np.exp(shifted) + return probabilities / probabilities.sum(axis=-1, keepdims=True) + + +def _dataset(*, signal: bool = True, seed: int = 17) -> SpatialReplayDataset: + rng = np.random.default_rng(seed) + n_rats = 6 + events_per_rat = 8 + n_events = n_rats * events_per_rat + n_time = 5 + n_bins = 12 + n_candidates = len(SPATIAL_CANDIDATE_NAMES) + rats = np.repeat([f"Rat{index}" for index in range(n_rats)], events_per_rat) + sessions = np.asarray( + [ + f"{rat}/Session{1 + (event_index % 2)}" + for event_index, rat in enumerate(rats) + ], + dtype=str, + ) + starts = 100.0 + np.arange(n_events, dtype=float) * 10.0 + fields = rng.normal(size=(n_events, n_candidates, n_bins)) + base = rng.uniform(0.5, 1.5, size=(n_events, n_bins)) + base /= base.sum(axis=1, keepdims=True) + + revision = fields[:, 0] + revision -= np.sum(base * revision, axis=1, keepdims=True) + revision /= np.sqrt( + np.sum(base * revision**2, axis=1, keepdims=True) + ) + log_emissions = np.empty((n_events, n_time, n_bins), dtype=float) + for event_index in range(n_events): + for time_index in range(n_time): + row = rng.normal(0.0, 0.12, n_bins) + if signal: + row += 3.0 * revision[event_index] + row -= np.max(row) + log_emissions[event_index, time_index] = row + + well_mass = _softmax(rng.normal(size=(n_events, 3))) + grid_x, grid_y = np.meshgrid(np.arange(4, dtype=float), np.arange(3, dtype=float)) + coordinates = np.column_stack((grid_x.ravel(), grid_y.ravel())) + return SpatialReplayDataset( + event_ids=tuple(f"event-{index:03d}" for index in range(n_events)), + rat_ids=np.asarray(rats, dtype=str), + session_ids=sessions, + event_start_s=starts, + event_end_s=starts + 0.12, + history_cutoff_s=starts - 0.01, + decoder_training_cutoff_s=starts - 1.0, + field_available_s=np.repeat( + (starts - 0.01)[:, None], + n_candidates, + axis=1, + ), + log_emissions=log_emissions, + log_emission_offsets=rng.normal(size=(n_events, n_time)), + time_mask=np.ones((n_events, n_time), dtype=bool), + active_spatial_mask=np.ones((n_events, n_bins), dtype=bool), + spatial_coordinates=np.broadcast_to( + coordinates[None, :, :], + (n_events, n_bins, 2), + ).copy(), + decoder_point_spread_cm=np.full(n_events, 0.5, dtype=float), + nuisance_base=base, + candidate_fields=fields, + candidate_available=np.ones((n_events, n_candidates), dtype=bool), + well_masses=well_mass, + well_ids=("well-0", "well-1", "well-2"), + ) + + +def _passing_gate() -> SpatialRecoveryGate: + pure = tuple( + SpatialRecoveryRecord( + generator=generator, + split_unit=split_unit, + selected_candidate=generator, + selected_margin=1.0, + selected_margin_lower=0.5, + decisive=True, + n_held_out_groups=4, + spatial_sigma_multiplier=1.0, + ) + for generator in SPATIAL_CANDIDATE_NAMES + for split_unit in ("leave_one_rat_out", "leave_one_session_out") + ) + mixture = tuple( + SpatialRecoveryRecord( + generator="smoothing_revision+td_error", + split_unit=split_unit, + selected_candidate="smoothing_revision", + selected_margin=0.0, + selected_margin_lower=-0.1, + decisive=False, + n_held_out_groups=4, + spatial_sigma_multiplier=1.0, + ) + for split_unit in ("leave_one_rat_out", "leave_one_session_out") + ) + null_records = tuple( + SpatialRecoveryRecord( + generator="null", + split_unit=split_unit, + selected_candidate="td_error", + selected_margin=0.0, + selected_margin_lower=-0.1, + decisive=False, + n_held_out_groups=6, + spatial_sigma_multiplier=1.0, + ) + for split_unit in ("leave_one_rat_out", "leave_one_session_out") + ) + return SpatialRecoveryGate( + pure_records=pure, + mixture_records=mixture, + null_records=null_records, + source_event_count=1, + common_event_count=1, + excluded_event_ids=(), + required_mixtures=("smoothing_revision+td_error",), + required_sigma_multipliers=(1.0,), + ) + + +def _config() -> SpatialComparisonConfig: + return SpatialComparisonConfig( + temperatures=(0.0, 0.5, 1.0, 2.0, 3.0, 4.0), + bootstrap_replicates=1000, + maximum_field_correlation=0.999, + seed=11, + ) + + +def _manifest( + directory: Path | None = None, + dataset: SpatialReplayDataset | None = None, +) -> ReplaySpatialManifest: + dataset_sha256 = "2" * 64 + dataset_manifest_file_sha256 = "6" * 64 + file_records_sha256 = "7" * 64 + report = { + "schema_version": "hipporeplayimm.pf-dataset-verification.v1", + "status": "pass", + "dataset_sha256": dataset_sha256, + "dataset_manifest_file_sha256": dataset_manifest_file_sha256, + "verified_file_count": 124, + "verified_total_bytes": 425_953_051, + "verified_session_count": 8, + "verified_file_records_sha256": file_records_sha256, + "missing_files": [], + "extra_files": [], + } + report_bytes = ( + json.dumps(report, indent=2, sort_keys=True) + "\n" + ).encode("utf-8") + report_sha256 = hashlib.sha256(report_bytes).hexdigest() + + route_parameters = {"median_window_s": 0.167, "gaussian_sigma_s": 0.1} + route_parameters_sha256 = hashlib.sha256( + json.dumps( + route_parameters, + sort_keys=True, + separators=(",", ":"), + ).encode("utf-8") + ).hexdigest() + route_payload = { + "analysis": "replay_behavior_route_primitives", + "producer_commit": "9" * 40, + "producer_clean_worktree": True, + "route_smoothing_scope": "within_completed_fill_interval", + "parameters": route_parameters, + "parameters_sha256": route_parameters_sha256, + "output_sha256": { + "replay_behavior_route_segments.csv": "b" * 64, + "replay_behavior_route_segment_points.csv": "c" * 64, + }, + } + route_bytes = ( + json.dumps(route_payload, indent=2, sort_keys=True) + "\n" + ).encode("utf-8") + route_sha256 = hashlib.sha256(route_bytes).hexdigest() + + audit_buffer = io.StringIO(newline="") + writer = csv.DictWriter( + audit_buffer, + fieldnames=("event_id", "session", "rat", "event_index"), + lineterminator="\n", + ) + writer.writeheader() + cohort: list[dict[str, object]] = [] + if dataset is not None: + for index, (event_id, session, rat) in enumerate( + zip( + dataset.event_ids, + np.asarray(dataset.session_ids, dtype=str), + np.asarray(dataset.rat_ids, dtype=str), + strict=True, + ) + ): + writer.writerow( + { + "event_id": event_id, + "session": session, + "rat": rat, + "event_index": index, + } + ) + cohort.append( + { + "event_id": event_id, + "session": str(session), + "event_index": index, + } + ) + audit_bytes = audit_buffer.getvalue().encode("utf-8") + audit_sha256 = hashlib.sha256(audit_bytes).hexdigest() + cohort_sha256 = hashlib.sha256( + ( + json.dumps(cohort, sort_keys=True, separators=(",", ":")) + + "\n" + ).encode("utf-8") + ).hexdigest() + + if directory is not None: + if dataset is None: + raise ValueError("dataset is required when writing provenance sidecars") + directory.mkdir(parents=True, exist_ok=True) + (directory / "replay_spatial_dataset_verification.json").write_bytes( + report_bytes + ) + (directory / "replay_spatial_route_manifest.json").write_bytes( + route_bytes + ) + (directory / "replay_spatial_event_audit.csv").write_bytes(audit_bytes) + + return ReplaySpatialManifest( + producer_repository="IPS-Stuttgart/HippoReplayDynamics", + producer_commit="1" * 40, + dataset_id="PfeifferFoster-open-field-2013", + dataset_sha256=dataset_sha256, + dataset_manifest_file_sha256=dataset_manifest_file_sha256, + dataset_verifier_report_file=( + "replay_spatial_dataset_verification.json" + ), + dataset_verifier_report_sha256=report_sha256, + dataset_verified_file_count=124, + dataset_verified_total_bytes=425_953_051, + dataset_verified_session_count=8, + dataset_verified_file_records_sha256=file_records_sha256, + route_manifest_file="replay_spatial_route_manifest.json", + route_manifest_file_sha256=route_sha256, + route_producer_commit="9" * 40, + route_producer_clean_worktree=True, + route_parameters_sha256=route_parameters_sha256, + route_segments_sha256="b" * 64, + route_points_sha256="c" * 64, + cohort_sha256=cohort_sha256, + event_audit_file="replay_spatial_event_audit.csv", + event_audit_sha256=audit_sha256, + event_selection_parameters_sha256="3" * 64, + behavior_field_parameters_sha256="4" * 64, + decoder_parameters_sha256=( + "a79fa8a1f55a964c4367853cc120efc9b742ec4e327c277ede78cfd6a277f20b" + ), + ) + +def test_signed_revision_field_is_kl_weighted_signed_and_pre_replay() -> None: + filtered = np.array([[0.8, 0.2], [0.5, 0.5]], dtype=float) + smoothed = np.array([[0.6, 0.4], [0.2, 0.8]], dtype=float) + ends = np.array([5.0, 9.0], dtype=float) + + mapping = np.array( + [ + [[1.0, 0.0, 0.0], [0.0, 1.0, 0.0]], + [[0.0, 1.0, 0.0], [0.0, 0.0, 1.0]], + ] + ) + result = build_signed_revision_field( + filtered, + smoothed, + ends, + event_start_s=10.0, + state_to_spatial=mapping, + recency_tau_s=4.0, + ) + + kl = np.sum(smoothed * np.log(smoothed / filtered), axis=1) + weights = kl * np.exp(-(10.0 - ends) / 4.0) + projected = np.einsum("hs,hsb->hb", smoothed - filtered, mapping) + expected = np.sum(weights[:, None] * projected, axis=0) + np.testing.assert_allclose(result.per_snippet_kl, kl, atol=1e-14) + np.testing.assert_allclose(result.snippet_weights, weights, atol=1e-14) + np.testing.assert_allclose(result.signed_field, expected, atol=1e-14) + np.testing.assert_allclose(result.signed_field.sum(), 0.0, atol=1e-14) + assert result.identifiable is True + + with pytest.raises(ValueError, match="end before replay"): + build_signed_revision_field( + filtered, + smoothed, + np.array([5.0, 10.1]), + event_start_s=10.0, + ) + + +def test_signed_revision_rejects_smoothing_mass_outside_filter_support() -> None: + with pytest.raises(ValueError, match="outside filtering support"): + build_signed_revision_field( + np.array([[1.0, 0.0]]), + np.array([[0.5, 0.5]]), + np.array([1.0]), + event_start_s=2.0, + ) + + +def test_dataset_rejects_future_candidate_or_decoder_evidence() -> None: + dataset = _dataset() + dataset.validate() + + future_field = np.asarray(dataset.field_available_s).copy() + future_field[0, 0] = dataset.event_start_s[0] + 0.1 + with pytest.raises(ValueError, match="before replay"): + replace(dataset, field_available_s=future_field).validate() + + future_decoder = np.asarray(dataset.decoder_training_cutoff_s).copy() + future_decoder[0] = dataset.event_start_s[0] + 0.1 + with pytest.raises(ValueError, match="decoder training"): + replace(dataset, decoder_training_cutoff_s=future_decoder).validate() + + +def test_loro_raw_emission_score_recovers_revision_only_with_gate() -> None: + dataset = _dataset(signal=True) + missing_gate = compare_spatial_replay_candidates(dataset, _config()) + assert missing_gate.winner == "smoothing_revision" + assert missing_gate.status == "abstain" + assert "recovery_gate_missing" in missing_gate.abstention_reasons + + result = compare_spatial_replay_candidates( + dataset, + _config(), + recovery_gate=_passing_gate(), + ) + + assert result.winner == "smoothing_revision" + assert result.runner_up != result.winner + assert result.winner_margin > 0.0 + assert result.winner_margin_ci[0] > 0.0 + assert all( + contrast.simultaneous_lower_bound > 0.0 + for contrast in result.target_contrasts + ) + assert result.target_candidate == "smoothing_revision" + assert all( + contrast.exact_one_sided_sign_flip_p <= 0.05 + for contrast in result.target_contrasts + ) + assert result.status == "identified" + assert result.rat_ids == ("Rat0", "Rat1", "Rat2", "Rat3", "Rat4", "Rat5") + assert result.rat_scores.shape == (6, len(SPATIAL_CANDIDATE_NAMES) + 1) + assert all(fold.n_sessions == 2 for fold in result.folds) + + +def test_four_rats_can_never_identify_at_95_percent() -> None: + dataset = _dataset(signal=True) + selected = np.asarray(dataset.rat_ids, dtype=str) < "Rat4" + four_rats = subset_spatial_replay_dataset(dataset, selected) + result = compare_spatial_replay_candidates( + four_rats, + _config(), + recovery_gate=_passing_gate(), + ) + + assert result.winner == "smoothing_revision" + assert result.status == "abstain" + assert "too_few_independent_rats" in result.abstention_reasons + assert "animal_sign_flip_resolution_insufficient" in result.abstention_reasons + assert all( + contrast.exact_one_sided_sign_flip_p >= 1.0 / 16.0 + for contrast in result.target_contrasts + ) + + +def test_candidate_differences_are_invariant_to_log_emission_offsets() -> None: + dataset = _dataset(signal=True) + baseline = compare_spatial_replay_candidates( + dataset, + _config(), + recovery_gate=_passing_gate(), + ) + shift = np.linspace(-100.0, 100.0, dataset.n_events)[:, None] + changed = replace( + dataset, + log_emission_offsets=dataset.log_emission_offsets + shift, + ) + shifted = compare_spatial_replay_candidates( + changed, + _config(), + recovery_gate=_passing_gate(), + ) + + assert shifted.winner == baseline.winner + assert shifted.runner_up == baseline.runner_up + np.testing.assert_allclose( + shifted.rat_scores + - shifted.rat_scores[:, [-1]], + baseline.rat_scores + - baseline.rat_scores[:, [-1]], + atol=1e-12, + ) + np.testing.assert_allclose( + shifted.winner_margin, + baseline.winner_margin, + atol=1e-12, + ) + + +def test_uninformative_emissions_force_uncertainty_abstention() -> None: + dataset = _dataset(signal=False) + zero = np.zeros_like(dataset.log_emissions) + result = compare_spatial_replay_candidates( + replace(dataset, log_emissions=zero), + _config(), + recovery_gate=_passing_gate(), + ) + + assert result.status == "abstain" + assert "smoothing_revision_contrast_uncertain" in result.abstention_reasons + assert result.winner_margin == pytest.approx(0.0, abs=1e-14) + + +def test_recovery_is_computed_from_loao_and_loso_emission_injections() -> None: + dataset = _dataset(signal=False) + gate = run_spatial_recovery_checks( + dataset, + _config(), + SpatialInjectionRecoveryConfig( + injection_temperature=4.0, + spatial_sigma_multipliers=(1.0,), + emission_noise_sd_nats=0.0, + mixtures=(("smoothing_revision", "td_error"),), + seed=29, + ), + ) + + assert len(gate.pure_records) == 2 * len(SPATIAL_CANDIDATE_NAMES) + assert {record.generator for record in gate.pure_records} == set( + SPATIAL_CANDIDATE_NAMES + ) + assert {record.split_unit for record in gate.pure_records} == { + "leave_one_rat_out", + "leave_one_session_out", + } + assert {record.generator for record in gate.mixture_records} == { + "smoothing_revision+td_error" + } + assert all(record.n_held_out_groups >= 4 for record in gate.pure_records) + assert all(np.isfinite(record.selected_margin) for record in gate.pure_records) + assert len(gate.null_records) == 2 + assert all(record.generator == "null" for record in gate.null_records) + assert all(not record.decisive for record in gate.null_records) + assert gate.source_event_count == dataset.n_events + assert gate.common_event_count == dataset.n_events + assert gate.excluded_event_ids == () + + +def test_recovery_subsets_to_exact_common_cohort_and_reports_exclusions() -> None: + dataset = _dataset(signal=False) + available = np.asarray(dataset.candidate_available).copy() + available[0, 0] = False + incomplete = replace(dataset, candidate_available=available) + gate = run_spatial_recovery_checks( + incomplete, + _config(), + SpatialInjectionRecoveryConfig( + injection_temperature=4.0, + spatial_sigma_multipliers=(1.0,), + emission_noise_sd_nats=0.0, + mixtures=(("smoothing_revision", "td_error"),), + seed=31, + ), + ) + + assert gate.source_event_count == dataset.n_events + assert gate.common_event_count == dataset.n_events - 1 + assert gate.excluded_event_ids == (dataset.event_ids[0],) + assert all(np.isfinite(record.selected_margin) for record in gate.pure_records) + + +def test_zero_complete_case_cohort_freezes_technical_abstention() -> None: + dataset = _dataset(signal=False) + unavailable = np.zeros_like(dataset.candidate_available, dtype=bool) + empty = replace(dataset, candidate_available=unavailable) + gate = run_spatial_recovery_checks( + empty, + _config(), + SpatialInjectionRecoveryConfig( + spatial_sigma_multipliers=(1.0,), + mixtures=(("smoothing_revision", "td_error"),), + ), + ) + result = compare_spatial_replay_candidates( + empty, + _config(), + recovery_gate=gate, + ) + + assert gate.common_event_count == 0 + assert gate.source_event_count == dataset.n_events + assert gate.excluded_event_ids == dataset.event_ids + assert gate.passed is False + assert result.common_event_count == 0 + assert result.status == "abstain" + assert "too_few_independent_rats" in result.abstention_reasons + assert "recovery_gate_failed" in result.abstention_reasons + + +def test_decisive_td_win_cannot_pass_as_mixture_abstention() -> None: + passing = _passing_gate() + decisive_td = tuple( + SpatialRecoveryRecord( + generator="smoothing_revision+td_error", + split_unit=split_unit, + selected_candidate="td_error", + selected_margin=0.4, + selected_margin_lower=0.2, + decisive=True, + n_held_out_groups=4, + spatial_sigma_multiplier=1.0, + ) + for split_unit in ("leave_one_rat_out", "leave_one_session_out") + ) + invalid = replace(passing, mixture_records=decisive_td) + + assert passing.passed is True + assert invalid.passed is False + + +def test_spatial_coordinates_are_required_and_cannot_leak_off_support() -> None: + dataset = _dataset() + missing = np.asarray(dataset.spatial_coordinates).copy() + missing[0, 0] = np.nan + with pytest.raises(ValueError, match="active spatial coordinates"): + replace(dataset, spatial_coordinates=missing).validate() + + +def test_predictor_and_later_outcome_artifacts_are_separate_and_hash_bound( + tmp_path, +) -> None: + dataset = _dataset(signal=True) + directory = tmp_path / "predictors" + manifest = _manifest(directory, dataset) + + frozen = write_spatial_predictor_artifact(directory, dataset, manifest) + loaded = load_spatial_predictor_artifact(directory) + assert loaded.predictor_sha256 == frozen.predictor_sha256 + assert loaded.manifest == manifest + assert loaded.dataset.event_ids == dataset.event_ids + np.testing.assert_allclose(loaded.dataset.log_emissions, dataset.log_emissions) + np.testing.assert_allclose( + loaded.dataset.spatial_coordinates, + dataset.spatial_coordinates, + ) + np.testing.assert_allclose( + loaded.dataset.decoder_point_spread_cm, + dataset.decoder_point_spread_cm, + ) + np.testing.assert_allclose(loaded.dataset.well_masses, dataset.well_masses) + + outcomes = LaterOutcomeTable( + event_ids=dataset.event_ids, + outcome_time_s=dataset.event_end_s + 1.0, + next_well_ids=tuple("well-0" for _ in dataset.event_ids), + valid=np.ones(dataset.n_events, dtype=bool), + ) + written_outcomes = write_later_outcome_artifact( + tmp_path / "outcomes", + loaded, + outcomes, + ) + loaded_outcomes = load_later_outcome_artifact( + tmp_path / "outcomes", + loaded, + ) + assert loaded_outcomes.predictor_sha256 == loaded.predictor_sha256 + assert loaded_outcomes.outcome_sha256 == written_outcomes.outcome_sha256 + assert loaded_outcomes.outcomes.next_well_ids == outcomes.next_well_ids + + too_early = replace( + outcomes, + outcome_time_s=np.asarray(dataset.event_end_s).copy(), + ) + with pytest.raises(ValueError, match="strictly after"): + too_early.validate(dataset) + + +def test_manifest_rejects_noncausal_selection_or_dirty_producer() -> None: + manifest = _manifest() + manifest.validate() + + with pytest.raises(ValueError, match="raw LFP"): + replace( + manifest, + event_selection_schedule="decoder_evidence_top_n", + ).validate() + with pytest.raises(ValueError, match="clean committed worktree"): + replace(manifest, producer_clean_worktree=False).validate() + with pytest.raises(ValueError, match="120-second point-spread"): + replace(manifest, decoder_parameters_sha256="5" * 64).validate() + + +def test_provenance_sidecar_tampering_is_detected(tmp_path) -> None: + dataset = _dataset() + directory = tmp_path / "predictors" + manifest = _manifest(directory, dataset) + write_spatial_predictor_artifact(directory, dataset, manifest) + + report = directory / "replay_spatial_dataset_verification.json" + report.write_bytes(report.read_bytes() + b" ") + with pytest.raises(ValueError, match="verifier report SHA-256"): + load_spatial_predictor_artifact(directory) + + manifest = _manifest(directory, dataset) + write_spatial_predictor_artifact(directory, dataset, manifest) + route = directory / "replay_spatial_route_manifest.json" + route.write_bytes(route.read_bytes() + b" ") + with pytest.raises(ValueError, match="route provenance manifest SHA-256"): + load_spatial_predictor_artifact(directory) + + manifest = _manifest(directory, dataset) + write_spatial_predictor_artifact(directory, dataset, manifest) + audit = directory / "replay_spatial_event_audit.csv" + audit.write_bytes(audit.read_bytes() + b" ") + with pytest.raises(ValueError, match="event audit SHA-256"): + load_spatial_predictor_artifact(directory) + + +def test_predictor_hash_tampering_is_detected(tmp_path) -> None: + dataset = _dataset() + directory = tmp_path / "predictors" + manifest = _manifest(directory, dataset) + write_spatial_predictor_artifact(directory, dataset, manifest) + with (directory / "replay_spatial_predictors.npz").open("ab") as handle: + handle.write(b"tamper") + + with pytest.raises(ValueError, match="SHA-256"): + load_spatial_predictor_artifact(directory)