#!/usr/bin/env python3 """Validate the bundled BEND-BCI Space tables and their paper provenance. The local checks run in the standalone Hugging Face Space. Passing ``--canonical-root`` additionally compares every manuscript-facing summary against ``paper/results`` in the main benchmark repository. """ from __future__ import annotations import argparse import hashlib import json from pathlib import Path from typing import Iterable import numpy as np import pandas as pd DATASETS = { "monkey", "allen_neuropixels", "speech", "mc_pacman", "ratinabox", } CONSISTENCY_DATASETS = DATASETS - {"mc_pacman"} LATENT_COORDINATE_SPACE = "per_session_whitened_reference_aligned_3d" FIGURE5_TARGET_SESSION = "sub-C_ses-CO-20150716_behavior+ecephys" # Coverage is defined in manuscript v7 Supplementary Table 2. EXPECTED_COVERAGE = { "clean_prediction_summary.csv": (115, 112), "robustness_summary.csv": (115, 112), "scalability_summary.csv": (115, 112), "consistency_summary.csv": (54, 46), "neuron_shap_summary.csv": (105, 105), "trial_shapley_summary.csv": (81, 81), "trial_shapley_retrain_summary.csv": (99, 99), } REQUIRED_COLUMNS = { "clean_prediction_summary.csv": { "model", "dataset", "status", "metric", "score", "decoder", }, "robustness_summary.csv": { "model", "dataset", "status", "metric", "noise_levels", "scores", "raw_auc", }, "scalability_summary.csv": { "model", "dataset", "status", "training_time_sec", "inference_time_sec", "peak_ram_gb", "peak_vram_gb", }, "consistency_summary.csv": { "model", "dataset", "is_active_model", "mean_r2", "n_sessions", "sessions", "latent_dim", "scoring_modes", "normalizations", }, "neuron_shap_summary.csv": { "model", "dataset", "is_active_model", "auc", "spearman_corr", "shap_mean_value", "shap_min_value", "shap_max_value", "shap_fraction_positive", "shap_fraction_negative", }, "trial_shapley_summary.csv": { "model", "dataset", "is_active_model", "analysis", "perturbation_auc", "rotation_angle_deg", "rotation_subspace_dim_spec", "trial_selection_mode", "converged", "shapley_mean_value", "shapley_min_value", "shapley_max_value", "shapley_fraction_positive", "shapley_fraction_negative", }, "trial_shapley_retrain_summary.csv": { "analysis", "model", "is_active_model", "condition", "metric", "score", }, "trial_historical_trajectories.csv": { "model", "target_session", "trial_index", "trial_id", "direction_index", "direction_label", "time_index", "target_x", "target_y", "current_only_x", "current_only_y", "historical_selected_x", "historical_selected_y", "current_only_r2", "historical_selected_r2", }, "latent_samples.csv": { "model", "dataset", "session", "session_label", "x", "y", "z", "condition", "trial_index", "time_index", "eval_time_index", "coordinate_space", "reference_session", "alignment", "landmark_type", "n_alignment_landmarks", "is_reference", "session_order", }, "latent_trajectories.csv": { "model", "dataset", "session", "session_label", "x", "y", "z", "condition", "time_index", "eval_time_index", "n_points", "coordinate_space", "reference_session", "alignment", "landmark_type", "n_alignment_landmarks", "is_reference", "session_order", }, } UNIQUE_KEYS = { "clean_prediction_summary.csv": ["model", "dataset"], "robustness_summary.csv": ["model", "dataset"], "scalability_summary.csv": ["model", "dataset"], "consistency_summary.csv": ["model", "dataset"], "neuron_shap_summary.csv": ["model", "dataset"], "trial_shapley_summary.csv": ["model", "dataset"], "trial_shapley_retrain_summary.csv": ["analysis", "model", "condition"], } CANONICAL_NAMES = { "clean_prediction_summary.csv": "metrics_summary.csv", "robustness_summary.csv": "robustness_summary.csv", "scalability_summary.csv": "scalability_summary.csv", "consistency_summary.csv": "consistency_summary.csv", "neuron_shap_summary.csv": "neuron_shap_summary.csv", "trial_shapley_summary.csv": "trial_shapley_summary.csv", "trial_shapley_retrain_summary.csv": "trial_shapley_retrain_summary.csv", } class ValidationError(RuntimeError): """Raised when Space data violates its manuscript-facing contract.""" 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 _active_mask(frame: pd.DataFrame) -> pd.Series: if "status" in frame.columns: return frame["status"].fillna("").eq("present") if "is_active_model" in frame.columns: return frame["is_active_model"].astype(str).str.lower().eq("true") return pd.Series(True, index=frame.index) def _require(condition: bool, message: str, errors: list[str]) -> None: if not condition: errors.append(message) def _same_values(left: pd.DataFrame, right: pd.DataFrame) -> None: pd.testing.assert_frame_equal( left.reset_index(drop=True), right.reset_index(drop=True), check_dtype=False, check_exact=True, check_categorical=False, ) def _validate_latent_cell( frame: pd.DataFrame, *, table_label: str, model: str, dataset: str, sessions: list[str], landmark_type: str, errors: list[str], ) -> None: """Validate reference/alignment metadata for one method-dataset cell.""" cell = frame[ frame["model"].astype(str).eq(model) & frame["dataset"].astype(str).eq(dataset) ] for session_order, session in enumerate(sessions): session_rows = cell[cell["session"].astype(str).eq(session)] if session_rows.empty: errors.append(f"{table_label}: missing {model}/{dataset}/{session}") continue expected_reference = "true" if session_order == 0 else "false" observed_reference = set( session_rows["is_reference"].dropna().astype(str).str.lower() ) _require( observed_reference == {expected_reference}, f"{table_label}: {model}/{dataset}/{session} is_reference " f"values {sorted(observed_reference)}", errors, ) expected_alignment = "identity" if session_order == 0 else "proper_similarity_procrustes" observed_alignment = set(session_rows["alignment"].dropna().astype(str)) _require( observed_alignment == {expected_alignment}, f"{table_label}: {model}/{dataset}/{session} alignment " f"values {sorted(observed_alignment)}", errors, ) observed_reference_sessions = set( session_rows["reference_session"].dropna().astype(str) ) _require( observed_reference_sessions == {sessions[0]}, f"{table_label}: {model}/{dataset}/{session} reference metadata " f"{sorted(observed_reference_sessions)}", errors, ) observed_landmarks = set(session_rows["landmark_type"].dropna().astype(str)) _require( observed_landmarks == {landmark_type}, f"{table_label}: {model}/{dataset}/{session} landmarks " f"{sorted(observed_landmarks)}", errors, ) orders = pd.to_numeric(session_rows["session_order"], errors="coerce") _require( orders.notna().all() and np.isfinite(orders.to_numpy()).all() and orders.eq(session_order).all(), f"{table_label}: {model}/{dataset}/{session} has invalid session_order", errors, ) landmark_counts = pd.to_numeric( session_rows["n_alignment_landmarks"], errors="coerce" ) _require( landmark_counts.notna().all() and np.isfinite(landmark_counts.to_numpy()).all() and landmark_counts.ge(3).all(), f"{table_label}: {model}/{dataset}/{session} has invalid landmark count", errors, ) def validate_local(data_dir: Path) -> dict[str, pd.DataFrame]: """Validate schemas, coverage, uniqueness, and fixed analysis conventions.""" frames: dict[str, pd.DataFrame] = {} errors: list[str] = [] for name, columns in REQUIRED_COLUMNS.items(): path = data_dir / name if not path.exists(): errors.append(f"missing required table: {path}") continue frame = pd.read_csv(path) missing = sorted(columns - set(frame.columns)) _require(not missing, f"{name}: missing columns {missing}", errors) if missing: continue frames[name] = frame manifest_path = data_dir / "release_manifest.json" if not manifest_path.exists(): errors.append(f"missing release manifest: {manifest_path}") else: try: manifest = json.loads(manifest_path.read_text(encoding="utf-8")) except (json.JSONDecodeError, OSError) as exc: errors.append(f"invalid release manifest: {exc}") manifest = {} if not isinstance(manifest, dict): errors.append("invalid release manifest: top-level value must be an object") manifest = {} raw_manifest_spaces = manifest.get("latent_coordinate_spaces", []) if not isinstance(raw_manifest_spaces, list): errors.append("release manifest: latent_coordinate_spaces must be a list") raw_manifest_spaces = [] elif not all(isinstance(value, str) for value in raw_manifest_spaces): errors.append("release manifest: coordinate-space values must be strings") raw_manifest_spaces = [] manifest_spaces = set(raw_manifest_spaces) _require( manifest_spaces == {LATENT_COORDINATE_SPACE}, f"release manifest: coordinate spaces {sorted(manifest_spaces)}", errors, ) manifest_tables = manifest.get("tables", {}) if not isinstance(manifest_tables, dict): errors.append("release manifest: tables must be an object") manifest_tables = {} for name in REQUIRED_COLUMNS: entry = manifest_tables.get(name) if not isinstance(entry, dict): errors.append(f"release manifest: missing table entry {name}") continue path = data_dir / name if not path.exists(): continue expected_hash = entry.get("sha256") actual_hash = _sha256(path) _require( expected_hash == actual_hash, f"release manifest: SHA-256 mismatch for {name}", errors, ) if name in frames: _require( entry.get("rows") == len(frames[name]), f"release manifest: row-count mismatch for {name}", errors, ) for name, (n_rows, n_active) in EXPECTED_COVERAGE.items(): if name not in frames: continue frame = frames[name] _require(len(frame) == n_rows, f"{name}: expected {n_rows} rows, found {len(frame)}", errors) active = int(_active_mask(frame).sum()) _require(active == n_active, f"{name}: expected {n_active} available rows, found {active}", errors) for name, keys in UNIQUE_KEYS.items(): if name not in frames or not set(keys).issubset(frames[name].columns): continue duplicates = frames[name].duplicated(keys, keep=False) _require(not duplicates.any(), f"{name}: duplicate keys for {keys}", errors) for name in ("clean_prediction_summary.csv", "robustness_summary.csv", "scalability_summary.csv"): if name in frames: observed = set(frames[name]["dataset"].dropna().astype(str)) _require(observed == DATASETS, f"{name}: dataset set is {sorted(observed)}", errors) if "clean_prediction_summary.csv" in frames: prediction = frames["clean_prediction_summary.csv"] _require(prediction["model"].nunique() == 23, "prediction: expected 23 methods", errors) metrics = set(prediction.loc[_active_mask(prediction), "metric"].dropna()) _require(metrics == {"accuracy", "r2"}, f"prediction: unexpected metrics {sorted(metrics)}", errors) if "consistency_summary.csv" in frames: consistency = frames["consistency_summary.csv"] active = consistency.loc[_active_mask(consistency)] observed = set(active["dataset"].dropna().astype(str)) _require(observed == CONSISTENCY_DATASETS, f"consistency: dataset set is {sorted(observed)}", errors) norms = set(active["normalizations"].dropna().astype(str)) _require(norms == {"per_session_whitening"}, f"consistency: unexpected normalization {sorted(norms)}", errors) if "neuron_shap_summary.csv" in frames: feature = frames["neuron_shap_summary.csv"] allen = feature[feature["dataset"].eq("allen_neuropixels")] _require(not allen.empty and allen["spearman_corr"].notna().all(), "feature attribution: Allen Spearman values missing", errors) other = feature[~feature["dataset"].eq("allen_neuropixels")] _require(other["auc"].notna().all(), "feature attribution: ROC-AUC values missing", errors) _require(feature["shap_min_value"].lt(0).any(), "feature attribution: signed negative values absent", errors) if "trial_shapley_summary.csv" in frames: trial = frames["trial_shapley_summary.csv"] _require(set(trial["analysis"].dropna()) == {"subspace_rotation"}, "trial valuation: noncanonical analysis", errors) angles = set(pd.to_numeric(trial["rotation_angle_deg"], errors="coerce").dropna()) _require(angles == {75.0}, f"trial valuation: rotation angles {sorted(angles)}", errors) dims = set(trial["rotation_subspace_dim_spec"].dropna().astype(str)) _require(dims == {"full"}, f"trial valuation: subspace specs {sorted(dims)}", errors) modes = set(trial["trial_selection_mode"].dropna().astype(str)) _require(modes == {"random"}, f"trial valuation: selection modes {sorted(modes)}", errors) _require(trial["perturbation_auc"].notna().all(), "trial valuation: detection AUC missing", errors) _require(trial["shapley_min_value"].lt(0).any(), "trial valuation: signed negative values absent", errors) if "trial_shapley_retrain_summary.csv" in frames: retrain = frames["trial_shapley_retrain_summary.csv"] expected = { "within_session_cleaning": {"mixed_full", "data_shapley", "oracle"}, "cross_session_old_trial_selection": { "target_only", "all_sessions", "oldonly_dshap_negative_removal", }, } observed = { analysis: set(group["condition"].dropna().astype(str)) for analysis, group in retrain.groupby("analysis") } _require(observed == expected, f"trial retraining: conditions {observed}", errors) if "trial_historical_trajectories.csv" in frames: historical = frames["trial_historical_trajectories.csv"] _require( len(historical) == 135 * 35, "historical trajectories: expected 135 trials × 35 time bins", errors, ) keys = ["trial_index", "time_index"] _require( not historical.duplicated(keys, keep=False).any(), f"historical trajectories: duplicate keys for {keys}", errors, ) _require( set(historical["model"].dropna().astype(str)) == {"rnn"}, "historical trajectories: expected the Figure 5e RNN example", errors, ) _require( set(historical["target_session"].dropna().astype(str)) == {FIGURE5_TARGET_SESSION}, "historical trajectories: unexpected target session", errors, ) _require( historical["trial_index"].nunique() == 135 and historical["trial_id"].nunique() == 135, "historical trajectories: expected 135 held-out trials and trial IDs", errors, ) time_counts = historical.groupby("trial_index")["time_index"].nunique() _require( len(time_counts) == 135 and time_counts.eq(35).all(), "historical trajectories: every trial must contain 35 time bins", errors, ) direction_indices = pd.to_numeric( historical["direction_index"], errors="coerce" ) _require( direction_indices.notna().all() and set(direction_indices.astype(int)) == set(range(8)), "historical trajectories: expected all eight reach directions", errors, ) for column in ( "target_x", "target_y", "current_only_x", "current_only_y", "historical_selected_x", "historical_selected_y", "current_only_r2", "historical_selected_r2", ): values = pd.to_numeric(historical[column], errors="coerce") _require( values.notna().all() and np.isfinite(values.to_numpy()).all(), f"historical trajectories: non-finite {column} values", errors, ) _require( historical["current_only_r2"].nunique() == 1 and historical["historical_selected_r2"].nunique() == 1, "historical trajectories: expected one R² value per condition", errors, ) for name in ("latent_samples.csv", "latent_trajectories.csv"): if name not in frames: continue latent = frames[name] observed = set(latent["dataset"].dropna().astype(str)) _require(observed.issubset(CONSISTENCY_DATASETS), f"{name}: unsupported datasets {sorted(observed - CONSISTENCY_DATASETS)}", errors) for column in ("x", "y", "z"): values = pd.to_numeric(latent[column], errors="coerce") _require( values.notna().all() and np.isfinite(values.to_numpy()).all(), f"{name}: non-finite {column} values", errors, ) for column in ( "coordinate_space", "reference_session", "alignment", "landmark_type", "n_alignment_landmarks", "is_reference", "session_order", ): _require( latent[column].notna().all(), f"{name}: missing {column} values", errors, ) spaces = set(latent["coordinate_space"].dropna().astype(str)) _require( spaces == {LATENT_COORDINATE_SPACE}, f"{name}: coordinate spaces {sorted(spaces)}", errors, ) if { "consistency_summary.csv", "latent_samples.csv", "latent_trajectories.csv" }.issubset(frames): consistency = frames["consistency_summary.csv"] consistency = consistency.loc[_active_mask(consistency)].copy() samples = frames["latent_samples.csv"] trajectories = frames["latent_trajectories.csv"] expected_sample_sessions: set[tuple[str, str, str]] = set() expected_trajectory_sessions: set[tuple[str, str, str]] = set() for row in consistency.itertuples(index=False): sessions = ( [] if pd.isna(row.sessions) else [item.strip() for item in str(row.sessions).split(";") if item.strip()] ) if not sessions: errors.append(f"consistency: {row.model}/{row.dataset} has no sessions") continue try: expected_n_sessions = int(row.n_sessions) except (TypeError, ValueError): expected_n_sessions = -1 _require( len(sessions) == expected_n_sessions, f"consistency: {row.model}/{row.dataset} lists {len(sessions)} " f"sessions but n_sessions={row.n_sessions}", errors, ) expected_sample_sessions.update( (str(row.model), str(row.dataset), session) for session in sessions ) if str(row.dataset) != "ratinabox": expected_trajectory_sessions.update( (str(row.model), str(row.dataset), session) for session in sessions ) _validate_latent_cell( samples, table_label="latent samples", model=str(row.model), dataset=str(row.dataset), sessions=sessions, landmark_type=str(row.scoring_modes), errors=errors, ) if str(row.dataset) != "ratinabox": _validate_latent_cell( trajectories, table_label="latent trajectories", model=str(row.model), dataset=str(row.dataset), sessions=sessions, landmark_type=str(row.scoring_modes), errors=errors, ) observed_sample_sessions = set( samples[["model", "dataset", "session"]].astype(str).itertuples(index=False, name=None) ) observed_trajectory_sessions = set( trajectories[["model", "dataset", "session"]].astype(str).itertuples(index=False, name=None) ) _require( observed_sample_sessions == expected_sample_sessions, "latent samples: active consistency session coverage differs", errors, ) _require( observed_trajectory_sessions == expected_trajectory_sessions, "latent trajectories: applicable consistency session coverage differs", errors, ) _require( len(observed_sample_sessions) == 173, f"latent samples: expected 173 sessions, found {len(observed_sample_sessions)}", errors, ) if errors: raise ValidationError("\n".join(f"- {item}" for item in errors)) return frames def validate_canonical(data_dir: Path, canonical_root: Path) -> None: """Require Space summaries to equal the current paper result tables.""" results_dir = canonical_root / "paper" / "results" errors: list[str] = [] for dashboard_name, paper_name in CANONICAL_NAMES.items(): dashboard_path = data_dir / dashboard_name paper_path = results_dir / paper_name if not dashboard_path.exists(): errors.append(f"missing dashboard table: {dashboard_path}") continue if not paper_path.exists(): errors.append(f"missing canonical table: {paper_path}") continue try: _same_values(pd.read_csv(dashboard_path), pd.read_csv(paper_path)) except AssertionError as exc: first_line = str(exc).splitlines()[0] if str(exc) else "values differ" errors.append(f"{dashboard_name} != {paper_name}: {first_line}") if errors: raise ValidationError("\n".join(f"- {item}" for item in errors)) def build_parser() -> argparse.ArgumentParser: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--data-dir", type=Path, default=Path(__file__).resolve().parent / "data") parser.add_argument( "--canonical-root", type=Path, help="Main benchmark repository root; enables exact paper/results comparisons.", ) return parser def main(argv: Iterable[str] | None = None) -> int: args = build_parser().parse_args(argv) frames = validate_local(args.data_dir) if args.canonical_root is not None: validate_canonical(args.data_dir, args.canonical_root.resolve()) print(f"Validated {len(frames)} BEND-BCI Space tables in {args.data_dir}") if args.canonical_root is not None: print("Canonical paper/results comparison passed") return 0 if __name__ == "__main__": raise SystemExit(main())