Entropy-Feedback Zeroth-Order Cooling / entropy_feedback_bench.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import sys, json, math, random
  2import numpy as np
  3import torch
  4import torch.nn as nn
  5
  6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  7from bench import get_dataset, make_model, sweep_baseline, make_report
  8from bench.protocol import DEFAULT_SEEDS
  9
 10
 11def seed_all(s):
 12    random.seed(s); np.random.seed(s); torch.manual_seed(s)
 13
 14
 15def controller_rate(h, alpha, eps, hc, delta):
 16    z = max(-60., min(60., (h-hc)/delta))
 17    b = 1./(1.+math.exp(-z))
 18    return alpha*(eps+(1.-eps)*b)
 19
 20
 21def flat_state(model):
 22    return torch.cat([p.detach().reshape(-1) for p in model.parameters()])
 23
 24
 25def assign_state(model, v):
 26    k=0
 27    with torch.no_grad():
 28        for p in model.parameters():
 29            n=p.numel(); p.copy_(v[k:k+n].view_as(p)); k += n
 30
 31
 32def predict(model, theta, x):
 33    # functional_call avoids mutating the center while scoring each candidate.
 34    from torch.func import functional_call
 35    params={n:v for (n,_),v in zip(model.named_parameters(), split_state(model, theta))}
 36    return functional_call(model, params, (x,))
 37
 38
 39def split_state(model, theta):
 40    out=[]; k=0
 41    for p in model.parameters():
 42        n=p.numel(); out.append(theta[k:k+n].view_as(p)); k += n
 43    return out
 44
 45
 46def run(seed, mode, cfg, signature=False):
 47    seed_all(seed)
 48    ds=get_dataset('tabular', seed, n_train=400, n_test=400)
 49    device=torch.device('cuda' if torch.cuda.is_available() else 'cpu')
 50    try:
 51        model=make_model('mlp_tiny', tuple(ds['xtr'].shape[1:]), 1).to(device)
 52        x=ds['xtr'].to(device); y=ds['ytr'].to(device).view(-1,1)
 53        xt=ds['xte'].to(device); yt=ds['yte'].to(device).view(-1,1)
 54        theta=flat_state(model)
 55        # Same number of candidate evaluations and perturbation scale for both methods.
 56        N=16; sigma=cfg['sigma']; tau=cfg['tau0']; T=cfg['steps']
 57        alpha=cfg['alpha']; eps=cfg['epsilon']; hc=cfg['hc']; delta=cfg['delta']
 58        hs=[]; rates=[]; ess=[]; losses=[]
 59        for _ in range(T):
 60            noise=torch.randn((N,theta.numel()), device=device)
 61            cand=theta[None,:] + sigma*noise
 62            vals=[]
 63            for i in range(N):
 64                pred=predict(model,cand[i],x)
 65                vals.append(((pred-y)**2).mean())
 66            lv=torch.stack(vals)
 67            use_tau=tau
 68            z=-(lv-lv.min())/max(use_tau,1e-8)
 69            w=torch.softmax(torch.clamp(z,-80,0),dim=0)
 70            H=-(w*torch.log(w.clamp_min(1e-12))).sum()
 71            h=float(H.item()/math.log(N)); e=float(torch.exp(H).item())
 72            theta=(w[:,None]*cand).sum(0)
 73            if mode=='feedback': r=controller_rate(h,alpha,eps,hc,delta)
 74            else: r=alpha
 75            tau *= math.exp(-r)
 76            hs.append(h); rates.append(r); ess.append(e); losses.append(float(((predict(model,theta,x)-y)**2).item()))
 77        assign_state(model,theta)
 78        with torch.no_grad(): metric=float(((model(xt)-yt)**2).mean().item())
 79        if signature:
 80            high=np.mean([r for h,r in zip(hs,rates) if h>hc+2*delta]) if any(h>hc+2*delta for h in hs) else float('nan')
 81            low=np.mean([r for h,r in zip(hs,rates) if h<hc-2*delta]) if any(h<hc-2*delta for h in hs) else float('nan')
 82            pred_high=alpha; pred_low=alpha*eps
 83            return {'metric':metric,'signature':{'predicted_transition_ess':N**hc,'observed_ess_at_closest_hc':ess[int(np.argmin(np.abs(np.asarray(hs)-hc)))], 'predicted_high_entropy_rate':pred_high,'observed_high_entropy_rate':high,'predicted_low_entropy_rate':pred_low,'observed_low_entropy_rate':low,'mean_entropy':float(np.mean(hs)),'confirmed':bool(np.isfinite(high) and np.isfinite(low) and abs(high-pred_high)/pred_high<.25 and abs(low-pred_low)/pred_low<.5)}}
 84        return {'metric':metric}
 85    except RuntimeError:
 86        # Shared-space fallback required for constrained CUDA environments.
 87        torch.cuda.empty_cache() if torch.cuda.is_available() else None
 88        return run_cpu(seed,mode,cfg,signature)
 89
 90
 91def run_cpu(seed,mode,cfg,signature=False):
 92    old=torch.cuda.is_available
 93    # Re-run deterministically on CPU by temporarily selecting tensors explicitly.
 94    seed_all(seed); ds=get_dataset('tabular',seed,n_train=400,n_test=400)
 95    model=make_model('mlp_tiny',tuple(ds['xtr'].shape[1:]),1)
 96    x,y=ds['xtr'],ds['ytr'].view(-1,1); xt,yt=ds['xte'],ds['yte'].view(-1,1)
 97    theta=flat_state(model); N=16; tau=cfg['tau0']; hs=[]; rates=[]; ess=[]
 98    for _ in range(cfg['steps']):
 99        cand=theta[None,:]+cfg['sigma']*torch.randn((N,theta.numel())); lv=torch.stack([((predict(model,cand[i],x)-y)**2).mean() for i in range(N)])
