"""Train the predeclared C5 Q2 architecture on Q1's exploratory index view.""" from __future__ import annotations import argparse import json import time import numpy as np import torch import train as q2_train from crg import INPUT_DIMS, StructuredGaussianImputer from data import ROOT, fit_preprocessor, load_official_splits, transform_split from train import ( RESULTS as ALIGNED_RESULTS, SEED, _fit_neural, _make_variant, assert_group_disjoint, evaluate, fit_imputer, fit_temperature, group_bootstrap, label_resolution_from_train, make_reliability_scenarios, seed_everything, sha256, split_calibration, tune_reliability_hparams, write_csv, ) RESULTS = ALIGNED_RESULTS.parent / "results_unaligned" SOURCE = ROOT / "E题数据" / "附件2-数据集特征文件" / "unaligned_50.pkl" def main() -> None: parser = argparse.ArgumentParser() parser.add_argument("--epochs", type=int, default=12) parser.add_argument("--imputer-epochs", type=int, default=8) parser.add_argument("--batch-size", type=int, default=64) parser.add_argument("--patience", type=int, default=3) parser.add_argument("--bootstrap-repeats", type=int, default=300) parser.add_argument("--seed", type=int, default=SEED) parser.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu") args = parser.parse_args() seed_everything(args.seed) device = torch.device(args.device) RESULTS.mkdir(parents=True, exist_ok=True) official = load_official_splits(SOURCE, version="unaligned_50") overlap = assert_group_disjoint(official) fit, heldout = split_calibration(official["train"], args.seed) reliability_validation, temperature_calibration = split_calibration(heldout, args.seed + 1, fraction=0.5) reliability_validation.name = "reliability_validation" temperature_calibration.name = "temperature_calibration" q2_train.DELTA_U = label_resolution_from_train(fit.regression_y) fitted = fit_preprocessor(fit) transformed = {name: transform_split(split, fitted) for name, split in official.items()} transformed["fit"] = transform_split(fit, fitted) transformed["reliability_validation"] = transform_split(reliability_validation, fitted) transformed["temperature_calibration"] = transform_split(temperature_calibration, fitted) np.savez_compressed(RESULTS / "preprocessor.npz", **{ f"{modality}_{stat}": value for modality, values in fitted.items() for stat, value in values.items() }) scenarios = make_reliability_scenarios(reliability_validation, args.seed + 906) print("official split sizes:", {k: v.n for k, v in official.items()}, flush=True) imputer = StructuredGaussianImputer(INPUT_DIMS).to(device) imputer_history = fit_imputer(imputer, transformed["fit"], fit, device, args.imputer_epochs, args.batch_size, args.seed + 1) torch.save({k: v.detach().cpu() for k, v in imputer.state_dict().items()}, RESULTS / "structured_imputer.pt") model = _make_variant("C5", imputer) model, history = _fit_neural( model, "C5", fit, official["valid"], transformed, device, args.epochs, args.batch_size, args.patience, np.random.default_rng(args.seed + 303), selection_split=reliability_validation, selection_arrays=transformed["reliability_validation"], selection_scenarios=scenarios, ) reliability, tuning_rows = tune_reliability_hparams( model, transformed["reliability_validation"], reliability_validation, scenarios, device, args.batch_size, "C5", seed=args.seed + 551, ) _, calibration_prediction = evaluate(model, transformed["temperature_calibration"], temperature_calibration, device, args.batch_size) temperature = fit_temperature(calibration_prediction["probabilities"], temperature_calibration.class_y) valid_metrics, _ = evaluate(model, transformed["valid"], official["valid"], device, args.batch_size, temperature=temperature) test_metrics, test_prediction = evaluate(model, transformed["test"], official["test"], device, args.batch_size, temperature=temperature) torch.save({k: v.detach().cpu() for k, v in model.state_dict().items()}, RESULTS / "crg_student.pt") (RESULTS / "validation_metrics.json").write_text( json.dumps({**valid_metrics, "model": "C5", "temperature": temperature}, indent=2), encoding="utf-8") (RESULTS / "test_metrics.json").write_text( json.dumps({**test_metrics, "model": "C5", "temperature": temperature}, indent=2), encoding="utf-8") write_csv(RESULTS / "reliability_hparam_tuning.csv", tuning_rows) write_csv(RESULTS / "training_history.csv", imputer_history + history) write_csv(RESULTS / "group_bootstrap_ci.csv", group_bootstrap(official["test"], test_prediction, args.bootstrap_repeats, args.seed + 44)) rows = [] for i, sample_id in enumerate(official["test"].ids): p = test_prediction["probabilities"][i] rows.append({ "sample_id": sample_id, "source_video_id": official["test"].groups[i], "true_class": int(official["test"].class_y[i]), "predicted_class": int(test_prediction["predicted_class"][i]), "true_sentiment": float(official["test"].regression_y[i]), "predicted_sentiment": float(test_prediction["predicted_score"][i]), "p_negative": float(p[0]), "p_neutral": float(p[1]), "p_positive": float(p[2]), }) write_csv(RESULTS / "test_predictions.csv", rows) manifest = { "scope": "exploratory unaligned_50 relative-index projection and prespecified C5 training", "physical_time_alignment": False, "input": str(SOURCE.relative_to(ROOT)), "input_sha256": sha256(SOURCE), "text_encoder": "official precomputed text field; revision not supplied", "q1_adapter_audit": {name: split.alignment_audit for name, split in official.items()}, "official_group_overlap": overlap, "internal_splits": {"fit": fit.n, "reliability_validation": reliability_validation.n, "temperature_calibration": temperature_calibration.n}, "model": "C5 fixed before this run; no unaligned architecture selection", "quality": "no quality scores in official file; q*=1 for visible rows and J_Q=0", "seed": args.seed, "device": str(device), "device_name": torch.cuda.get_device_name(device) if device.type == "cuda" else "CPU", "torch_version": torch.__version__, "epochs_limit": args.epochs, "trained_c5_epochs": len(history), "selected_c5_epoch": int(min(history, key=lambda row: row["inner_selection_nll"])["epoch"]), "imputer_epochs": args.imputer_epochs, "batch_size": args.batch_size, "patience": args.patience, "selected_reliability": reliability, "temperature": temperature, "bootstrap_repeats": args.bootstrap_repeats, "attachment3": "not inferred: unaligned files lack numerical text and trusted lengths", "validation_metrics": valid_metrics, "test_metrics": test_metrics, "completed_utc": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()), } (RESULTS / "run_manifest.json").write_text(json.dumps(manifest, ensure_ascii=False, indent=2), encoding="utf-8") stale_teacher = RESULTS / "teacher.pt" if stale_teacher.exists(): stale_teacher.unlink() print("C5 unaligned complete:", json.dumps({ "accuracy": test_metrics["accuracy"], "macro_f1": test_metrics["macro_f1"], "mae": test_metrics["regression_mae"], "temperature": temperature, }), flush=True) if __name__ == "__main__": main()