STL-Robust Mixture-of-Experts Gating / bench_stl_moe.py

Unverified

Raw ⬇ ZIP
  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()