Wasserstein-Controlled Gaussian-Mixture Rollouts / stage2_bench.py

Failed on benchmark

Raw ⬇ ZIP
  1import sys, json, math, random
  2import numpy as np
  3import torch
  4import torch.nn as nn
  5
  6sys.path.insert(0, "/home/maxwelhelp/all/math2nn")
  7from bench import get_dataset, make_model, train_model, evaluate, sweep_baseline, make_report
  8
  9SEEDS = tuple(range(8))
 10SWEEP_SEEDS = tuple(range(4))
 11EPOCHS = 15
 12NTRAIN, NTEST = 800, 300
 13BATCH = 128
 14
 15# Same GRU backbone and hidden size as bench.models.rnn_small.
 16class MixtureRNN(nn.Module):
 17    def __init__(self, hidden=64, modes=2):
 18        super().__init__()
 19        self.rnn = nn.GRU(3, hidden, batch_first=True)
 20        self.head = nn.Linear(hidden, 3*modes)  # logits, means, log standard deviations
 21        self.modes = modes
 22    def forward(self, x):
 23        seq = x.view(x.shape[0], -1, 3)
 24        _, h = self.rnn(seq)
 25        return self.head(h[-1])
 26
 27def seed_all(seed):
 28    random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
 29    if torch.cuda.is_available():
 30        try: torch.cuda.manual_seed_all(seed)
 31        except Exception: pass
 32
 33def device():
 34    return "cuda" if torch.cuda.is_available() else "cpu"
 35
 36def mixture_loss(raw, y):
 37    k = raw.shape[1] // 3
 38    logits, means, logstd = raw[:, :k], raw[:, k:2*k], raw[:, 2*k:]
 39    logstd = logstd.clamp(-5.0, 2.0)
 40    z = (y - means) / logstd.exp()
 41    lp = -0.5*z*z - logstd - 0.5*math.log(2*math.pi)
 42    return -(torch.log_softmax(logits, 1) + lp).logsumexp(1).mean()
 43
 44def train_mixture(seed, lr, epochs=EPOCHS, return_model=False):
 45    seed_all(seed); ds = get_dataset("dynamics", seed, NTRAIN, NTEST)
 46    net = MixtureRNN(); dev = device()
 47    try:
 48        net.to(dev); x, y = ds["xtr"].to(dev), ds["ytr"].to(dev)
 49        opt = torch.optim.Adam(net.parameters(), lr=lr)
 50        for _ in range(epochs):
 51            net.train(); perm = torch.randperm(len(x), device=dev)
 52            for i in range(0, len(x), BATCH):
 53                ix = perm[i:i+BATCH]; loss = mixture_loss(net(x[ix]), y[ix])
 54                opt.zero_grad(); loss.backward(); opt.step()
 55        net.eval()
 56        with torch.no_grad():
 57            raw = net(ds["xte"].to(dev)); k=2
 58            w = torch.softmax(raw[:,:k],1); mu=raw[:,k:2*k]
 59            pred=(w*mu).sum(1,keepdim=True)
 60            mse=((pred-ds["yte"].to(dev))**2).mean().item()
 61        return (mse, net, ds, raw.detach().cpu()) if return_model else mse
 62    except RuntimeError:
 63        # Small CPU retry mirrors the harness's GPU-fallback intent.
 64        seed_all(seed); net = MixtureRNN(); net.to("cpu")
 65        x,y=ds["xtr"],ds["ytr"]; opt=torch.optim.Adam(net.parameters(),lr=lr)
 66        for _ in range(epochs):
 67            perm=torch.randperm(len(x))
 68            for i in range(0,len(x),BATCH):
 69                ix=perm[i:i+BATCH]; loss=mixture_loss(net(x[ix]),y[ix])
 70                opt.zero_grad(); loss.backward(); opt.step()
 71        with torch.no_grad():
 72            raw=net(ds["xte"]); w=torch.softmax(raw[:,:2],1); mu=raw[:,2:4]
 73            mse=((w*mu).sum(1,keepdim=True)-ds["yte"]).pow(2).mean().item()
 74        return (mse,net,ds,raw) if return_model else mse
 75
 76def baseline_fn(cfg):
 77    def run(seed):
 78        seed_all(seed); ds=get_dataset("dynamics",seed,NTRAIN,NTEST)
 79        net=make_model("rnn_small",ds["input_shape"],ds["out_dim"])
 80        _, metric, _=train_model(net,ds,epochs=EPOCHS,lr=cfg["lr"],batch=BATCH,log=lambda *_:None)
 81        return metric
 82    return run
 83
 84def idea_fn(cfg):
 85    return lambda seed: train_mixture(seed,cfg["lr"])
 86
 87def signature(seed, lr):
 88    mse, net, ds, raw = train_mixture(seed,lr,return_model=True)
 89    w=torch.softmax(raw[:,:2],1); mu=raw[:,2:4]; sd=raw[:,4:6].clamp(-5,2).exp()
 90    sep=(mu[:,0]-mu[:,1]).abs(); avg_sd=(w*sd).sum(1)
 91    # Model-derived chance estimate at theta<=0 versus observed test frequency.
 92    cdf=0.5*(1+torch.erf((-mu)/(sd*math.sqrt(2))))
 93    p_mix=(w*cdf).sum(1).numpy(); p_gauss=(0.5*(1+torch.erf((-(w*mu).sum(1))/(torch.sqrt((w*(sd**2+(mu-(w*mu).sum(1,keepdim=True))**2)).sum(1))*math.sqrt(2))))).numpy()
 94    obs=(ds["yte"].numpy().reshape(-1)<=0).astype(float)
 95    return {"n":len(obs),"test_mse":mse,"predicted_mean_mode_separation":float(sep.mean()),"predicted_mean_component_sd":float(avg_sd.mean()),"predicted_bimodal_fraction":float((sep>2*avg_sd).float().mean()),"mixture_chance_abs_error":float(abs(p_mix.mean()-obs.mean())),"moment_gaussian_chance_abs_error":float(abs(p_gauss.mean()-obs.mean())),"observed_event_rate":float(obs.mean()),"confirmed":bool(sep.mean().item()>0.05 and abs(p_mix.mean()-obs.mean()) < abs(p_gauss.mean()-obs.mean()))}
 96
 97def main():
 98    grid=[{"lr":1e-3},{"lr":3e-3},{"lr":1e-2}]
 99    base=sweep_baseline(baseline_fn,grid,seeds=SWEEP_SEEDS)
