Files
modeling_zhaocui/deep_learning/Q2/q2/train_math_protocol.py
T

514 lines
22 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Retrain the two maintained Q2 models under the math/Q2 V2 protocol.
The model architectures and joint CE + SmoothL1 objective stay unchanged.
Training masks, official splits, validation scenarios, and final-test handling
follow the corresponding math/Q2 protocol where those choices apply.
"""
from __future__ import annotations
import argparse
import csv
import hashlib
import json
import math
import random
import time
from collections import Counter, defaultdict
from pathlib import Path
from typing import Any
import numpy as np
import torch
import torch.nn.functional as F
from sklearn.metrics import accuracy_score, f1_score, mean_absolute_error, mean_squared_error
from torch import nn
from .data import ATTACHMENT2, RobustStats, Split, apply_robust_stats, fit_robust_stats
from .evaluate_math_protocol import (
AURC_BOOTSTRAP_SEED,
BOOTSTRAP_REPS,
CURVE_MODES,
METHODS,
SCENARIO_SEED,
TEST_BOOTSTRAP_SEED,
actual_additional_rates,
aurc_from_curve,
continuous_mask,
curve_scenarios,
load_splits,
make_scenarios,
metrics,
scenario_seed,
sha256,
write_csv,
)
from .models import AlignedFusionModel
from .mofe import MixtureOfFusionExperts
from .train_mofe import EARLYCONCAT, MODEL_CONFIG, MOFE7_MLP, _predict
from .train_compare import _loss, seed_everything
Q2_ROOT = Path(__file__).resolve().parents[1]
OUTPUT_DIR = Q2_ROOT / "outputs" / "followups" / "R03_math_protocol_retraining"
SEED = 20260924
TRAIN_MASK_SEED = 20261227
BATCH_SIZE = 64
EPOCH_LIMIT = 12
PATIENCE = 3
LEARNING_RATE = 3e-4
WEIGHT_DECAY = 1e-3
SELECTION_SCENARIOS = ("0.0/none", "0.3/single", "0.3/sync", "0.5/async")
TRAIN_RATES = (0.0, 0.1, 0.3, 0.5, 0.7)
TRAIN_MODES = ("single", "sync", "partial", "async")
def device_for(name: str) -> torch.device:
if name == "auto":
return torch.device("cuda" if torch.cuda.is_available() else "cpu")
return torch.device(name)
def set_deterministic(seed: int) -> None:
seed_everything(seed)
torch.set_num_threads(4)
torch.backends.cudnn.deterministic = True
torch.backends.cudnn.benchmark = False
def build_model(method: str, dims: tuple[int, int, int], device: torch.device) -> nn.Module:
if method == EARLYCONCAT:
return AlignedFusionModel("concat", dims=dims).to(device)
if method == MOFE7_MLP:
return MixtureOfFusionExperts(dims=dims, **MODEL_CONFIG).to(device)
raise ValueError(f"unknown method: {method}")
def model_state(model: nn.Module, method: str) -> dict[str, Any]:
state: dict[str, Any] = {
"method": method,
"dims": tuple(int(x) for x in model_dims(model)),
"state_dict": model.state_dict(),
"seed": SEED,
"protocol": "math/Q2 V2 adapted deterministic-model training",
}
if method == EARLYCONCAT:
state["kind"] = "concat"
else:
state["config"] = MODEL_CONFIG
return state
def model_dims(model: nn.Module) -> tuple[int, int, int]:
if isinstance(model, AlignedFusionModel):
return tuple(layer[0].in_features for layer in model.projections) # type: ignore[return-value]
if isinstance(model, MixtureOfFusionExperts):
return tuple(layer[0].in_features for layer in model.private_projections) # type: ignore[return-value]
raise TypeError(type(model))
def train_masks_for_epoch(split: Split, epoch: int) -> tuple[np.ndarray, Counter[str]]:
"""Sample reproducible math-protocol rates/patterns per training example."""
rows: list[np.ndarray] = []
counts: Counter[str] = Counter()
for sample_id, observed in zip(split.ids, split.mask):
rng = np.random.default_rng(scenario_seed(TRAIN_MASK_SEED + SEED, sample_id, f"train/{epoch}"))
rate = float(rng.choice(TRAIN_RATES))
mode = str(rng.choice(TRAIN_MODES))
key = f"{rate:.1f}/{mode}"
counts[key] += 1
row = continuous_mask(observed, rate, mode, rng)
rows.append(row)
return np.stack(rows), counts
def _batched_loss(
model: nn.Module,
split: Split,
masks: np.ndarray,
device: torch.device,
batch_size: int,
) -> float:
model.eval()
losses: list[float] = []
weights: list[int] = []
with torch.inference_mode():
for start in range(0, split.n, batch_size):
end = min(start + batch_size, split.n)
xs = tuple(torch.as_tensor(x[start:end], dtype=torch.float32, device=device) for x in split.x)
mb = torch.as_tensor(masks[start:end], dtype=torch.bool, device=device)
y_cls = torch.as_tensor(split.y_cls[start:end], dtype=torch.long, device=device)
y_reg = torch.as_tensor(split.y_reg[start:end], dtype=torch.float32, device=device)
losses.append(float(_loss(model(xs, mb), y_cls, y_reg).item()))
weights.append(end - start)
return float(np.average(losses, weights=weights))
def selection_loss(model: nn.Module, valid: Split, scenarios: dict[str, np.ndarray], device: torch.device) -> float:
return float(np.mean([
_batched_loss(model, valid, scenarios[key], device, BATCH_SIZE)
for key in SELECTION_SCENARIOS
]))
def train_one(
method: str,
train: Split,
valid: Split,
valid_scenarios: dict[str, np.ndarray],
orders: list[np.ndarray],
output_dir: Path,
device: torch.device,
) -> tuple[nn.Module, int, list[dict[str, Any]], Counter[str]]:
set_deterministic(SEED)
model = build_model(method, tuple(x.shape[-1] for x in train.x), device)
optimizer = torch.optim.AdamW(model.parameters(), lr=LEARNING_RATE, weight_decay=WEIGHT_DECAY)
xs = tuple(torch.as_tensor(x, dtype=torch.float32, device=device) for x in train.x)
y_cls = torch.as_tensor(train.y_cls, dtype=torch.long, device=device)
y_reg = torch.as_tensor(train.y_reg, dtype=torch.float32, device=device)
checkpoint_path = output_dir / "model_best.pt"
history: list[dict[str, Any]] = []
train_mask_counts: Counter[str] = Counter()
best_loss = math.inf
best_epoch = 0
stale = 0
for epoch in range(1, EPOCH_LIMIT + 1):
model.train()
epoch_masks, epoch_counts = train_masks_for_epoch(train, epoch)
train_mask_counts.update(epoch_counts)
batch_losses: list[float] = []
order = orders[epoch - 1]
for start in range(0, train.n, BATCH_SIZE):
indices_np = order[start:start + BATCH_SIZE]
indices = torch.as_tensor(indices_np, dtype=torch.long, device=device)
mb = torch.as_tensor(epoch_masks[indices_np], dtype=torch.bool, device=device)
output = model(tuple(x.index_select(0, indices) for x in xs), mb)
loss = _loss(output, y_cls.index_select(0, indices), y_reg.index_select(0, indices))
optimizer.zero_grad(set_to_none=True)
loss.backward()
nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
optimizer.step()
batch_losses.append(float(loss.detach().item()))
valid_selection_loss = selection_loss(model, valid, valid_scenarios, device)
row = {
"method": method,
"seed": SEED,
"epoch": epoch,
"train_loss": float(np.mean(batch_losses)),
"valid_selection_loss": valid_selection_loss,
"valid_clean_loss": _batched_loss(model, valid, valid.mask, device, BATCH_SIZE),
}
history.append(row)
print(
f"[{method}] epoch={epoch:02d} train={row['train_loss']:.4f} "
f"valid_selection={valid_selection_loss:.4f} clean={row['valid_clean_loss']:.4f}",
flush=True,
)
if valid_selection_loss < best_loss - 1e-4:
best_loss = valid_selection_loss
best_epoch = epoch
stale = 0
torch.save(model_state(model, method) | {"best_epoch": best_epoch}, checkpoint_path)
else:
stale += 1
if stale >= PATIENCE:
break
saved = torch.load(checkpoint_path, map_location=device, weights_only=False)
model.load_state_dict(saved["state_dict"])
model.eval()
write_csv(output_dir / "training_history.csv", history)
return model, best_epoch, history, train_mask_counts
def _group_map(ids: list[str]) -> tuple[list[str], dict[str, np.ndarray]]:
source_ids = [sample_id.split("$_$", 1)[0] for sample_id in ids]
groups = sorted(set(source_ids))
mapping = {
group: np.flatnonzero(np.asarray([source == group for source in source_ids]))
for group in groups
}
return groups, mapping
def test_group_bootstrap(test: Split, predictions: dict[str, dict[str, np.ndarray]]) -> list[dict[str, Any]]:
groups, mapping = _group_map(test.ids)
rng = np.random.default_rng(TEST_BOOTSTRAP_SEED)
draws: dict[str, list[float]] = defaultdict(list)
for _ in range(BOOTSTRAP_REPS):
selected = rng.choice(groups, size=len(groups), replace=True)
indices = np.concatenate([mapping[group] for group in selected])
values = {
method: metrics(test, predictions[method]["logits"], predictions[method]["intensity"], indices)
for method in METHODS
}
for name in values[EARLYCONCAT]:
draws[name].append(values[MOFE7_MLP][name] - values[EARLYCONCAT][name])
point = {
name: metrics(test, predictions[MOFE7_MLP]["logits"], predictions[MOFE7_MLP]["intensity"])[name]
- metrics(test, predictions[EARLYCONCAT]["logits"], predictions[EARLYCONCAT]["intensity"])[name]
for name in draws
}
return [{
"comparison": f"{MOFE7_MLP} minus {EARLYCONCAT}",
"metric": name,
"delta": point[name],
"bootstrap_ci_2p5": float(np.quantile(values, 0.025)),
"bootstrap_ci_97p5": float(np.quantile(values, 0.975)),
"bootstrap_probability_delta_gt_0": float(np.mean(np.asarray(values) > 0)),
"replicates": BOOTSTRAP_REPS,
"resampling_unit": "source video id",
"paired": True,
"seed": TEST_BOOTSTRAP_SEED,
} for name, values in draws.items()]
def validation_aurc_bootstrap(
valid: Split,
predictions: dict[tuple[str, str], dict[str, np.ndarray]],
rates_by_sample: dict[str, np.ndarray],
) -> list[dict[str, Any]]:
groups, mapping = _group_map(valid.ids)
rng = np.random.default_rng(AURC_BOOTSTRAP_SEED)
deltas: dict[str, list[float]] = {mode: [] for mode in CURVE_MODES}
def score(method: str, mode: str, indices: np.ndarray) -> float:
keys = curve_scenarios(mode)
xs = [float(np.nanmean(rates_by_sample[key][indices])) for key in keys]
ys = [
float(np.abs(valid.y_reg[indices] - predictions[(method, key)]["intensity"][indices]).mean())
for key in keys
]
return aurc_from_curve(xs, ys)
for _ in range(BOOTSTRAP_REPS):
selected = rng.choice(groups, size=len(groups), replace=True)
indices = np.concatenate([mapping[group] for group in selected])
for mode in CURVE_MODES:
deltas[mode].append(score(MOFE7_MLP, mode, indices) - score(EARLYCONCAT, mode, indices))
rows = []
for mode in CURVE_MODES:
all_indices = np.arange(valid.n)
values = deltas[mode]
rows.append({
"mask_mode": mode,
"delta_aurc_mae_mofe_minus_earlyconcat": score(MOFE7_MLP, mode, all_indices) - score(EARLYCONCAT, mode, all_indices),
"bootstrap_ci_2p5": float(np.quantile(values, 0.025)),
"bootstrap_ci_97p5": float(np.quantile(values, 0.975)),
"bootstrap_probability_delta_lt_0": float(np.mean(np.asarray(values) < 0)),
"replicates": BOOTSTRAP_REPS,
"resampling_unit": "source video id",
"paired": True,
"seed": AURC_BOOTSTRAP_SEED,
})
return rows
def run(device_name: str = "auto", output_dir: Path = OUTPUT_DIR) -> None:
if output_dir.exists() and any(output_dir.iterdir()):
raise FileExistsError(f"refusing to overwrite non-empty result directory: {output_dir}")
output_dir.mkdir(parents=True, exist_ok=True)
device = device_for(device_name)
if device.type == "cuda" and not torch.cuda.is_available():
raise RuntimeError("CUDA was requested but is unavailable")
feature_path = ATTACHMENT2 / "aligned_50.pkl"
raw_splits = load_splits(feature_path)
train_raw, valid_raw, test_raw = raw_splits["train"], raw_splits["valid"], raw_splits["test"]
stats = fit_robust_stats(train_raw)
train, valid, test = (apply_robust_stats(s, stats) for s in (train_raw, valid_raw, test_raw))
stats_path = output_dir / "aligned_robust_stats.npz"
stats.save(stats_path)
dims = tuple(int(x.shape[-1]) for x in train.x)
valid_scenarios = make_scenarios(valid, SCENARIO_SEED)
if len(valid_scenarios) != 42:
raise ValueError(f"expected 42 controlled scenarios, got {len(valid_scenarios)}")
rates_by_sample = actual_additional_rates(valid.mask, valid_scenarios)
set_deterministic(SEED)
order_rng = np.random.default_rng(SEED + 809)
orders = [order_rng.permutation(train.n) for _ in range(EPOCH_LIMIT)]
best_epochs: dict[str, int] = {}
training_rows: list[dict[str, Any]] = []
mask_count_rows: list[dict[str, Any]] = []
parameter_rows: list[dict[str, Any]] = []
for method in METHODS:
model_dir = output_dir / "models" / method / f"seed_{SEED}"
model_dir.mkdir(parents=True, exist_ok=True)
model, best_epoch, history, mask_counts = train_one(
method, train, valid, valid_scenarios, orders, model_dir, device
)
best_epochs[method] = best_epoch
training_rows.extend(history)
parameter_rows.append({
"method": method,
"parameters_total": sum(p.numel() for p in model.parameters()),
"parameters_trainable": sum(p.numel() for p in model.parameters() if p.requires_grad),
"best_epoch": best_epoch,
})
for key, count in sorted(mask_counts.items()):
mask_count_rows.append({"method": method, "seed": SEED, "rate_mode": key, "sample_epoch_assignments": count})
del model
if torch.cuda.is_available():
torch.cuda.empty_cache()
write_csv(output_dir / "training_history.csv", training_rows)
write_csv(output_dir / "training_mask_distribution.csv", mask_count_rows)
write_csv(output_dir / "parameter_count.csv", parameter_rows)
# Reload the selected checkpoints, then conduct one final official-test pass.
test_predictions: dict[str, dict[str, np.ndarray]] = {}
test_rows: list[dict[str, Any]] = []
condition_predictions: dict[tuple[str, str], dict[str, np.ndarray]] = {}
condition_rows: list[dict[str, Any]] = []
for method in METHODS:
checkpoint_path = output_dir / "models" / method / f"seed_{SEED}" / "model_best.pt"
saved = torch.load(checkpoint_path, map_location=device, weights_only=False)
model = build_model(method, dims, device)
model.load_state_dict(saved["state_dict"])
model.eval()
test_prediction = _predict(model, test, test.mask, device, BATCH_SIZE)
test_predictions[method] = test_prediction
test_rows.append({
"method": method,
"seed": SEED,
"best_epoch": best_epochs[method],
"n_test": test.n,
**metrics(test, test_prediction["logits"], test_prediction["intensity"]),
})
for scenario, masks in valid_scenarios.items():
prediction = _predict(model, valid, masks, device, BATCH_SIZE)
condition_predictions[(method, scenario)] = prediction
condition_rows.append({
"method": method,
"seed": SEED,
"scenario": scenario,
"realized_additional_global_rate": float(np.nanmean(rates_by_sample[scenario])),
"n_valid": valid.n,
**metrics(valid, prediction["logits"], prediction["intensity"]),
})
print(f"[valid/{method}] {scenario} done", flush=True)
del model
if torch.cuda.is_available():
torch.cuda.empty_cache()
write_csv(output_dir / "official_test_metrics_by_seed.csv", test_rows)
write_csv(output_dir / "official_test_paired_bootstrap.csv", test_group_bootstrap(test, test_predictions))
write_csv(output_dir / "controlled_metrics_by_scenario.csv", condition_rows)
test_summary = []
for method in METHODS:
row = next(r for r in test_rows if r["method"] == method)
for metric in ("accuracy", "macro_f1", "mae", "rmse", "pearson"):
test_summary.append({"method": method, "metric": metric, "mean": row[metric], "sd_across_seeds": 0.0, "n_seeds": 1})
write_csv(output_dir / "official_test_summary.csv", test_summary)
aurc_rows: list[dict[str, Any]] = []
for method in METHODS:
for mode in CURVE_MODES:
keys = curve_scenarios(mode)
xs = [float(np.nanmean(rates_by_sample[key])) for key in keys]
ys = [
float(np.abs(valid.y_reg - condition_predictions[(method, key)]["intensity"]).mean())
for key in keys
]
aurc_rows.append({
"method": method,
"seed": SEED,
"mask_mode": mode,
"aurc_mae": aurc_from_curve(xs, ys),
"rates_realized": json.dumps(xs),
})
write_csv(output_dir / "aurc_mae_by_mode_seed.csv", aurc_rows)
write_csv(output_dir / "aurc_mae_paired_bootstrap.csv", validation_aurc_bootstrap(valid, condition_predictions, rates_by_sample))
manifest = {
"experiment": "Retrained EarlyConcat and MoFE-7 + MLP Router using math/Q2 V2-compatible protocol",
"created_unix": time.time(),
"device": str(device),
"cuda_device": torch.cuda.get_device_name(0) if device.type == "cuda" else None,
"feature_file": str(feature_path),
"feature_sha256": sha256(feature_path),
"representation": "official aligned_50 ordered positions; not Q1 physical-time bins",
"train_valid_test_counts": {name: split.n for name, split in raw_splits.items()},
"source_video_groups": {name: len({sid.split("$_$", 1)[0] for sid in split.ids}) for name, split in raw_splits.items()},
"official_group_splits_disjoint": True,
"train_only_scaler": str(stats_path),
"scaler_fit": "median and 1.4826*MAD on observed training rows only; zero-MAD fallback to std then 1",
"seed": SEED,
"model_seeds": [SEED],
"training_configuration": {
"epoch_limit": EPOCH_LIMIT,
"early_stopping_patience": PATIENCE,
"batch_size": BATCH_SIZE,
"optimizer": "AdamW",
"learning_rate": LEARNING_RATE,
"weight_decay": WEIGHT_DECAY,
"gradient_clip_norm": 1.0,
"early_stopping_metric": "mean validation joint CE + 0.5*SmoothL1 over 0.0/none, 0.3/single, 0.3/sync, 0.5/async",
"architecture_preserved": {
EARLYCONCAT: "EarlyConcat + BiGRU",
MOFE7_MLP: "MoFE-7 + MLP Router",
},
"objective": "cross entropy + 0.5 * SmoothL1(intensity/3, label/3); same objective for both methods",
"training_corruption": {
"rates": list(TRAIN_RATES),
"patterns": list(TRAIN_MODES),
"preserve_at_least_fraction_per_selected_modality": 0.2,
"generator_seed": TRAIN_MASK_SEED,
"same_sample_masks_and_batch_orders_across_models": True,
},
},
"validation_protocol": {
"scenario_seed": SCENARIO_SEED,
"scenario_count": len(valid_scenarios),
"same_fixed_masks_for_both_models": True,
"scenario_design": "math/Q2 42 controlled continuous-mask scenarios regenerated on each sample's original observation mask",
"selection_scenarios": list(SELECTION_SCENARIOS),
"selection_note": "Deterministic-model adaptation; uses joint supervised loss instead of C5's probabilistic selection NLL.",
"aurc": "normalized trapezoidal MAE area over realized equal-modality-weighted additional missing rate for single/sync/partial/async at 0/.1/.3/.5/.7",
},
"test_protocol": {
"official_test_final_clean_passes": 1,
"test_used_for_training_or_checkpoint_selection": False,
"metrics": ["accuracy", "macro_f1", "mae", "rmse", "pearson"],
"paired_group_bootstrap_replicates": BOOTSTRAP_REPS,
"bootstrap_unit": "source video id",
"bootstrap_seed": TEST_BOOTSTRAP_SEED,
},
}
(output_dir / "run_manifest.json").write_text(json.dumps(manifest, indent=2), encoding="utf-8")
(output_dir / "hypothesis.md").write_text(
"# R03: 按 math/Q2 V2 口径重训两种保留模型\n\n"
"## 假设\n\n"
"在保持 EarlyConcat + BiGRU 与 MoFE-7 + MLP Router 结构及共同监督目标不变的情况下,"
"使用数学方案中的官方划分、连续块缺失训练和 42 个固定验证情景,可以公平比较两种模型的干净测试表现与缺失鲁棒性。\n\n"
"## 唯一实验改动\n\n"
"相对现有检查点,本轮重新训练时将缺失训练改为 0/10/30/50/70% 与 single/sync/partial/async,"
"每个被选模态至少保留 20% 观测;训练和批次顺序在两个模型间配对。数学方案中的 C5 概率损失不适用于现有确定性分类/回归头,"
"因此保留项目既有的 CE + 0.5 SmoothL1 联合目标。\n\n"
"## 数据使用\n\n"
"标准化器只在官方训练集观测行上拟合;官方验证集只用于早停与缺失评估;官方测试集在全部检查点确定后做一次干净评估。\n",
encoding="utf-8",
)
print(f"wrote retraining results to {output_dir}", flush=True)
print(f"train/valid/test={train.n}/{valid.n}/{test.n}; device={device}; best_epochs={best_epochs}", flush=True)
for row in test_rows:
print(
f"{row['method']}: Acc={row['accuracy']:.4f} Macro-F1={row['macro_f1']:.4f} "
f"MAE={row['mae']:.4f} RMSE={row['rmse']:.4f} Pearson={row['pearson']:.4f}",
flush=True,
)
if __name__ == "__main__":
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--device", default="auto", choices=("auto", "cuda", "cpu"))
parser.add_argument("--output-dir", type=Path, default=OUTPUT_DIR)
arguments = parser.parse_args()
run(device_name=arguments.device, output_dir=arguments.output_dir)