Coordinate Path-Integral Joint Gibbs Policy / run_bench.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import sys, json, random
  2from pathlib import Path
  3import numpy as np
  4import torch
  5import torch.nn as nn
  6import torch.nn.functional as F
  7sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  8from bench import sweep_baseline, evaluate, make_report
  9from joint_action_values import get_dataset
 10
 11SEEDS = tuple(range(8))
 12A = torch.linspace(-1., 1., 21)
 13W = np.array([[.8, -.4, .3], [-.3, .7, .5]], np.float32)
 14
 15def seed_all(s):
 16    random.seed(s); np.random.seed(s); torch.manual_seed(s)
 17    if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)
 18
 19def data(seed, n=400):
 20    rng = np.random.default_rng(seed)
 21    state = rng.normal(size=(n, 3)).astype(np.float32)
 22    target = np.tanh(state @ W.T + .12 * rng.normal(size=(n, 2))).astype(np.float32)
 23    action = rng.uniform(-1., 1., size=(n, 2)).astype(np.float32)
 24    q = np.stack([-(action[:, 0]-target[:, 0])**2 + .75*action[:, 0]*action[:, 1],
 25                  -(action[:, 1]-target[:, 1])**2 - .25*action[:, 0]*action[:, 1]], 1).astype(np.float32)
 26    return state, action, q, target
 27
 28class QNet(nn.Module):
 29    def __init__(self):
 30        super().__init__()
 31        self.net = nn.Sequential(nn.Linear(5, 64), nn.Tanh(), nn.Linear(64, 64), nn.Tanh(), nn.Linear(64, 2))
 32    def forward(self, s, a): return self.net(torch.cat([s, a], 1))
 33
 34def path_energy(model, s, u, K=5, create_graph=False):
 35    """Fixed 1->2 coordinate path, midpoint quadrature, from trained q gradients."""
 36    phi = torch.zeros(len(s), device=s.device)
 37    for j in range(K):
 38        r = (j + .5) / K
 39        z = torch.cat([s, u[:, 0:1]*r, torch.zeros_like(u[:, 1:2])], 1).requires_grad_(True)
 40        q = model(z[:, :3], z[:, 3:]); g = torch.autograd.grad(q[:, 0].sum(), z, create_graph=create_graph, retain_graph=True)[0][:, 3]
 41        phi = phi + u[:, 0] * g / K
 42        z = torch.cat([s, u[:, 0:1], u[:, 1:2]*r], 1).requires_grad_(True)
 43        q = model(z[:, :3], z[:, 3:]); g = torch.autograd.grad(q[:, 1].sum(), z, create_graph=create_graph, retain_graph=True)[0][:, 4]
 44        phi = phi + u[:, 1] * g / K
 45    return phi
 46
 47def train(cfg, seed, intervention):
 48    seed_all(seed)
 49    device = 'cuda' if torch.cuda.is_available() else 'cpu'
 50    try:
 51        s,a,y,_ = data(seed); net=QNet().to(device)
 52        s=torch.tensor(s,device=device); a=torch.tensor(a,device=device); y=torch.tensor(y,device=device)
 53        opt=torch.optim.Adam(net.parameters(),lr=cfg['lr'],weight_decay=cfg['wd'])
 54        for _ in range(cfg['epochs']):
 55            q=net(s,a); loss=F.mse_loss(q,y)
 56            if intervention:
 57                phi=path_energy(net,s,a,K=cfg['K'],create_graph=True)
 58                loss=loss+cfg['lam']*F.mse_loss(phi,q.sum(1))
 59            opt.zero_grad(); loss.backward(); opt.step()
 60        st,_,_,target=data(seed+5000,400); st=torch.tensor(st,device=device)
 61        grid=torch.cartesian_prod(A,A).to(device); batch=128; acts=[]
 62        net.eval()
 63        for lo in range(0,len(st),batch):
 64            ss=st[lo:lo+batch]; ns=len(ss)
 65            uu=grid.unsqueeze(0).expand(ns,-1,-1).reshape(-1,2)
 66            sx=ss.unsqueeze(1).expand(-1,len(grid),-1).reshape(-1,3)
 67            with torch.no_grad(): q=net(sx,uu).reshape(ns,len(grid),2)
 68            if not intervention:
 69                # standard independent greedy response from each learned q_i
 70                i1=q[:,:,0].argmax(1); i2=q[:,:,1].argmax(1)
 71                acts.append(torch.stack([grid[i1,0],grid[i2,1]],1))
 72            else:
 73                # one coherent tempered joint Gibbs distribution
 74                ph=path_energy(net,sx,uu,K=cfg['K']).reshape(ns,len(grid))
 75                p=torch.softmax(ph/cfg['temp'],1)
 76                acts.append(p @ grid)
 77        pred=torch.cat(acts).cpu().numpy()
 78        return float(np.mean((pred-target)**2)), net.cpu()
 79    except Exception as e:
 80        print('fallback/error:',repr(e))
 81        return float('nan'), None
 82
 83def base_fn(cfg): return lambda s: train(cfg,s,False)[0]
 84
 85def signature():
 86    val,net=train({'lr':.003,'wd':0.,'epochs':10,'K':5,'lam':0.,'temp':.25},0,True)
 87    if net is None: return {'confirmed':False,'error':'training failed'}
 88    s,_,_,_=data(5000,64); s=torch.tensor(s); u=torch.rand(64,2)*2-1
 89    with torch.no_grad():
 90        # Estimate cross-partial mismatch by finite differences of F_i.
 91        h=1e-3; z=torch.cat([s,u],1)
 92    def F_at(uu):
 93        z=torch.cat([s,uu],1).requires_grad_(True); q=net(z[:,:3],z[:,3:]); out=[]
 94        for i in range(2): out.append(torch.autograd.grad(q[:,i].sum(),z,retain_graph=True)[0][:,3+i])
 95        return torch.stack(out,1)
 96    fp=F_at(u+torch.tensor([[0.,h]])); fm=F_at(u-torch.tensor([[0.,h]])); dF1=(fp[:,0]-fm[:,0])/(2*h)
 97    fp=F_at(u+torch.tensor([[h,0.]])); fm=F_at(u-torch.tensor([[h,0.]])); dF2=(fp[:,1]-fm[:,1])/(2*h)
 98    predicted=float(torch.abs(dF1-dF2).mean().item()/3.)
 99    observed=float(torch.sqrt(torch.mean((path_energy(net,s,u)-path_energy_reverse(net,s,u))**2)).item()) if False else float('nan')
