Coordinate Path-Integral Joint Gibbs Policy / run_bench.py
Mechanism confirmed, baseline not beaten
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()