Time-Shell Long-Horizon Decoder / bench_run.py
Beats tuned baseline
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()