Train ATI-HO and finalize project outputs
This commit is contained in:
@@ -1,102 +1,98 @@
|
||||
"""Run the saved Q2 student on the aligned, unlabeled attachment-3 cases."""
|
||||
"""Run a saved Q2 model on the unlabeled Attachment 3 cases."""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import time
|
||||
import csv
|
||||
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 RESULTS, _make_variant, infer_attachment3, reencode_attachment3, validate_attachment3_predictions, write_csv
|
||||
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:
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
manifest_path = RESULTS / "run_manifest.json"
|
||||
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((RESULTS / "validation_metrics.json").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("run_manifest.json does not identify a selected model")
|
||||
raise ValueError(f"no selected_model recorded in {calibration_path}")
|
||||
|
||||
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)
|
||||
imputer.load_state_dict(torch.load(results_dir / "structured_imputer.pt", map_location=device, weights_only=True))
|
||||
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)
|
||||
model.load_state_dict(torch.load(results_dir / "crg_student.pt", map_location=device, weights_only=True))
|
||||
|
||||
with np.load(RESULTS / "preprocessor.npz", allow_pickle=False) as archive:
|
||||
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)
|
||||
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)
|
||||
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({
|
||||
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()),
|
||||
"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)
|
||||
}
|
||||
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__":
|
||||
|
||||
Reference in New Issue
Block a user