"""Run the saved Q2 student on the aligned, unlabeled attachment-3 cases.""" from __future__ import annotations import json import time import csv import numpy as np import torch from crg import INPUT_DIMS, MODALITIES, StructuredGaussianImputer from train import RESULTS, _make_variant, infer_attachment3, reencode_attachment3, validate_attachment3_predictions, write_csv def main() -> None: device = torch.device("cuda" if torch.cuda.is_available() else "cpu") manifest_path = RESULTS / "run_manifest.json" manifest = json.loads(manifest_path.read_text(encoding="utf-8")) calibration = json.loads((RESULTS / "validation_metrics.json").read_text(encoding="utf-8")) selected = calibration.get("selected_model", manifest.get("selected_model")) if not selected: raise ValueError("run_manifest.json does not identify a selected model") imputer = StructuredGaussianImputer(INPUT_DIMS).to(device) imputer_state = torch.load(RESULTS / "structured_imputer.pt", map_location=device, weights_only=True) imputer.load_state_dict(imputer_state) model = _make_variant(selected, imputer).to(device) state = torch.load(RESULTS / "crg_student.pt", map_location=device, weights_only=True) model.load_state_dict(state) with np.load(RESULTS / "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) 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) inference_by_id = {row["case_id"]: row for row in inference_audit} write_csv(RESULTS / "attachment3_predictions.csv", predictions) write_csv(RESULTS / "attachment3_audit.csv", [ {**source, **inference_by_id[source["case_id"]]} for source in source_audit ]) # The training script can finish and persist all labeled-evaluation outputs # before an unlabeled attachment export fails. Reconcile the manifest from # those completed artifacts so the standalone export is safely rerunnable. group_risk_rows = list(csv.DictReader((RESULTS / "group_risk_tuning.csv").open(encoding="utf-8-sig", newline=""))) selected_risk = next((row for row in group_risk_rows if row.get("selected", "").lower() == "true"), None) reliability_rows = list(csv.DictReader((RESULTS / "reliability_hparam_tuning.csv").open(encoding="utf-8-sig", newline=""))) # split_calibration's generic internal names are canonicalized in train.py; # repair artifacts from runs produced before that naming fix as well. for row in group_risk_rows: if row.get("selection_split") == "fit": row["selection_split"] = "reliability_validation" for row in reliability_rows: if row.get("selection_split") == "fit": row["selection_split"] = "reliability_validation" write_csv(RESULTS / "group_risk_tuning.csv", group_risk_rows) write_csv(RESULTS / "reliability_hparam_tuning.csv", reliability_rows) if selected_risk: risk_values = (float(selected_risk["lambda_group"]), float(selected_risk["group_temperature"])) manifest["group_risk_hyperparameters"]["selected"] = list(risk_values) manifest["loss"]["selected_group_risk"] = list(risk_values) manifest["group_risk_hyperparameters"]["selection_split"] = "reliability_validation" manifest["reliability_hyperparameters"]["selected_by_model"] = { row["model"]: [float(row[key]) for key in ("rho_imp", "lambda_u", "lambda_gap", "lambda_span")] for row in reliability_rows if row.get("selected", "").lower() == "true" and (not row.get("risk_candidate_selected") or row["risk_candidate_selected"].lower() == "true") } test_metrics = json.loads((RESULTS / "test_metrics.json").read_text(encoding="utf-8")) manifest["selected_model"] = selected manifest["final_test_metrics"] = test_metrics manifest["calibration"]["temperature"] = temperature manifest["calibration"]["valid_used_for_selection"] = True manifest["calibration"]["test_used_for_selection_or_calibration"] = False manifest["training_configuration"].update({ "student_epoch_limit": 12, "imputer_epochs": 8, "batch_size": 64, "early_stopping_patience": 3, }) manifest["imputer"]["epochs"] = 8 manifest.update({ "completed_utc": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()), "attachment3_cases": len(cases), "attachment3_prediction_file": "attachment3_predictions.csv", "attachment3_audit_file": "attachment3_audit.csv", "attachment3_labeled_metrics": None, "quality_flags": {m: "unavailable; q*=1 fallback for visible rows, unknown flag retained" for m in MODALITIES}, "neutral_output": "exact zero when neutral is the predicted class; no near-zero threshold", }) manifest_path.write_text(json.dumps(manifest, ensure_ascii=False, indent=2), encoding="utf-8") print(f"Wrote {len(predictions)} unlabeled attachment-3 predictions to {RESULTS}", flush=True) if __name__ == "__main__": main()