Positive-real rational resolvent mixer / bench_resolvent.py
Failed on benchmark
1import json
2import random
3import sys
4from pathlib import Path
5
6import numpy as np
7import torch
8import torch.nn as nn
9
10sys.path.insert(0, "/home/maxwelhelp/all/math2nn")
11from bench import get_dataset, make_report, make_model, sweep_baseline, evaluate, train_model, count_params
12
13SEEDS = tuple(range(8))
14# Identical union of learning rates on both systems satisfies search-space parity.
15GRID = [{"lr": 0.001}, {"lr": 0.003}, {"lr": 0.006}]
16EPOCHS = 10
17NTRAIN, NTEST, BATCH = 1000, 250, 128
18
19
20def seed_all(seed):
21 random.seed(seed)
22 np.random.seed(seed)
23 torch.manual_seed(seed)
24 if torch.cuda.is_available():
25 torch.cuda.manual_seed_all(seed)
26
27
28class ResolventRNN(nn.Module):
29 """rnn_small with a positive-real rational resolvent on its final hidden state."""
30 def __init__(self, out_dim=1, hidden=64, rank=16, eta=0.8, epsilon=1e-3):
31 super().__init__()
32 self.rnn = nn.GRU(3, hidden, batch_first=True)
33 self.head = nn.Linear(hidden, out_dim)
34 self.eta, self.epsilon = eta, epsilon
35 p = torch.diag(torch.tensor([1.] * (rank // 2) + [0.] * (rank - rank // 2)))
36 self.register_buffer("P1", p)
37 self.register_buffer("P2", torch.eye(rank) - p)
38 self.B = nn.Parameter(torch.randn(rank, hidden) / np.sqrt(hidden))
39 self.C = nn.Parameter(torch.randn(hidden, rank) / np.sqrt(rank))
40 self.log_a = nn.Parameter(torch.tensor(np.log(0.35)))
41 self.log_d = nn.Parameter(torch.tensor(np.log(0.12)))
42 self.w = nn.Parameter(torch.zeros(2, hidden))
43 self._last_ratio = 0.0
44 self._last_solve_failed = False
45
46 def forward(self, x):
47 seq = x.view(x.shape[0], -1, 3)
48 try:
49 _, h = self.rnn(seq)
50 except RuntimeError:
51 old = torch.backends.cudnn.enabled
52 torch.backends.cudnn.enabled = False
53 try:
54 _, h = self.rnn(seq)
55 finally:
56 torch.backends.cudnn.enabled = old
57 v = h[-1]
58 q = v.detach()
59 z = self.epsilon + torch.nn.functional.softplus(q @ self.w.T)
60 a = torch.nn.functional.softplus(self.log_a) + 1e-4
61 d = torch.nn.functional.softplus(self.log_d)
62 M = a * torch.eye(self.B.shape[0], device=v.device, dtype=v.dtype)[None]
63 M = M + z[:, 0, None, None] * self.P1 + z[:, 1, None, None] * self.P2
64 rhs = self.B[None] @ v[:, :, None]
65 try:
66 u = torch.linalg.solve(M, rhs).squeeze(-1)
67 H_v = d * v + torch.bmm(self.C[None].expand(v.shape[0], -1, -1), u[..., None]).squeeze(-1)
68 out = torch.linalg.solve(torch.eye(v.shape[1], device=v.device, dtype=v.dtype)[None] + self.eta * (d * torch.eye(v.shape[1], device=v.device, dtype=v.dtype)[None] + self.C[None].expand(v.shape[0], -1, -1) @ torch.linalg.solve(M, self.B[None].expand(v.shape[0], -1, -1))), v[..., None]).squeeze(-1)
69 self._last_ratio = float((out.norm(dim=1) / (v.norm(dim=1) + 1e-8)).max().detach().cpu())
70 self._last_solve_failed = False
71 except RuntimeError:
72 self._last_solve_failed = True
73 out = v
74 return self.head(out)
75
76
77def train_one(seed, idea, lr, collect=False):
78 seed_all(seed)
79 ds = get_dataset("dynamics", seed, n_train=NTRAIN, n_test=NTEST)
80 if idea:
81 net = ResolventRNN(out_dim=1, hidden=64, rank=16, eta=0.8)
82 else:
83 net = make_model("rnn_small", ds["input_shape"], ds["out_dim"])
84 trained, metric, history = train_model(net, ds, epochs=EPOCHS, lr=lr, batch=BATCH)
85 if collect and trained is not None and idea:
86 trained.eval()
87 dev = next(trained.parameters()).device
88 with torch.no_grad():
89 _ = trained(ds["xte"][:128].to(dev))
90 return metric, {"max_hidden_resolvent_ratio": trained._last_ratio, "failed_linear_solves": int(trained._last_solve_failed), "params": count_params(trained)}
91 return metric
92
93
94def main():
95 # Baseline standard-practice sweep on four seeds, then canonical full rerun.
96 base_block = sweep_baseline(lambda cfg: (lambda s: train_one(s, False, cfg["lr"])), GRID, seeds=(0, 1, 2, 3))
97 # Idea is run at the same three settings; report the best setting based on the same sweep seeds.
98 idea_sweep = []
99 for cfg in GRID:
100 r = evaluate(lambda s, lr=cfg["lr"]: train_one(s, True, lr), seeds=(0, 1, 2, 3))
101 idea_sweep.append({"cfg": cfg, "mean": r["mean"]})
102 best_cfg = min(idea_sweep, key=lambda x: x["mean"])["cfg"]
103 idea_res = evaluate(lambda s: train_one(s, True, best_cfg["lr"]), seeds=SEEDS)
104 sigs = [train_one(s, True, best_cfg["lr"], collect=True)[1] for s in SEEDS]
105 ratios = [x["max_hidden_resolvent_ratio"] for x in sigs]
106 signature = {
107 "prediction": "positive-real resolvent is nonexpansive: max ||R v||/||v|| <= 1",
108 "observed_max_ratio_mean": float(np.mean(ratios)),
109 "observed_max_ratio_max": float(np.max(ratios)),
110 "failed_linear_solves_total": int(sum(x["failed_linear_solves"] for x in sigs)),
111 "confirmed": bool(np.max(ratios) <= 1.05 and sum(x["failed_linear_solves"] for x in sigs) == 0),
112 }
113 base_block["idea_union_sweep"] = idea_sweep
114 report = make_report("dynamics", "rnn_small", base_block, idea_res, {"mechanism_signature": signature, "idea_best_cfg": best_cfg, "epochs": EPOCHS, "n_train": NTRAIN, "n_test": NTEST})
115 Path("bench_report.json").write_text(json.dumps(report, indent=2))
116 print(json.dumps(report, indent=2))
117
118
119if __name__ == "__main__":
120 main()