STL-Robust Mixture-of-Experts Gating / bench_stl_moe.py
Unverified
1import sys, json, random
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 get_dataset, train_model, evaluate, sweep_baseline, make_report
9
10SEED = 1448
11
12def seed_all(seed):
13 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
14 if torch.cuda.is_available():
15 try: torch.cuda.manual_seed_all(seed)
16 except Exception: pass
17
18class RobustMoE(nn.Module):
19 """Matched recurrent MoE. beta=0 is ordinary learned softmax routing."""
20 def __init__(self, beta=0.0, hidden=64, experts=4):
21 super().__init__()
22 self.beta = float(beta); self.experts = experts
23 self.rnn = nn.GRU(3, hidden, batch_first=True)
24 self.router = nn.Linear(hidden, experts)
25 self.transitions = nn.Parameter(torch.zeros(experts, experts))
26 self.heads = nn.ModuleList([nn.Linear(hidden, 1) for _ in range(experts)])
27 self.register_buffer('expert_bias', torch.tensor([-0.15, -0.05, 0.05, 0.15]))
28
29 def _robustness(self, last, pred):
30 # Smooth STL-style G(position >= -1.5 AND |omega| <= 3).
31 # A short constant-velocity rollout makes the temporal minimum differentiable.
32 th, om = last[:, 0], last[:, 1]
33 hs = torch.arange(1, 7, device=last.device, dtype=last.dtype)[None, :]
34 future_th = pred[:, :, None] + 0.05 * hs * om[:, None, None]
35 r_pos = future_th + 1.5
36 r_vel = 3.0 - torch.abs(om[:, None, None].expand_as(future_th))
37 atoms = torch.minimum(r_pos, r_vel)
38 tau = 0.12
39 return -tau * torch.logsumexp(-atoms / tau, dim=2)
40
41 def forward(self, x):
42 seq = x.view(x.shape[0], -1, 3)
43 _, h = self.rnn(seq)
44 h = h[-1]
45 pred = torch.cat([head(h) for head in self.heads], dim=1)
46 rho = self._robustness(seq[:, -1, :], pred)
47 pi = torch.softmax(self.router(h), dim=1)
48 # T[j,i] = softmax_i(A[j,i] + beta*rho_i), then pi' = pi @ T.
49 T = torch.softmax(self.transitions[None, :, :] + self.beta * rho[:, None, :], dim=2)
50 post = torch.bmm(pi[:, None, :], T).squeeze(1)
51 return (post * pred).sum(1, keepdim=True)
52
53
54def make_train(cfg):
55 def fn(seed):
56 seed_all(seed)
57 d = get_dataset('dynamics', seed, n_train=800, n_test=400)
58 net = RobustMoE(beta=cfg.get('beta', 0.0))
59 _, metric, _ = train_model(net, d, epochs=12, lr=cfg['lr'], batch=128,
60 weight_decay=0.0, log=lambda *a, **k: None)
61 return metric
62 return fn
63
64
65def signature(seed, beta, lr):
66 seed_all(seed); d = get_dataset('dynamics', seed, n_train=800, n_test=400)
67 net = RobustMoE(beta=beta)
68 net, _, _ = train_model(net, d, epochs=12, lr=lr, batch=128, log=lambda *a, **k: None)
69 net = net.cpu(); net.eval(); device = torch.device('cpu'); x = d['xte'][:256]
70 with torch.no_grad():
71 seq=x.view(x.shape[0],-1,3); _,h=net.rnn(seq); h=h[-1]
72 pred=torch.cat([q(h) for q in net.heads],1)
73 rho=net._robustness(seq[:,-1,:],pred)
74 pi=torch.softmax(net.router(h),1)
75 T=torch.softmax(net.transitions[None,:,:]+beta*rho[:,None,:],2)
76 post=torch.bmm(pi[:,None,:],T).squeeze(1)
77 # Across samples, test the claimed positive robustness dependence after
78 # controlling for the learned transition-logit difference.
79 i,j=0,1
80 observed=(torch.log(post[:,i]+1e-8)-torch.log(post[:,j]+1e-8)).cpu().numpy()
81 delta=(rho[:,i]-rho[:,j]).cpu().numpy()
82 slope=float(np.polyfit(delta, observed, 1)[0])
83 predicted=float(beta)
84 corr=float(np.corrcoef(delta, observed)[0,1])
85 return {'beta': beta, 'predicted_slope': predicted,
86 'observed_slope': slope, 'robustness_logodds_correlation': corr,
87 'n_samples': len(delta), 'confirmed': bool(beta > 0 and slope > 0 and corr > 0.05)}
88
89
90def main():
91 # Cheap algebra sanity check is independent of trained scoring.
92 A=np.array([0.4,-0.3,0.1,-0.2]); r=np.array([.2,-.75,.55,-.1]); b=np.linspace(0,4,9)
93 logits=A[None,:]+b[:,None]*r[None,:]
94 odds=logits[:,2]-logits[:,1]
95 math_check={'predicted_slope':float(r[2]-r[1]),
96 'observed_slope':float(np.polyfit(b,odds,1)[0]),
97 'max_identity_error':float(np.max(np.abs(odds-(A[2]-A[1]+b*(r[2]-r[1])))))}
98 # Baseline sweep includes every lr used in the shared search space.
99 grid=[{'lr':v} for v in (1e-3,3e-3,6e-3)]
100 base=sweep_baseline(make_train, grid)
101 best_lr=base['best_cfg']['lr']
102 # Same-sized idea sweep, with beta=0 retained as matched control.
103 idea_grid=[{'lr':best_lr,'beta':v} for v in (0.0,0.75,1.5)]
104 tried=[]
105 for cfg in idea_grid:
106 r=evaluate(make_train(cfg), seeds=(0,1,2,3))
107 tried.append({'cfg':cfg,'mean':r['mean']})
108 best_beta=min(tried,key=lambda z:z['mean'])['cfg']['beta']
109 idea=evaluate(make_train({'lr':best_lr,'beta':best_beta}))
110 # Include sweep evidence without changing the canonical report schema.
111 extra={'track_choice':'dynamics: pendulum control and long-term stability structure',
112 'math_check':math_check,
113 'idea_sweep':tried,
114 'mechanism_signature':signature(0,best_beta,best_lr)}
115 report=make_report('dynamics','rnn_small',base,idea,extra)
116 report['idea']['selected_cfg']={'lr':best_lr,'beta':best_beta}
117 Path('bench_report.json').write_text(json.dumps(report,indent=2))
118 print(json.dumps(report,indent=2))
119
120if __name__=='__main__': main()