Inverse-Gain Structured Privileged Distillation / experiment.py

✓✓ Beats tuned baseline

Raw ⬇ ZIP
 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()