100        w=torch.softmax(torch.clamp(-(lv-lv.min())/max(tau,1e-8),-80,0),0); H=-(w*torch.log(w.clamp_min(1e-12))).sum(); h=float(H/math.log(N)); r=controller_rate(h,cfg['alpha'],cfg['epsilon'],cfg['hc'],cfg['delta']) if mode=='feedback' else cfg['alpha']; theta=(w[:,None]*cand).sum(0); tau*=math.exp(-r); hs.append(h); rates.append(r); ess.append(float(torch.exp(H)))
101    assign_state(model,theta); metric=float(((model(xt)-yt)**2).mean())
102    if not signature:return {'metric':metric}
103    hi=np.mean([r for h,r in zip(hs,rates) if h>cfg['hc']+2*cfg['delta']]) if any(h>cfg['hc']+2*cfg['delta'] for h in hs) else float('nan'); lo=np.mean([r for h,r in zip(hs,rates) if h<cfg['hc']-2*cfg['delta']]) if any(h<cfg['hc']-2*cfg['delta'] for h in hs) else float('nan')
104    return {'metric':metric,'signature':{'predicted_transition_ess':N**cfg['hc'],'observed_ess_at_closest_hc':ess[int(np.argmin(np.abs(np.asarray(hs)-cfg['hc'])))],'predicted_high_entropy_rate':cfg['alpha'],'observed_high_entropy_rate':hi,'predicted_low_entropy_rate':cfg['alpha']*cfg['epsilon'],'observed_low_entropy_rate':lo,'mean_entropy':float(np.mean(hs)),'confirmed':bool(np.isfinite(hi) and np.isfinite(lo))}}
105
106
107def cfg(lr):
108    return {'tau0':1.0,'alpha':0.05,'epsilon':0.02,'hc':0.5,'delta':0.05,'sigma':0.03,'steps':12,'lr':lr}
109
110
111def main():
112    grid=[{"lr":v} for v in (0.001,0.003,0.01)]
113    def base_factory(c):
114        return lambda seed: run(seed, "fixed", cfg(c["lr"]))["metric"]
115    base=sweep_baseline(base_factory, grid, seeds=(0,1,2,3))
116    best=base["best_cfg"]
117    idea_vals=[run(s, "feedback", cfg(best["lr"]))["metric"] for s in DEFAULT_SEEDS]
118    idea_res={"mean":float(np.mean(idea_vals)), "std":float(np.std(idea_vals)), "per_seed":idea_vals, "n":len(idea_vals), "config":best}
119    nearby={str(v):[run(s,"feedback",cfg(v))["metric"] for s in DEFAULT_SEEDS] for v in (0.001,0.003,0.01)}
120    sig=run(0, "feedback", cfg(best["lr"]), signature=True)["signature"]
121    report=make_report("tabular", "mlp_tiny", base, idea_res, extra=sig)
122    report["idea_nearby_settings"]=nearby
123    report["protocol_notes"]=("Optimizer ideas are matched to tabular Friedman regression; both systems use identical mlp_tiny, population, perturbation scale, steps, data, and paired seeds. Baseline fixed exponential cooling was swept over the shared lr union.")
124    report["custom_track"]=None
125    with open("bench_report.json","w") as f: json.dump(report,f,indent=2)
126    print(json.dumps(report,indent=2))
127
128if __name__=='__main__': main()