CoolFace
Modelpublic

Cccccz/HY

sourceHugging Faceupdated 11d agoView on Hugging Face
0likes
train_evaluate_predictor_confidence.py282 linesDownload Raw Back to tools
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 
Cccccz/HY · CoolFace