100    # Equal-size intervention sweep over exactly the baseline union of learning rates.
101    idea_trials=[]
102    for cfg in grid:
103        r=evaluate(idea_fn(cfg),seeds=SWEEP_SEEDS)
104        idea_trials.append({"cfg":cfg,"mean":r["mean"]})
105    best_cfg=min(idea_trials,key=lambda z:z["mean"])["cfg"]
106    idea=evaluate(idea_fn(best_cfg),seeds=SEEDS)
107    # Nearby settings are explicitly run on all paired seeds for transparent reporting.
108    nearby={str(c["lr"]):evaluate(idea_fn(c),seeds=SEEDS) for c in grid}
109    base["idea_union_sweep_note"]="baseline evaluated at every intervention learning rate; final baseline is its tuned best config"
110    rep=make_report("dynamics","rnn_small",base,idea,extra={"prediction":"A learned mixture should retain separated predictive modes and improve threshold-probability calibration when the trained task is multimodal.","idea_sweep":idea_trials,"idea_nearby_full":nearby,"trained_model_signature":signature(SEEDS[0],best_cfg["lr"])})
111    rep["selection"]={"idea_best_cfg":best_cfg,"epochs":EPOCHS,"n_train":NTRAIN,"n_test":NTEST,"structural_match":"dynamics: controlled pendulum multi-step target"}
112    with open("bench_report.json","w") as f: json.dump(rep,f,indent=2)
113    print(json.dumps(rep,indent=2))
114
115if __name__=="__main__": main()