import json, random from pathlib import Path import numpy as np import torch from torch import nn SEED=2854 random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED) torch.set_num_threads(4) try: device=torch.device('cuda' if torch.cuda.is_available() else 'cpu') if device.type=='cuda': torch.cuda.get_device_properties(0) except Exception: device=torch.device('cpu') def rot(x,k): if k==0: return x if k==1: return torch.stack((-x[...,1],x[...,0]),-1) if k==2: return -x return torch.stack((x[...,1],-x[...,0]),-1) def transforms(x): return torch.stack([rot(x,k) for k in range(4)],0) def target(s): r2=(s*s).sum(-1,keepdim=True) return s*(1.0+0.25*r2) def value_target(s): return (s*s).sum(-1,keepdim=True) class Student(nn.Module): def __init__(self): super().__init__(); self.body=nn.Sequential(nn.Linear(2,48),nn.Tanh(),nn.Linear(48,48),nn.Tanh()) self.policy=nn.Linear(48,2); self.value=nn.Linear(48,1) def forward(self,x): h=self.body(x); return self.policy(h),self.value(h) def train(lam_eq, lam_v, steps=900, n=256): s=torch.rand(n,2,device=device)*2-1; at=target(s); vt=value_target(s) model=Student().to(device); opt=torch.optim.Adam(model.parameters(),lr=3e-3) for _ in range(steps): idx=torch.randint(0,n,(64,),device=device); x=s[idx]; y=at[idx]; vy=vt[idx] pred,val=model(x); loss=((pred-y)**2).mean()+.15*((val-vy)**2).mean() xs=transforms(x); pg,vg=model(xs.reshape(-1,2)); pg=pg.reshape(4,-1,2); vg=vg.reshape(4,-1,1) loss=loss+lam_eq*((pg-transforms(pred))**2).mean()+lam_v*((vg-val.unsqueeze(0))**2).mean() opt.zero_grad(); loss.backward(); opt.step() return model @torch.no_grad() def metrics(model,n=4096): s=torch.rand(n,2,device=device)*2-1; xs=transforms(s); p,_=model(s) pg,v=model(xs.reshape(-1,2)); pg=pg.reshape(4,n,2); v=v.reshape(4,n,1) eq=((pg-transforms(p)).norm(dim=-1)/(p.norm(dim=-1).clamp_min(1e-6))).mean().item() val_inv=(v-v[0:1]).abs().mean().item() canon=((p-target(s))**2).mean().item(); transformed=((pg-transforms(target(s))).pow(2).mean()).item() return dict(eq_error=eq,value_invariance=val_inv,canonical_mse=canon,transformed_mse=transformed) def main(): x=torch.tensor([[.2,-.7],[.8,.1]],device=device); tx=transforms(x) core_action=(transforms(target(x))-target(tx)).abs().max().item() core_value=(value_target(tx)-value_target(x).unsqueeze(0)).abs().max().item() eq_lams=[0.0,0.03,0.1,0.3,1.0]; eq_results=[] for l in eq_lams: torch.manual_seed(SEED+int(l*1000)); eq_results.append({'lambda_eq':l,**metrics(train(l,0.0))}) v_lams=[0.0,0.03,0.1,0.3]; v_results=[] for l in v_lams: torch.manual_seed(SEED+100+int(l*1000)); v_results.append({'lambda_v':l,**metrics(train(0.3,l))}) baseline=eq_results[0]; idea=eq_results[-1] eq_noninc=all(eq_results[i+1]['eq_error'] <= eq_results[i]['eq_error']*1.12 for i in range(4)) fivefold=baseline['eq_error']/max(idea['eq_error'],1e-12) v_noninc=all(v_results[i+1]['value_invariance'] <= v_results[i]['value_invariance']*1.15 for i in range(3)) out={'device':str(device),'seed':SEED,'core_math':{'target_equivariance_max_abs':core_action,'target_value_invariance_max_abs':core_value},'predictions':{'P1':'exact target residuals should be zero','P2':'eq_error decreases with lambda_eq','P3':'value_invariance decreases with lambda_v'},'equivariance_sweep':eq_results,'value_sweep':v_results,'checks':{'p1_pass':core_action<1e-6 and core_value<1e-6,'p2_monotone_pass':eq_noninc,'p2_fold_reduction':fivefold,'p2_expected_at_least_5x':fivefold>=5.0,'p3_monotone_pass':v_noninc},'baseline':baseline,'idea':idea} Path('results.json').write_text(json.dumps(out,indent=2)); print(json.dumps(out,indent=2)) if __name__=='__main__': main()