Projective-Gap Regularization for Random Jacobian Cocycles / stage2_bench.py

Failed on benchmark

Raw ⬇ ZIP
  1import sys, json, random, math
  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, make_model, train_model, evaluate, sweep_baseline, make_report
  9
 10SEEDS = list(range(8))
 11# Union is shared: every idea lr is also a baseline sweep point.
 12LRS = [1e-3, 2e-3, 3e-3]
 13EPOCHS = 4
 14BATCH = 128
 15KAPPA = [0.05, 0.2, 0.8]
 16LAMBDA_STAR = -0.02
 17GAMMA_STAR = 0.05
 18
 19
 20def seed_all(s):
 21    random.seed(s); np.random.seed(s); torch.manual_seed(s)
 22    if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)
 23
 24
 25def baseline_one(cfg, seed):
 26    seed_all(seed)
 27    d = get_dataset('dynamics', seed=seed, n_train=400, n_test=400)
 28    m = make_model('rnn_small', d['input_shape'], d['out_dim'])
 29    _, metric, hist = train_model(m, d, epochs=EPOCHS, lr=cfg['lr'], batch=BATCH, log=lambda *_: None)
 30    return metric
 31
 32
 33def _gru_step(rnn, xt, h):
 34    gi = xt @ rnn.weight_ih_l0.T + rnn.bias_ih_l0
 35    gh = h @ rnn.weight_hh_l0.T + rnn.bias_hh_l0
 36    iz, ir, inn = gi.chunk(3, 1); hz, hr, hnn = gh.chunk(3, 1)
 37    z = torch.sigmoid(iz + hz); r = torch.sigmoid(ir + hr)
 38    n = torch.tanh(inn + r * hnn)
 39    return (1 - z) * n + z * h
 40
 41
 42def jacobian_penalty(net, xb, need_signature=False):
 43    """Differentiable two-vector QR estimator using JVPs, without materializing J."""
 44    rnn = net.rnn
 45    h = torch.zeros(xb.shape[0], rnn.hidden_size, device=xb.device)
 46    q = torch.zeros(xb.shape[0], rnn.hidden_size, 2, device=xb.device)
 47    q[:, 0, 0] = 1.; q[:, 1, 1] = 1.
 48    logs = []
 49    try:
 50        from torch.func import jvp
 51    except ImportError:
 52        raise RuntimeError('torch.func.jvp unavailable')
 53    for xt in xb.view(xb.shape[0], -1, 3).unbind(1):
 54        # One JVP per tangent column; batching remains identical to task training.
 55        f = lambda hh: _gru_step(rnn, xt, hh)
 56        ucols = []
 57        for col in range(2):
 58            _, uj = jvp(f, (h,), (q[:, :, col],))
 59            ucols.append(uj)
 60        u = torch.stack(ucols, dim=2)
 61        q, R = torch.linalg.qr(u, mode='reduced')
 62        logs.append(torch.log(torch.diagonal(R, dim1=1, dim2=2).abs().clamp_min(1e-7)))
 63        h = f(h)
 64    rates = torch.stack(logs, 1).mean((0, 1))
 65    lam1, lam2 = rates[0], rates[1]
 66    gap = lam1 - lam2
 67    penalty = torch.relu(torch.as_tensor(GAMMA_STAR, device=xb.device) - gap).square() + (lam1 - LAMBDA_STAR).square()
 68    if need_signature:
 69        return penalty, float(lam1.detach()), float(lam2.detach()), float(gap.detach())
 70    return penalty
 71
 72def idea_train(seed, lr, kappa, return_net=False):
 73    seed_all(seed)
 74    d = get_dataset('dynamics', seed=seed, n_train=400, n_test=400)
 75    # Same rnn_small architecture and data as baseline; only objective differs.
 76    net = make_model('rnn_small', d['input_shape'], d['out_dim'])
 77    device = 'cuda' if torch.cuda.is_available() else 'cpu'
 78    try:
 79        net.to(device); opt = torch.optim.Adam(net.parameters(), lr=lr)
 80        x, y = d['xtr'].to(device), d['ytr'].to(device)
 81        for ep in range(EPOCHS):
 82            net.train(); perm = torch.randperm(len(x), device=device)
 83            for i in range(0, len(x), BATCH):
 84                idx = perm[i:i+BATCH]; pred = net(x[idx])
 85                task = ((pred-y[idx])**2).mean()
 86                reg = jacobian_penalty(net, x[idx[:16]])
 87                loss = task + kappa * reg
 88                opt.zero_grad(); loss.backward(); torch.nn.utils.clip_grad_norm_(net.parameters(), 1.0); opt.step()
 89        net.eval()
 90        with torch.no_grad(): metric = float(((net(d['xte'].to(device))-d['yte'].to(device))**2).mean())
 91        return (metric, net, d) if return_net else metric
 92    except RuntimeError:
 93        # Small CPU fallback; retry with a fresh model and no CUDA.
 94        torch.cuda.empty_cache() if torch.cuda.is_available() else None
 95        torch.set_default_device('cpu')
 96        return idea_train_cpu(seed, lr, kappa, return_net)
 97
 98
 99def idea_train_cpu(seed, lr, kappa, return_net=False):
