Inverse-Gain Structured Privileged Distillation / experiment.py
Beats tuned baseline
1import json, random, math
2from pathlib import Path
3import numpy as np
4import torch
5from torch import nn
6
7SEED=2086
8np.random.seed(SEED); random.seed(SEED); torch.manual_seed(SEED)
9torch.set_num_threads(4)
10
11
12def math_checks():
13 rng=np.random.default_rng(SEED); n=200000; Tc=.05; k1=1.7
14 a=rng.choice([.25,.5,1.,2.],n); d=rng.normal(0,.8,n)
15 xp=rng.normal(0,2,n); v=rng.normal(0,2,n); r=rng.normal(0,2,n)
16 x=xp+Tc*(a*v+d); delta=(x-xp)/Tc; e=r-x; z=k1*e-delta; q=1/a
17 expert=(k1*e-d)/a; wrapped=v+q*z
18 qerr=rng.normal(0,.08,n); wrapped_bad=v+(q+qerr)*z
19 sigmas=np.array([.01,.03,.06,.12,.24]); maes=[]
20 for s in sigmas: maes.append(np.mean(np.abs(rng.normal(0,s,n)*z)))
21 slope=np.polyfit(sigmas,maes,1)[0]
22 # Explicit parameter sweeps: cancellation must not depend on disturbance,
23 # and action error must scale with inverse-gain error and |z|.
24 disturbance_sweep=[]
25 for amp in [0., .1, .5, 1., 2., 4.]:
26 dd=rng.normal(0,amp,n)
27 xx=xp+Tc*(a*v+dd); dz=(xx-xp)/Tc; zz=k1*(r-xx)-dz
28 ee=(k1*(r-xx)-dd)/a
29 ww=v+(1/a)*zz
30 disturbance_sweep.append({'disturbance_std':amp,'max_cancellation_error':float(np.max(abs(ee-ww)))})
31 q_scales=np.array([.01,.03,.06,.12,.24])
32 q_sweep=[]
33 for s in q_scales:
34 qe=rng.normal(0,s,n); ae=qe*z
35 q_sweep.append({'q_error_std':float(s),'action_mae':float(np.mean(abs(ae))),
36 'predicted_action_mae':float(s*np.mean(abs(z)))})
37 return dict(identity_max=float(np.max(abs(expert-wrapped))),factor_identity_max=float(np.max(abs((wrapped_bad-expert)-qerr*z))),sigmas=sigmas.tolist(),maes=np.asarray(maes).tolist(),slope=float(slope),predicted_slope=float(np.mean(abs(z))*math.sqrt(2/math.pi)),disturbance_sweep=disturbance_sweep,q_error_sweep=q_sweep)
38
39
40def make_data(episodes=180, T=80, train=True):
41 rng=np.random.default_rng(SEED+(0 if train else 99)); X=[]; Y=[]
42 for _ in range(episodes):
43 a_seq=rng.choice([.25,.5,1.,2.],T); d_seq=np.repeat(rng.normal(0,.35,T//8+1),8)[:T]
44 # hold-out uses interpolated gains/disturbance patterns but unseen combinations
45 if not train: a_seq=np.array([.35, .75, 1.5, 1.9])[rng.integers(0,4,T)]
46 x1=0.; v=0.; prevx=0.; prev_u=0.; hist=[]
47 for k in range(T):
48 r=1.2*math.sin(.035*k)+.4*math.sin(.11*k+(_%3)); a=float(a_seq[k]); d=float(d_seq[k])
49 x1n=x1+.05*(a*v+d); delta=(x1n-x1)/.05; e=r-x1n; z=1.5*e-delta
50 q=1/a; target=v+q*z
51 feat=[x1n,v,r,prev_u,delta,e]
52 X.append(feat); Y.append(target)
53 # teacher-forced state transition using expert reference
54 u=2.2*(target-v); v=v+.05*u; prev_u=u; prevx=x1; x1=x1n
55 return torch.tensor(X,dtype=torch.float32),torch.tensor(Y,dtype=torch.float32).unsqueeze(1)
56
57class MLP(nn.Module):
58 def __init__(self):
59 super().__init__(); self.net=nn.Sequential(nn.Linear(6,48),nn.Tanh(),nn.Linear(48,48),nn.Tanh(),nn.Linear(48,1))
60 def forward(self,x): return self.net(x)
61class GRU(nn.Module):
62 def __init__(self, structured=False):
63 super().__init__(); self.structured=structured; self.gru=nn.GRU(6,32,batch_first=True); self.head=nn.Linear(32,1)
64 def forward(self,x):
65 h=self.gru(x)[0][:,-1]; out=self.head(h)
66 if self.structured:
67 # q interval is [0.5,4], and final feature indices are delta,e; wrapper target
68 q=.5+3.5*torch.sigmoid(out); return x[:,-1,1:2]+q*(1.5*x[:,-1,5:6]-x[:,-1,4:5])
69 return out
70
71def train_eval(X,Y,XT,YT,kind):
72 if kind=='mlp': model=MLP(); xx=X
73 else: model=GRU(kind=='structured'); xx=X.reshape(-1,1,6)
74 opt=torch.optim.Adam(model.parameters(),lr=2e-3); lossfn=nn.SmoothL1Loss()
75 for ep in range(35):
76 p=torch.randperm(len(Y))
77 for i in range(0,len(Y),128):
78 j=p[i:i+128]; pred=model(xx[j]); loss=lossfn(pred,Y[j]); opt.zero_grad(); loss.backward(); opt.step()
79 with torch.no_grad():
80 pred=model(XT if kind=='mlp' else XT.reshape(-1,1,6)); err=(pred-YT).numpy().ravel()
81 return float(np.sqrt(np.mean(err**2))),float(np.mean(np.abs(err)))
82
83def main():
84 checks=math_checks(); X,Y=make_data(); XT,YT=make_data(70,80,False)
85 results={}
86 for k in ['mlp','gru','structured']: results[k]=train_eval(X,Y,XT,YT,k)
87 out={'math_checks':checks,'action_rmse_mae':results,'notes':'Exact identities use the sampled plant equation; learning comparison is teacher-forced offline action prediction.'}
88 Path('results.json').write_text(json.dumps(out,indent=2))
89 print(json.dumps(out,indent=2))
90if __name__=='__main__': main()