import sys, json, time from pathlib import Path import numpy as np import torch import torch.nn as nn sys.path.insert(0, "/home/maxwelhelp/all/math2nn") from bench import train_model, sweep_baseline, make_report from custom_track import get_dataset SEEDS = tuple(range(8)) LR_GRID = [1e-3, 3e-3, 6e-3] SCALE_GRID = [0.5, 1.0] EPOCHS = 18 BATCH = 128 class HorizonDecoder(nn.Module): def __init__(self, horizon, mode="shell", scale=1.0): super().__init__() self.horizon, self.mode, self.scale = horizon, mode, scale self.gru = nn.GRU(1, 32, batch_first=True) self.h = nn.Linear(32, 32) self.e = nn.Parameter(torch.randn(horizon, 8) * .04) self.beta = nn.Parameter(torch.linspace(1., .1, horizon).unsqueeze(1), requires_grad=False) in_dim = 32 + 8 + (1 if mode == "shell" else horizon) self.readout = nn.Sequential(nn.Linear(in_dim, 48), nn.ReLU(), nn.Linear(48, 1)) def forward(self, x): _, hn = self.gru(x.unsqueeze(-1)) h = torch.tanh(self.h(hn[-1])) k = self.horizon emb = self.e.unsqueeze(0).expand(x.shape[0], -1, -1) if self.mode == "shell": b = self.beta[:, 0] B = torch.flip(torch.cumsum(torch.flip(b, (0,)), 0), (0,)) coeff = (B / k * self.scale).view(1, k, 1).expand(x.shape[0], -1, -1) else: coeff = (self.beta[:, 0] / k * self.scale).view(1, 1, k).expand(x.shape[0], k, -1) hh = h.unsqueeze(1).expand(-1, k, -1) z = torch.cat([hh, emb, coeff], dim=-1) return self.readout(z).squeeze(-1) def make_system(mode, scale): return HorizonDecoder(16, mode=mode, scale=scale) def train_one(mode, cfg, seed, capture=False): torch.manual_seed(seed); np.random.seed(seed) raw = get_dataset(seed, 400, 200) ds = {k: (torch.from_numpy(v) if isinstance(v, np.ndarray) else v) for k,v in raw.items()} net = make_system(mode, cfg["scale"]) t0 = time.perf_counter() trained, metric, hist = train_model(net, ds, epochs=EPOCHS, lr=cfg["lr"], batch=BATCH, log=lambda *_: None) elapsed = time.perf_counter() - t0 if trained is None: return float("nan") if capture: return trained, ds, float(metric), elapsed return float(metric) def baseline_factory(cfg): return lambda seed: train_one("dense", cfg, seed) def idea_factory(cfg): return lambda seed: train_one("shell", cfg, seed) def signature(cfg): net, ds, metric, elapsed = train_one("shell", cfg, 0, capture=True) net = net.cpu() net.eval() dev = torch.device("cpu") x = ds["xte"][:64].to(dev) with torch.no_grad(): base = net(x).detach().cpu().numpy() # Re-test the trained mechanism: changing a late coefficient should # affect early shells, while equal total coefficient preserves B_0. old = net.beta.detach().clone() b = old[:,0].clone(); perm = b.flip(0) net.beta[:,0].copy_(perm) changed = net(x).detach().cpu().numpy() net.beta[:,0].copy_(b) early = float(np.mean(np.abs(changed[:,0] - base[:,0]))) late = float(np.mean(np.abs(changed[:,-1] - base[:,-1]))) return {"prediction": "reverse cumulative shell makes early outputs more sensitive to later coefficients", "predicted": {"early_over_late_sensitivity": ">1"}, "observed": {"early_abs_change": early, "late_abs_change": late, "early_over_late_sensitivity": early / (late + 1e-12), "trained_test_mse": metric}, "confirmed": bool(early > late), "forward_seconds": elapsed} def main(): # Equal union: every idea lr/scale is also evaluated for the dense baseline. grid = [{"lr": lr, "scale": scale} for lr in LR_GRID for scale in SCALE_GRID] base = sweep_baseline(baseline_factory, grid, seeds=(0,1,2,3)) idea_cfgs = [base["best_cfg"], {"lr": 1e-3, "scale": 1.0}, {"lr": 6e-3, "scale": 0.5}] idea_runs = [] for cfg in idea_cfgs: r = {"cfg": cfg, "result": __import__('bench').evaluate(idea_factory(cfg), seeds=SEEDS)} idea_runs.append(r) best = min(idea_runs, key=lambda r: r["result"]["mean"]) rep = make_report("multihorizon_sequence", "custom_gru_decoder", base, best["result"], {"signature": signature(best["cfg"]), "idea_sweep": idea_runs, "complexity_prediction": {"shell": "O(Kd)", "dense": "O(K^2d)"}, "custom_track": {"name": "multihorizon_sequence", "file": "custom_track.py", "domain": "sequence"}}) Path("bench_report.json").write_text(json.dumps(rep, indent=2)) print(json.dumps(rep, indent=2)) if __name__ == "__main__": main()