import json, random, math from pathlib import Path import numpy as np import torch from torch import nn SEED=2086 np.random.seed(SEED); random.seed(SEED); torch.manual_seed(SEED) torch.set_num_threads(4) def math_checks(): rng=np.random.default_rng(SEED); n=200000; Tc=.05; k1=1.7 a=rng.choice([.25,.5,1.,2.],n); d=rng.normal(0,.8,n) xp=rng.normal(0,2,n); v=rng.normal(0,2,n); r=rng.normal(0,2,n) x=xp+Tc*(a*v+d); delta=(x-xp)/Tc; e=r-x; z=k1*e-delta; q=1/a expert=(k1*e-d)/a; wrapped=v+q*z qerr=rng.normal(0,.08,n); wrapped_bad=v+(q+qerr)*z sigmas=np.array([.01,.03,.06,.12,.24]); maes=[] for s in sigmas: maes.append(np.mean(np.abs(rng.normal(0,s,n)*z))) slope=np.polyfit(sigmas,maes,1)[0] # Explicit parameter sweeps: cancellation must not depend on disturbance, # and action error must scale with inverse-gain error and |z|. disturbance_sweep=[] for amp in [0., .1, .5, 1., 2., 4.]: dd=rng.normal(0,amp,n) xx=xp+Tc*(a*v+dd); dz=(xx-xp)/Tc; zz=k1*(r-xx)-dz ee=(k1*(r-xx)-dd)/a ww=v+(1/a)*zz disturbance_sweep.append({'disturbance_std':amp,'max_cancellation_error':float(np.max(abs(ee-ww)))}) q_scales=np.array([.01,.03,.06,.12,.24]) q_sweep=[] for s in q_scales: qe=rng.normal(0,s,n); ae=qe*z q_sweep.append({'q_error_std':float(s),'action_mae':float(np.mean(abs(ae))), 'predicted_action_mae':float(s*np.mean(abs(z)))}) 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) def make_data(episodes=180, T=80, train=True): rng=np.random.default_rng(SEED+(0 if train else 99)); X=[]; Y=[] for _ in range(episodes): a_seq=rng.choice([.25,.5,1.,2.],T); d_seq=np.repeat(rng.normal(0,.35,T//8+1),8)[:T] # hold-out uses interpolated gains/disturbance patterns but unseen combinations if not train: a_seq=np.array([.35, .75, 1.5, 1.9])[rng.integers(0,4,T)] x1=0.; v=0.; prevx=0.; prev_u=0.; hist=[] for k in range(T): r=1.2*math.sin(.035*k)+.4*math.sin(.11*k+(_%3)); a=float(a_seq[k]); d=float(d_seq[k]) x1n=x1+.05*(a*v+d); delta=(x1n-x1)/.05; e=r-x1n; z=1.5*e-delta q=1/a; target=v+q*z feat=[x1n,v,r,prev_u,delta,e] X.append(feat); Y.append(target) # teacher-forced state transition using expert reference u=2.2*(target-v); v=v+.05*u; prev_u=u; prevx=x1; x1=x1n return torch.tensor(X,dtype=torch.float32),torch.tensor(Y,dtype=torch.float32).unsqueeze(1) class MLP(nn.Module): def __init__(self): super().__init__(); self.net=nn.Sequential(nn.Linear(6,48),nn.Tanh(),nn.Linear(48,48),nn.Tanh(),nn.Linear(48,1)) def forward(self,x): return self.net(x) class GRU(nn.Module): def __init__(self, structured=False): super().__init__(); self.structured=structured; self.gru=nn.GRU(6,32,batch_first=True); self.head=nn.Linear(32,1) def forward(self,x): h=self.gru(x)[0][:,-1]; out=self.head(h) if self.structured: # q interval is [0.5,4], and final feature indices are delta,e; wrapper target q=.5+3.5*torch.sigmoid(out); return x[:,-1,1:2]+q*(1.5*x[:,-1,5:6]-x[:,-1,4:5]) return out def train_eval(X,Y,XT,YT,kind): if kind=='mlp': model=MLP(); xx=X else: model=GRU(kind=='structured'); xx=X.reshape(-1,1,6) opt=torch.optim.Adam(model.parameters(),lr=2e-3); lossfn=nn.SmoothL1Loss() for ep in range(35): p=torch.randperm(len(Y)) for i in range(0,len(Y),128): j=p[i:i+128]; pred=model(xx[j]); loss=lossfn(pred,Y[j]); opt.zero_grad(); loss.backward(); opt.step() with torch.no_grad(): pred=model(XT if kind=='mlp' else XT.reshape(-1,1,6)); err=(pred-YT).numpy().ravel() return float(np.sqrt(np.mean(err**2))),float(np.mean(np.abs(err))) def main(): checks=math_checks(); X,Y=make_data(); XT,YT=make_data(70,80,False) results={} for k in ['mlp','gru','structured']: results[k]=train_eval(X,Y,XT,YT,k) out={'math_checks':checks,'action_rmse_mae':results,'notes':'Exact identities use the sampled plant equation; learning comparison is teacher-forced offline action prediction.'} Path('results.json').write_text(json.dumps(out,indent=2)) print(json.dumps(out,indent=2)) if __name__=='__main__': main()