"""Run a saved Q2 model on the unlabeled Attachment 3 cases.""" from __future__ import annotations import argparse import json import time from pathlib import Path import numpy as np import torch from data_paths import PROJECT_ROOT from model.crg import INPUT_DIMS, MODALITIES, StructuredGaussianImputer from .train import ( _make_variant, infer_attachment3, reencode_attachment3, validate_attachment3_predictions, write_csv, ) DEFAULT_RESULTS_DIR = PROJECT_ROOT / "experiments" / "q2" / "unaligned_math_all_b128" DEFAULT_OUTPUT_DIR = PROJECT_ROOT / "output" / "q2" def main() -> None: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--input-version", choices=("aligned_50", "unaligned_50"), default="unaligned_50") parser.add_argument("--results-dir", type=Path, default=DEFAULT_RESULTS_DIR, help="saved Q2 checkpoint and calibration directory") parser.add_argument("--output-dir", type=Path, default=DEFAULT_OUTPUT_DIR) 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) results_dir = args.results_dir.expanduser().resolve() output_dir = args.output_dir.expanduser().resolve() manifest_path = results_dir / "run_manifest.json" calibration_path = results_dir / "validation_metrics.json" manifest = json.loads(manifest_path.read_text(encoding="utf-8")) calibration = json.loads(calibration_path.read_text(encoding="utf-8")) selected = calibration.get("selected_model", manifest.get("selected_model")) if not selected: raise ValueError(f"no selected_model recorded in {calibration_path}") imputer = StructuredGaussianImputer(INPUT_DIMS).to(device) imputer.load_state_dict(torch.load(results_dir / "structured_imputer.pt", map_location=device, weights_only=True)) model = _make_variant(selected, imputer).to(device) model.load_state_dict(torch.load(results_dir / "crg_student.pt", map_location=device, weights_only=True)) with np.load(results_dir / "preprocessor.npz", allow_pickle=False) as archive: fitted = {m: {k: archive[f"{m}_{k}"].copy() for k in ("mean", "std")} for m in MODALITIES} priors = manifest["attachment3_low_information_priors"] temperature = float(calibration["temperature"]) class_prior = np.asarray(priors["class_probability_values"], dtype=np.float64) magnitude_priors = np.asarray((priors["negative_beta"], priors["positive_beta"]), dtype=np.float32) cases, source_audit = reencode_attachment3(device, input_version=args.input_version) predictions, inference_audit = infer_attachment3( model, cases, fitted, device, temperature, class_prior, magnitude_priors, ) validate_attachment3_predictions([case["case_id"] for case in cases], predictions) output_dir.mkdir(parents=True, exist_ok=True) inference_by_id = {row["case_id"]: row for row in inference_audit} predictions_path = output_dir / "attachment3_predictions.csv" audit_path = output_dir / "attachment3_audit.csv" manifest_out_path = output_dir / "attachment3_prediction_manifest.json" write_csv(predictions_path, predictions) write_csv(audit_path, [{**source, **inference_by_id[source["case_id"]]} for source in source_audit]) try: results_reference = results_dir.relative_to(PROJECT_ROOT).as_posix() except ValueError: results_reference = "external checkpoint directory" prediction_manifest = { "task": "unlabeled Attachment 3 inference", "input_version": args.input_version, "selected_model": selected, "checkpoint_run": results_reference, "prediction_count": len(predictions), "temperature": temperature, "labels_available": False, "prediction_file": predictions_path.name, "audit_file": audit_path.name, "completed_utc": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()), } manifest_out_path.write_text(json.dumps(prediction_manifest, ensure_ascii=False, indent=2), encoding="utf-8") print(f"Wrote {len(predictions)} unlabeled Attachment 3 predictions to {output_dir}", flush=True) if __name__ == "__main__": main()