100    # reverse path is explicitly evaluated below without changing trained weights
101    def reverse():
102        out=torch.zeros(64)
103        for j in range(5):
104            r=(j+.5)/5
105            z=torch.cat([s,torch.zeros_like(u[:,0:1]),u[:,1:2]*r],1).requires_grad_(True); q=net(z[:,:3],z[:,3:]); g=torch.autograd.grad(q[:,1].sum(),z,retain_graph=True)[0][:,4]; out+=u[:,1]*g/5
106            z=torch.cat([s,u[:,0:1]*r,u[:,1:2]],1).requires_grad_(True); q=net(z[:,:3],z[:,3:]); g=torch.autograd.grad(q[:,0].sum(),z,retain_graph=True)[0][:,3]; out+=u[:,0]*g/5
107        return out
108    observed=float(torch.sqrt(torch.mean((path_energy(net,s,u)-reverse())**2)).item())
109    return {'predicted_rms':predicted,'observed_rms':observed,'confirmed':bool(abs(observed-predicted)<=max(.05,.5*predicted))}
110
111def main():
112    grid=[{'lr':x,'wd':w,'epochs':10} for x in (.001,.003,.01) for w in (0.,)]
113    base=sweep_baseline(base_fn,grid)
114    configs=[dict(base['best_cfg'],lam=l,K=K,temp=t) for l,K,t in ((.01,3,.15),(.05,5,.25),(.2,7,.4))]
115    results=[evaluate(lambda s,c=c: train(c,s,True)[0],SEEDS) for c in configs]
116    idea=min(results,key=lambda r:r['mean'])
117    report=make_report('joint_action_values','local_q_mlp',base,idea,{'custom_track':{'name':'joint_action_values','file':'joint_action_values.py','domain':'multi_agent_rl'},'mechanism_signature':signature()})
118    Path('bench_report.json').write_text(json.dumps(report,indent=2)); print(json.dumps(report,indent=2))
119if __name__=='__main__': main()