Cccccz/HY
0
1#!/usr/bin/env python32"""Train C0/C1/C2 confidence heads and report held-out risk coverage."""3 4from __future__ import annotations5 6import argparse7import csv8import json9import math10import os11import random12from pathlib import Path13 14import matplotlib.pyplot as plt15import numpy as np16import torch17import torch.distributed as dist18import torch.nn.functional as F19from scipy.stats import pearsonr, spearmanr20from torch.nn.parallel import DistributedDataParallel as DDP21from torch.utils.data import DataLoader, TensorDataset22from torch.utils.data.distributed import DistributedSampler23 24from models.confidence_head import PredictorConfidenceHead25 26 27TARGETS = ("c0", "c1", "c2")28RATIOS = (0.0, 0.05, 0.1, 0.2, 0.25, 1 / 3, 0.4, 0.5, 0.6, 0.75, 1.0)29 30 31def load_features(directory):32 shards = [torch.load(path, map_location="cpu", weights_only=False) for path in sorted(directory.glob("features_rank_*.pt"))]33 if not shards:34 raise FileNotFoundError(f"No feature shards in {directory}")35 names = ("pooled_hidden", "log_residual", "true_c0", "true_c1", "true_c2", "scheduler_coeff", "step_id")36 result = {name: torch.cat([shard[name] for shard in shards]) for name in names}37 result["metadata"] = sum((shard["metadata"] for shard in shards), [])38 order = sorted(range(len(result["metadata"])), key=lambda i: tuple(result["metadata"][i][k] for k in ("case_id", "action_id", "chunk_id", "step_id")))39 index = torch.tensor(order)40 for name in names:41 result[name] = result[name][index]42 result["metadata"] = [result["metadata"][i] for i in order]43 return result44 45 46def train_head(47 features, target_name, device, epochs, batch_size, seed, rank, world_size,48 mlp_dims,49):50 torch.manual_seed(seed)51 model = PredictorConfidenceHead(52 features["pooled_hidden"].shape[1], mlp_dims=mlp_dims53 ).to(device)54 optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=1e-4)55 dataset = TensorDataset(56 features["pooled_hidden"].float(), features["step_id"].long(),57 features["log_residual"].float(),58 torch.log(features[f"true_{target_name}"].float() + 1e-8),59 )60 sampler = DistributedSampler(dataset, num_replicas=world_size, rank=rank, shuffle=True, seed=seed) if world_size > 1 else None61 loader = DataLoader(dataset, batch_size=batch_size, shuffle=sampler is None, sampler=sampler)62 wrapped = DDP(model, device_ids=[device.index], broadcast_buffers=False) if world_size > 1 else model63 wrapped.train()64 for epoch in range(epochs):65 if sampler is not None:66 sampler.set_epoch(epoch)67 total = count = 068 for pooled, step, residual, target in loader:69 pooled, step, residual, target = (x.to(device) for x in (pooled, step, residual, target))70 optimizer.zero_grad(set_to_none=True)71 loss = F.smooth_l1_loss(wrapped(pooled, step, residual), target)72 loss.backward()73 optimizer.step()74 total += float(loss.detach()) * len(pooled)75 count += len(pooled)76 if rank == 0 and (epoch == 0 or (epoch + 1) % 10 == 0):77 print(f"{target_name} epoch {epoch + 1}/{epochs} loss={total / count:.6f}", flush=True)78 if world_size > 1:79 dist.barrier()80 return model.eval()81 82 83@torch.inference_mode()84def predict(model, features, device):85 result = []86 for start in range(0, len(features["step_id"]), 512):87 stop = start + 51288 result.append(model(89 features["pooled_hidden"][start:stop].float().to(device),90 features["step_id"][start:stop].long().to(device),91 features["log_residual"][start:stop].float().to(device),92 ).cpu())93 return torch.cat(result).numpy()94 95 96def correlation(x, y, function):97 if len(x) < 2 or np.all(x == x[0]) or np.all(y == y[0]):98 return float("nan")99 return float(function(x, y).statistic)100 101 102def metric_rows(split, predictions, features):103 rows = []104 steps = features["step_id"].numpy()105 for target, score in predictions.items():106 truth = np.log(features[f"true_{target}"].numpy() + 1e-8)107 groups = [("overall", np.ones(len(steps), dtype=bool))] + [(f"P{s}", steps == s) for s in (1, 2, 3)]108 for label, mask in groups:109 rows.append({110 "model": f"Confidence {target.upper()}", "split": split, "step": label,111 "spearman": correlation(score[mask], truth[mask], spearmanr),112 "pearson": correlation(score[mask], truth[mask], pearsonr),113 "log_mae": float(np.abs(score[mask] - truth[mask]).mean()),114 })115 return rows116 117 118def remaining_risk(risk, score, ratio):119 count = min(len(risk), int(round(ratio * len(risk))))120 if count == 0:121 return float(risk.mean())122 selected = np.argpartition(score, len(score) - count)[-count:]123 keep = np.ones(len(risk), dtype=bool)124 keep[selected] = False125 return float(np.where(keep, risk, 0.0).mean())126 127 128def risk_rows(train, validation, predictions, seed, risk_target):129 risk = validation[f"true_{risk_target}"].numpy()130 train_steps = train["step_id"].numpy()131 train_risk = train[f"true_{risk_target}"].numpy()132 means = {step: train_risk[train_steps == step].mean() for step in (1, 2, 3)}133 methods = {134 "Timestep-only": np.array([means[int(s)] for s in validation["step_id"]]),135 **{136 f"Confidence {target.upper()}": score137 for target, score in predictions.items()138 },139 "Oracle": risk,140 }141 rows = []142 rng = np.random.default_rng(seed)143 for ratio in RATIOS:144 count = min(len(risk), int(round(ratio * len(risk))))145 random_risks = []146 for _ in range(100):147 chosen = rng.choice(len(risk), count, replace=False) if count else []148 keep = np.ones(len(risk), dtype=bool)149 keep[chosen] = False150 random_risks.append(float(np.where(keep, risk, 0.0).mean()))151 rows.append({"method": "Random", "fallback_ratio": ratio, "remaining_risk": np.mean(random_risks), "std_if_random": np.std(random_risks)})152 for method, score in methods.items():153 rows.append({"method": method, "fallback_ratio": ratio, "remaining_risk": remaining_risk(risk, score, ratio), "std_if_random": ""})154 return rows155 156 157def write_csv(path, rows):158 path.parent.mkdir(parents=True, exist_ok=True)159 with path.open("w", newline="", encoding="utf-8") as handle:160 writer = csv.DictWriter(handle, fieldnames=list(rows[0]))161 writer.writeheader()162 writer.writerows(rows)163 164 165def make_plots(validation, risks, predictions, output, risk_target):166 output.mkdir(parents=True, exist_ok=True)167 plt.figure(figsize=(7.2, 5.2))168 confidence_method = f"Confidence {risk_target.upper()}"169 for method in ("Random", "Timestep-only", confidence_method, "Oracle"):170 rows = [row for row in risks if row["method"] == method]171 plt.plot([100 * row["fallback_ratio"] for row in rows], [row["remaining_risk"] for row in rows], marker="o", label=method)172 plt.xlabel("Fallback ratio (%)")173 plt.ylabel(f"Remaining {risk_target.upper()} risk")174 plt.legend()175 plt.tight_layout()176 plt.savefig(output / "risk_coverage.png", dpi=180)177 plt.close()178 scores = predictions[risk_target]179 count = int(round(len(scores) / 3))180 chosen = np.argpartition(scores, len(scores) - count)[-count:]181 steps = validation["step_id"].numpy()[chosen]182 plt.figure(figsize=(5.5, 4.4))183 plt.bar(("P1", "P2", "P3"), [float((steps == s).mean()) for s in (1, 2, 3)])184 plt.ylabel("Fraction of selected fallbacks")185 plt.tight_layout()186 plt.savefig(output / "fallback_distribution_33pct.png", dpi=180)187 plt.close()188 189 190def main():191 parser = argparse.ArgumentParser(description=__doc__)192 parser.add_argument("--train_features", type=Path, required=True)193 parser.add_argument("--validation_features", type=Path, required=True)194 parser.add_argument("--output_dir", type=Path, required=True)195 parser.add_argument("--epochs", type=int, default=40)196 parser.add_argument("--batch_size", type=int, default=256)197 parser.add_argument("--seed", type=int, default=0)198 parser.add_argument("--targets", nargs="+", choices=TARGETS, default=list(TARGETS))199 parser.add_argument("--risk_target", choices=TARGETS, default="c1")200 parser.add_argument("--mlp_dims", nargs="+", type=int, default=[256])201 args = parser.parse_args()202 if args.risk_target not in args.targets:203 parser.error("--risk_target must be included in --targets")204 if any(dim <= 0 for dim in args.mlp_dims):205 parser.error("--mlp_dims values must be positive")206 targets = tuple(dict.fromkeys(args.targets))207 mlp_dims = tuple(args.mlp_dims)208 random.seed(args.seed); np.random.seed(args.seed); torch.manual_seed(args.seed)209 world_size = int(os.environ.get("WORLD_SIZE", 1)); rank = int(os.environ.get("RANK", 0)); local_rank = int(os.environ.get("LOCAL_RANK", 0))210 if world_size > 1:211 torch.cuda.set_device(local_rank); dist.init_process_group("nccl")212 device = torch.device("cuda", local_rank) if torch.cuda.is_available() else torch.device("cpu")213 train = load_features(args.train_features); validation = load_features(args.validation_features)214 if rank == 0:215 for child in ("checkpoints", "plots", "predictions", "tables"):216 (args.output_dir / child).mkdir(parents=True, exist_ok=True)217 if world_size > 1: dist.barrier()218 train_predictions = {}; validation_predictions = {}219 for target in targets:220 model = train_head(221 train, target, device, args.epochs, args.batch_size, args.seed,222 rank, world_size, mlp_dims,223 )224 if rank == 0:225 torch.save(226 {227 "model": model.state_dict(),228 "target": target,229 "hidden_dim": train["pooled_hidden"].shape[1],230 "timestep_dim": 32,231 "mlp_dims": list(mlp_dims),232 },233 args.output_dir / "checkpoints" / f"confidence_{target}.pt",234 )235 train_predictions[target] = predict(model, train, device)236 validation_predictions[target] = predict(model, validation, device)237 if world_size > 1: dist.barrier()238 if rank != 0:239 dist.destroy_process_group(); return240 metrics = metric_rows("train", train_predictions, train) + metric_rows("validation", validation_predictions, validation)241 risks = risk_rows(242 train, validation, validation_predictions, args.seed, args.risk_target243 )244 write_csv(args.output_dir / "tables" / "confidence_metrics.csv", metrics)245 write_csv(args.output_dir / "tables" / "risk_coverage.csv", risks)246 prediction_rows = [{**item, **{f"pred_{target}": float(validation_predictions[target][i]) for target in targets}, **{f"true_{target}": float(validation[f"true_{target}"][i]) for target in targets}} for i, item in enumerate(validation["metadata"])]247 write_csv(args.output_dir / "predictions" / "validation_predictions.csv", prediction_rows)248 make_plots(249 validation, risks, validation_predictions, args.output_dir / "plots",250 args.risk_target,251 )252 at_33 = {row["method"]: float(row["remaining_risk"]) for row in risks if math.isclose(row["fallback_ratio"], 1 / 3)}253 confidence_method = f"Confidence {args.risk_target.upper()}"254 target_metrics = {row["step"]: row for row in metrics if row["split"] == "validation" and row["model"] == confidence_method}255 scores = validation_predictions[args.risk_target]; count = int(round(len(scores) / 3)); chosen = np.argpartition(scores, len(scores) - count)[-count:]256 chosen_mask = np.zeros(len(scores), dtype=bool); chosen_mask[chosen] = True; steps = validation["step_id"].numpy()257 summary = {258 "protocol": "WorldPlay fulltrain100 (8400 transitions); independent validation25 (2100 transitions); no validation tuning",259 "targets": list(targets), "risk_target": args.risk_target,260 "mlp_dims": list(mlp_dims),261 "train_samples": len(train["step_id"]), "validation_samples": len(validation["step_id"]),262 "risk_target_validation_metrics": target_metrics,263 "remaining_risk_at_33pct": at_33,264 "gain_vs_random": (at_33["Random"] - at_33[confidence_method]) / at_33["Random"],265 "oracle_gap": (at_33[confidence_method] - at_33["Oracle"]) / max(at_33["Oracle"], 1e-12),266 "routing_at_33pct": {f"P{s}": {"selection_fraction": float((steps[chosen_mask] == s).mean()), "within_step_fallback_rate": float(chosen_mask[steps == s].mean())} for s in (1, 2, 3)},267 }268 if args.risk_target == "c1":269 summary["c1_validation_metrics"] = target_metrics270 (args.output_dir / "summary.json").write_text(json.dumps(summary, indent=2, sort_keys=True) + "\n")271 (args.output_dir / "summary.md").write_text(272 "# HY-WorldPlay offline confidence validation\n\n" + summary["protocol"] + "\n\n" +273 f"## {args.risk_target.upper()} validation Spearman\n\n" + "\n".join(f"- {key}: {value['spearman']:.4f}" for key, value in target_metrics.items()) +274 f"\n\n## 33.3% fallback remaining {args.risk_target.upper()} risk\n\n" + "\n".join(f"- {key}: {value:.8g}" for key, value in at_33.items()) + "\n"275 )276 print(json.dumps(summary, indent=2), flush=True)277 if world_size > 1: dist.destroy_process_group()278 279 280if __name__ == "__main__":281 main()282 