Files
modeling_zhaocui/math/Q2/train_unaligned_c5.py

159 lines
7.8 KiB
Python

"""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()