Time-Shell Long-Horizon Decoder / bench_run.py

✓✓ Beats tuned baseline

Raw ⬇ ZIP
  1import sys, json, time
  2from pathlib import Path
  3import numpy as np
  4import torch
  5import torch.nn as nn
  6
  7sys.path.insert(0, "/home/maxwelhelp/all/math2nn")
  8from bench import train_model, sweep_baseline, make_report
  9
 10from custom_track import get_dataset
 11
 12SEEDS = tuple(range(8))
 13LR_GRID = [1e-3, 3e-3, 6e-3]
 14SCALE_GRID = [0.5, 1.0]
 15EPOCHS = 18
 16BATCH = 128
 17
 18class HorizonDecoder(nn.Module):
 19    def __init__(self, horizon, mode="shell", scale=1.0):
 20        super().__init__()
 21        self.horizon, self.mode, self.scale = horizon, mode, scale
 22        self.gru = nn.GRU(1, 32, batch_first=True)
 23        self.h = nn.Linear(32, 32)
 24        self.e = nn.Parameter(torch.randn(horizon, 8) * .04)
 25        self.beta = nn.Parameter(torch.linspace(1., .1, horizon).unsqueeze(1), requires_grad=False)
 26        in_dim = 32 + 8 + (1 if mode == "shell" else horizon)
 27        self.readout = nn.Sequential(nn.Linear(in_dim, 48), nn.ReLU(), nn.Linear(48, 1))
 28
 29    def forward(self, x):
 30        _, hn = self.gru(x.unsqueeze(-1))
 31        h = torch.tanh(self.h(hn[-1]))
 32        k = self.horizon
 33        emb = self.e.unsqueeze(0).expand(x.shape[0], -1, -1)
 34        if self.mode == "shell":
 35            b = self.beta[:, 0]
 36            B = torch.flip(torch.cumsum(torch.flip(b, (0,)), 0), (0,))
 37            coeff = (B / k * self.scale).view(1, k, 1).expand(x.shape[0], -1, -1)
 38        else:
 39            coeff = (self.beta[:, 0] / k * self.scale).view(1, 1, k).expand(x.shape[0], k, -1)
 40        hh = h.unsqueeze(1).expand(-1, k, -1)
 41        z = torch.cat([hh, emb, coeff], dim=-1)
 42        return self.readout(z).squeeze(-1)
 43
 44def make_system(mode, scale):
 45    return HorizonDecoder(16, mode=mode, scale=scale)
 46
 47def train_one(mode, cfg, seed, capture=False):
 48    torch.manual_seed(seed); np.random.seed(seed)
 49    raw = get_dataset(seed, 400, 200)
 50    ds = {k: (torch.from_numpy(v) if isinstance(v, np.ndarray) else v) for k,v in raw.items()}
 51    net = make_system(mode, cfg["scale"])
 52    t0 = time.perf_counter()
 53    trained, metric, hist = train_model(net, ds, epochs=EPOCHS, lr=cfg["lr"], batch=BATCH, log=lambda *_: None)
 54    elapsed = time.perf_counter() - t0
 55    if trained is None: return float("nan")
 56    if capture: return trained, ds, float(metric), elapsed
 57    return float(metric)
 58
 59def baseline_factory(cfg):
 60    return lambda seed: train_one("dense", cfg, seed)
 61
 62def idea_factory(cfg):
 63    return lambda seed: train_one("shell", cfg, seed)
 64
 65def signature(cfg):
 66    net, ds, metric, elapsed = train_one("shell", cfg, 0, capture=True)
 67    net = net.cpu()
 68    net.eval()
 69    dev = torch.device("cpu")
 70    x = ds["xte"][:64].to(dev)
 71    with torch.no_grad():
 72        base = net(x).detach().cpu().numpy()
 73        # Re-test the trained mechanism: changing a late coefficient should
 74        # affect early shells, while equal total coefficient preserves B_0.
 75        old = net.beta.detach().clone()
 76        b = old[:,0].clone(); perm = b.flip(0)
 77        net.beta[:,0].copy_(perm)
 78        changed = net(x).detach().cpu().numpy()
 79        net.beta[:,0].copy_(b)
 80    early = float(np.mean(np.abs(changed[:,0] - base[:,0])))
 81    late = float(np.mean(np.abs(changed[:,-1] - base[:,-1])))
 82    return {"prediction": "reverse cumulative shell makes early outputs more sensitive to later coefficients",
 83            "predicted": {"early_over_late_sensitivity": ">1"},
 84            "observed": {"early_abs_change": early, "late_abs_change": late,
 85                         "early_over_late_sensitivity": early / (late + 1e-12), "trained_test_mse": metric},
 86            "confirmed": bool(early > late), "forward_seconds": elapsed}
 87
 88def main():
 89    # Equal union: every idea lr/scale is also evaluated for the dense baseline.
 90    grid = [{"lr": lr, "scale": scale} for lr in LR_GRID for scale in SCALE_GRID]
 91    base = sweep_baseline(baseline_factory, grid, seeds=(0,1,2,3))
 92    idea_cfgs = [base["best_cfg"], {"lr": 1e-3, "scale": 1.0}, {"lr": 6e-3, "scale": 0.5}]
 93    idea_runs = []
 94    for cfg in idea_cfgs:
 95        r = {"cfg": cfg, "result": __import__('bench').evaluate(idea_factory(cfg), seeds=SEEDS)}
 96        idea_runs.append(r)
 97    best = min(idea_runs, key=lambda r: r["result"]["mean"])
 98    rep = make_report("multihorizon_sequence", "custom_gru_decoder", base, best["result"],
 99                      {"signature": signature(best["cfg"]), "idea_sweep": idea_runs,
100                       "complexity_prediction": {"shell": "O(Kd)", "dense": "O(K^2d)"},
101                       "custom_track": {"name": "multihorizon_sequence", "file": "custom_track.py", "domain": "sequence"}})
102    Path("bench_report.json").write_text(json.dumps(rep, indent=2))
103    print(json.dumps(rep, indent=2))
104
105if __name__ == "__main__": main()