from __future__ import annotations import argparse import csv from pathlib import Path from typing import Any import numpy as np import torch from data_paths import PROJECT_ROOT from q2.deep_learning.q2.data import RobustStats from ..run_experiments import _read_attachment4 from .evaluate import ( MODEL_SEEDS, _attachment_predictions_and_explanations, _attachment_split, _load_ensemble, ) from .owen import hierarchical_owen_one SCALER_PATH = PROJECT_ROOT / "experiments" / "q2" / "unaligned_deep_two_b128" / "unaligned_50_robust_stats.npz" DEFAULT_OUTPUT = PROJECT_ROOT / "output" / "q3" / "ati_ho" def _write_csv(path: Path, rows: list[dict[str, Any]]) -> None: if not rows: raise ValueError(f"no rows to write: {path}") path.parent.mkdir(parents=True, exist_ok=True) with path.open("w", encoding="utf-8-sig", newline="") as stream: writer = csv.DictWriter(stream, fieldnames=list(rows[0])) writer.writeheader() writer.writerows(rows) def main() -> None: parser = argparse.ArgumentParser(description="Regenerate the official Q3 Attachment 4 predictions and explanations.") parser.add_argument("--output-dir", type=Path, default=DEFAULT_OUTPUT) parser.add_argument("--device", choices=("auto", "cpu", "cuda"), default="auto") args = parser.parse_args() if args.device == "cuda" and not torch.cuda.is_available(): parser.error("CUDA was requested but is not available") device_name = "cuda" if args.device == "auto" and torch.cuda.is_available() else args.device if device_name == "auto": device_name = "cpu" device = torch.device(device_name) cases, _ = _read_attachment4("unaligned_50") stats = RobustStats.load(SCALER_PATH) attachment = _attachment_split(cases, stats) dims = tuple(int(values.shape[-1]) for values in attachment.x) models = _load_ensemble("A0", MODEL_SEEDS, dims, device) predictions, explanations, _, _ = _attachment_predictions_and_explanations( "A0", models, cases, attachment, device ) local_rows: list[dict[str, Any]] = [] for index, case in enumerate(cases): xs = tuple(torch.as_tensor(values[index:index + 1], dtype=torch.float32, device=device) for values in attachment.x) mask = torch.as_tensor(attachment.mask[index:index + 1], dtype=torch.bool, device=device) result = hierarchical_owen_one( models, xs, mask, seed=20260926 + index, start_permutations=8, max_permutations=64 ) for modality, label in enumerate(("T", "A", "V")): for bin_index, (left, right) in enumerate(result["bin_slices"]): local_rows.append({ "case_id": case["case_id"], "modality": label, "relative_bin": bin_index, "relative_position_start": left / 50.0, "relative_position_end": right / 50.0, "local_owen_margin_contribution": float(result["contribution"][modality, bin_index]), "owen_standard_error": float(result["standard_error"][modality, bin_index]), "permutations": result["permutations"], "stopping_status": result["stopping_status"], "physical_time_alignment": False, }) _write_csv(args.output_dir / "attachment4_predictions.csv", predictions) _write_csv(args.output_dir / "attachment4_explanations.csv", explanations) _write_csv(args.output_dir / "attachment4_local_evidence.csv", local_rows) print(f"Wrote {len(predictions)} predictions, {len(explanations)} explanations, and {len(local_rows)} local-evidence rows to {args.output_dir}") if __name__ == "__main__": main()