87 lines
3.7 KiB
Python
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()
|