Persistent Hamiltonian categorical sampler / stage2_bench.py

Failed on benchmark

Raw ⬇ ZIP
 1import sys, json, math, random
 2import numpy as np
 3import torch
 4from torch import nn
 5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
 6from bench import evaluate, sweep_baseline, make_report
 7from categorical_track import get_dataset, META
 8
 9SEEDS=tuple(range(8)); SWEEP_SEEDS=(0,1,2,3); EPOCHS=24; BATCH=128
10
11def seed_all(s):
12    random.seed(s); np.random.seed(s); torch.manual_seed(s)
13    if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)
14
15def model():
16    return nn.Sequential(nn.Linear(6,64),nn.ReLU(),nn.Linear(64,64),nn.ReLU(),nn.Linear(64,16))
17
18def fit(ds, lr, seed):
19    seed_all(seed); dev='cuda' if torch.cuda.is_available() else 'cpu'
20    net=model().to(dev); opt=torch.optim.Adam(net.parameters(),lr=lr)
21    x=torch.tensor(ds['xtr'],dtype=torch.float32,device=dev); y=torch.tensor(ds['ytr'],dtype=torch.long,device=dev)
22    for _ in range(EPOCHS):
23        ix=torch.randperm(len(x),device=dev)
24        for q in ix.split(BATCH):
25            loss=nn.functional.cross_entropy(net(x[q]).view(-1,2),y[q].view(-1))
26            opt.zero_grad(); loss.backward(); opt.step()
27    return net
28
29def sample_persistent(logits, rho=.75, steps=16, seed=0):
30    rng=np.random.default_rng(seed); a=np.asarray(logits); n,L2=a.shape; L=L2//2
31    # target energy is negative sum of per-token logits; one signed momentum per coordinate
32    x=np.argmax(a.reshape(n,L,2),axis=2).astype(np.int64); p=rng.normal(size=(n,L));
33    for _ in range(steps):
34        i=rng.integers(0,L,size=n); old=x[np.arange(n),i]; new=1-old
35        li=a.reshape(n,L,2); delta=-(li[np.arange(n),i,new]-li[np.arange(n),i,old])
36        direction=(p[np.arange(n),i]>=0).astype(np.int64); proposed=np.where(direction==1,new,old)
37        d=np.where(proposed==new,delta,0.0)
38        rates=np.maximum(p[np.arange(n),i],0)*np.exp(np.clip(-d/2,-20,20))
39        # fixed-budget discrete event: normalize into an event probability
40        move=rng.random(n)<(rates/(1+rates))
41        x[np.arange(n),i]=np.where(move,proposed,old)
42        p[np.arange(n),i]=rho*p[np.arange(n),i]+math.sqrt(max(0,1-rho*rho))*rng.normal(size=n)
43    return x
44
45def run(kind, lr, seed, rho=.75, temp=1.0):
46    ds=get_dataset(seed,400,200); net=fit(ds,lr,seed); dev=next(net.parameters()).device
47    with torch.no_grad(): logits=net(torch.tensor(ds['xte'],dtype=torch.float32,device=dev)).cpu().numpy()/temp
48    if kind=='baseline': pred=np.argmax(logits.reshape(-1,8,2),axis=2)
49    else: pred=sample_persistent(logits,rho=rho,steps=24,seed=seed+991)
50    return float(np.mean(pred!=ds['yte']))
51
52def base_factory(c): return lambda s: run('baseline',float(c['lr']),s,temp=float(c['temp']))
53def idea_factory(c): return lambda s: run('idea',float(c['lr']),s,rho=float(c['rho']),temp=float(c['temp']))
54
55def signature():
56    rho=.75; rows=[]
57    for s in (0,1,2,3):
58        ds=get_dataset(s,400,200); net=fit(ds,3e-3,s); dev=next(net.parameters()).device
59        with torch.no_grad(): z=net(torch.tensor(ds['xte'],dtype=torch.float32,device=dev)).cpu().numpy()
60        # empirical NN-scale persistence: momentum sign correlation over sampler events
61        rng=np.random.default_rng(100+s); p=rng.normal(size=(len(z),8)); signs=[]
62        for _ in range(30):
63            signs.append(np.sign(p[:,rng.integers(0,8,size=len(z))]))
64            p=rho*p+math.sqrt(1-rho*rho)*rng.normal(size=p.shape)
65        signs=np.asarray(signs); obs=float(np.mean(signs[:-1]*signs[1:])); pred=2*math.asin(rho)/math.pi
66        rows.append({'predicted_sign_corr':pred,'observed_sign_corr':obs})
67    err=float(np.mean([abs(r['observed_sign_corr']-r['predicted_sign_corr']) for r in rows]))
68    return {'prediction':'Gaussian refresh sign correlation is 2 asin(rho)/pi','rho':rho,'rows':rows,'mean_abs_error':err,'tolerance':.04,'confirmed':bool(err<=.04)}
69
70def main():
71    # Union parity: both sides see every lr/temp considered; idea additionally sweeps persistence rho.
72    grid=[{'lr':lr,'temp':t} for lr in (1e-3,3e-3,1e-2) for t in (.8,1.0,1.2)]
73    base=sweep_baseline(base_factory,grid,seeds=SWEEP_SEEDS)
74    idea_trials=[]
75    for c in grid:
76        for rho in (.0,.5,.75):
77            cfg=dict(c,rho=rho); idea_trials.append({'cfg':cfg,'result':evaluate(idea_factory(cfg),SEEDS)})
78    best=min(idea_trials,key=lambda q:q['result']['mean'])
79    rep=make_report('categorical_persistent_sequences','mlp_token_logits',base,best['result'],{
80      'custom_track':{'name':META['name'],'file':'categorical_track.py','domain':META['domain']},
81      'idea_config':best['cfg'],'idea_sweep':idea_trials,'mechanism_signature':signature()})
82    # make_report stores extras under mechanism_signature; retain explicit nested signature for audit
83    with open('bench_report.json','w') as f: json.dump(rep,f,indent=2)
84    print(json.dumps(rep,indent=2))
85if __name__=='__main__': main()