Exact-Jacobian Flow Controller / bench_run.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import sys, json, random, math
  2from pathlib import Path
  3import numpy as np
  4import torch
  5from torch import nn
  6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  7from bench import get_dataset, train_model, evaluate, sweep_baseline, make_report
  8
  9SEEDS = tuple(range(8))
 10SWEEP_SEEDS = (0, 1, 2, 3)
 11EPOCHS = 12
 12NTRAIN, NTEST = 800, 300
 13
 14class Encoder(nn.Module):
 15    def __init__(self, hidden=64):
 16        super().__init__()
 17        self.rnn = nn.GRU(3, hidden, batch_first=True)
 18        self._no_cudnn = False
 19    def forward(self, x):
 20        seq = x.view(x.shape[0], -1, 3)
 21        try:
 22            _, h = self.rnn(seq)
 23        except RuntimeError:
 24            self._no_cudnn = True
 25        if self._no_cudnn:
 26            old = torch.backends.cudnn.enabled; torch.backends.cudnn.enabled = False
 27            try: _, h = self.rnn(seq)
 28            finally: torch.backends.cudnn.enabled = old
 29        return h[-1]
 30
 31class Baseline(nn.Module):
 32    def __init__(self):
 33        super().__init__(); self.enc = Encoder()
 34        self.head = nn.Sequential(nn.Linear(64, 32), nn.ReLU(), nn.Linear(32, 1))
 35    def forward(self, x): return self.head(self.enc(x))
 36
 37class Coupling(nn.Module):
 38    def __init__(self, cond, mask, clamp=1.0):
 39        super().__init__(); self.register_buffer('mask', mask); self.clamp=clamp
 40        self.net=nn.Sequential(nn.Linear(2+cond,32),nn.Tanh(),nn.Linear(32,32),nn.Tanh(),nn.Linear(32,4))
 41        nn.init.zeros_(self.net[-1].weight); nn.init.zeros_(self.net[-1].bias)
 42    def params(self, x, c):
 43        o=self.net(torch.cat([x*self.mask,c],-1)); s,b=o[...,:2],o[...,2:]
 44        m=1-self.mask; return self.clamp*torch.tanh(s)*m,b*m
 45    def forward(self,x,c):
 46        s,b=self.params(x,c); m=1-self.mask
 47        return x*self.mask+m*(x*torch.exp(s)+b), s.sum(-1)
 48    def inverse(self,y,c):
 49        s,b=self.params(y,c); m=1-self.mask
 50        return y*self.mask+m*(y-b)*torch.exp(-s), -s.sum(-1)
 51
 52class ExactJacobianController(nn.Module):
 53    def __init__(self, K=4):
 54        super().__init__(); self.enc=Encoder(); self.K=K
 55        self.layers=nn.ModuleList([Coupling(64, torch.tensor([float((k+1)%2),float(k%2)])) for k in range(K)])
 56        self.out=nn.Linear(64,2)
 57        nn.init.zeros_(self.out.weight); nn.init.zeros_(self.out.bias)
 58    def flow(self,z,c):
 59        ld=torch.zeros(z.shape[0],device=z.device)
 60        for layer in self.layers: z,a=layer(z,c); ld=ld+a
 61        return z,ld
 62    def inv(self,y,c):
 63        ld=torch.zeros(y.shape[0],device=y.device)
 64        for layer in reversed(self.layers): y,a=layer.inverse(y,c); ld=ld+a
 65        return y,ld
 66    def forward(self,x):
 67        c=self.enc(x); z=self.out(c); y, _ = self.flow(z,c)
 68        return y[:, :1]
 69    def reconstruction_error(self, x):
 70        with torch.no_grad():
 71            c=self.enc(x); z=self.out(c); y,ld=self.flow(z,c); zr,_=self.inv(y,c)
 72            return float((zr-z).abs().max()), float(ld.abs().mean())
 73    def jacobian_signature(self,x):
 74        c=self.enc(x[:1]).detach(); z=self.out(c).detach().requires_grad_(True)
 75        def fn(u):
 76            return self.flow(u.unsqueeze(0),c)[0][0]
 77        J=torch.autograd.functional.jacobian(fn,z[0])
 78        sign, actual=torch.linalg.slogdet(J)
 79        with torch.no_grad(): _, analytic=self.flow(z.detach(),c)
 80        return float(abs(actual-analytic[0]).item()), float(sign.item())
 81
 82def seed_all(s):
 83    random.seed(s); np.random.seed(s); torch.manual_seed(s)
 84    if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)
 85
 86def ds(seed): return get_dataset('dynamics', seed, n_train=NTRAIN, n_test=NTEST)
 87
 88def train_one(kind, cfg, seed, capture=False):
 89    seed_all(seed)
 90    model = Baseline() if kind=='baseline' else ExactJacobianController(K=cfg.get('K',4))
 91    net, metric, hist = train_model(model, ds(seed), epochs=EPOCHS, lr=cfg['lr'], batch=128, log=lambda *a: None)
 92    if net is None: return float('inf')
 93    if capture: return float(metric), net
 94    return float(metric)
 95
 96def main():
 97    # Shared union: every learning rate considered for the idea is swept for baseline.
 98    grid=[{'lr':1e-3},{'lr':3e-3},{'lr':1e-2}]
 99    base=sweep_baseline(lambda cfg: (lambda seed: train_one('baseline',cfg,seed)), grid, seeds=SWEEP_SEEDS)
100    idea_cfgs=[{'lr':base['best_cfg']['lr'],'K':4}, {'lr':1e-3,'K':4}, {'lr':1e-2,'K':4}]
101    # Select idea config on the same sweep seeds, then evaluate its selected config fully.
102    idea_sweep=[]
103    for cfg in idea_cfgs:
104        r=evaluate(lambda s: train_one('idea',cfg,s), SWEEP_SEEDS)
105        idea_sweep.append({'cfg':cfg,'mean':r['mean']})
106    best_idea=min(idea_sweep,key=lambda q:q['mean'])['cfg']
107    idea=evaluate(lambda s: train_one('idea',best_idea,s), SEEDS)
108    # Re-test trained systems for a behavioral mechanism signature.
109    metric, model=train_one('idea',best_idea,0,capture=True)
110    d=ds(0); device=next(model.parameters()).device; xb=d['xte'][:16].to(device)
111    recon, ldmean=model.reconstruction_error(xb)
112    jacerr, sign=model.jacobian_signature(xb)
113    signature={'prediction':'analytic triangular inverse and logdet remain exact after training',
114               'reconstruction_max_abs':recon,'observed_mean_abs_logdet':ldmean,
115               'observed_jacobian_logdet_abs_error':jacerr,'jacobian_sign':sign,
116               'confirmed': bool(recon < 2e-5 and jacerr < 2e-5 and sign > 0)}
117    base['idea_union_sweep']=idea_sweep
118    report=make_report('dynamics','rnn_small',base,idea,extra=signature)
119    report['protocol_notes']={'epochs':EPOCHS,'n_train':NTRAIN,'n_test':NTEST,'idea_configs':idea_cfgs,
120      'structural_match':'controlled pendulum multi-step rollout; flow is trained end-to-end with shared GRU encoder'}
121    Path('bench_report.json').write_text(json.dumps(report,indent=2))
122    print(json.dumps(report,indent=2))
123if __name__=='__main__': main()