benchdash / validate_data.py
josephsoo's picture
Align Space with manuscript figures
77d097c
Raw
History Blame
24.5 kB
#!/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())