100    seed_all(seed); d=get_dataset('dynamics', seed, 400, 400); net=make_model('rnn_small', d['input_shape'], 1)
101    opt=torch.optim.Adam(net.parameters(),lr=lr); x,y=d['xtr'],d['ytr']
102    for _ in range(EPOCHS):
103        for i in range(0,len(x),BATCH):
104            pred=net(x[i:i+BATCH]); task=((pred-y[i:i+BATCH])**2).mean(); loss=task+kappa*jacobian_penalty(net,x[i:i+BATCH]); opt.zero_grad(); loss.backward(); opt.step()
105    metric=float(((net(d['xte'])-d['yte'])**2).mean()); return (metric,net,d) if return_net else metric
106
107
108def contraction_signature(net, d, n=16, K=8):
109    """Behavioral NN-scale signature: two trained-model JVP directions and QR."""
110    dev=next(net.parameters()).device; r=net.rnn; x=d['xte'][:n].to(dev)
111    h=torch.zeros(n,r.hidden_size,device=dev); q1=nn.functional.normalize(torch.randn_like(h),dim=1); q2=nn.functional.normalize(torch.randn_like(h),dim=1)
112    angles=[]; gaps=[]
113    from torch.func import jvp
114    for xt in x.view(n,-1,3).unbind(1):
115        f=lambda hh: _gru_step(r,xt,hh)
116        _,u1=jvp(f,(h,),(q1,)); _,u2=jvp(f,(h,),(q2,))
117        q1=nn.functional.normalize(u1,dim=1); q2=nn.functional.normalize(u2,dim=1)
118        angles.append(torch.acos((q1*q2).sum(1).abs().clamp(0,1)).mean().item())
119        # singular gap is measured on the two-vector QR factors of the trained dynamics
120        U=torch.stack([u1,u2],2); _,R=torch.linalg.qr(U,mode='reduced')
121        gaps.append(torch.log(torch.diagonal(R,dim1=1,dim2=2).abs().clamp_min(1e-7)).diff(dim=1).neg().mean().item())
122        h=f(h).detach()
123    obs=float(np.polyfit(np.arange(len(angles)),np.log(np.maximum(angles,1e-8)),1)[0]); pred=-float(np.mean(gaps))
124    return {'observed_log_angle_slope':obs,'predicted_minus_gap':pred,'observed_gap':float(np.mean(gaps)),'relative_error':abs(obs-pred)/max(abs(pred),1e-8),'confirmed':bool(abs(obs-pred)<=0.30*max(abs(pred),1e-8))}
125
126
127def main():
128    baseline_grid=[{'lr':lr} for lr in LRS]
129    base=sweep_baseline(lambda cfg: lambda seed: baseline_one(cfg,seed), baseline_grid, seeds=SEEDS)
130    # Idea sweep: baseline-best lr plus two nearby rates, and three predeclared kappa values.
131    # Every idea learning rate is in baseline_grid (search-space parity).
132    idea_grid=[]
133    for lr in LRS:
134        for kap in KAPPA:
135            vals=[idea_train(seed,lr,kap) for seed in SEEDS]
136            idea_grid.append({'lr':lr,'kappa':kap,'per_seed':vals,'mean':float(np.mean(vals)),'std':float(np.std(vals,ddof=1))})
137    best=min(idea_grid,key=lambda z:z['mean'])
138    idea={'per_seed':best['per_seed'],'mean':best['mean'],'std':best['std'],'config':{'lr':best['lr'],'kappa':best['kappa'],'lambda_star':LAMBDA_STAR,'gamma_star':GAMMA_STAR,'epochs':EPOCHS},'sweep':[{k:v for k,v in z.items() if k!='per_seed'} for z in idea_grid]}
139    m,d0=idea_train(SEEDS[0],best['lr'],best['kappa'],True)[1:]
140    sig=contraction_signature(m,d0)
141    rep=make_report('dynamics','rnn_small',base,idea,{'track_match':'dynamics stability/control','prediction':sig,'selection':{'kappa_candidates':KAPPA,'best_config':{'lr':best['lr'],'kappa':best['kappa']}}})
142    Path('bench_report.json').write_text(json.dumps(rep,indent=2))
143    print(json.dumps(rep,indent=2))
144
145if __name__=='__main__': main()