Files

87 lines
3.7 KiB
Python

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()