Orbit-Consistent Equivariant Distillation / run_experiment.py

Mechanism failed

Raw ⬇ ZIP
 1import json, random
 2from pathlib import Path
 3import numpy as np
 4import torch
 5from torch import nn
 6
 7SEED=2854
 8random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
 9torch.set_num_threads(4)
10try:
11    device=torch.device('cuda' if torch.cuda.is_available() else 'cpu')
12    if device.type=='cuda': torch.cuda.get_device_properties(0)
13except Exception:
14    device=torch.device('cpu')
15
16def rot(x,k):
17    if k==0: return x
18    if k==1: return torch.stack((-x[...,1],x[...,0]),-1)
19    if k==2: return -x
20    return torch.stack((x[...,1],-x[...,0]),-1)
21def transforms(x): return torch.stack([rot(x,k) for k in range(4)],0)
22
23def target(s):
24    r2=(s*s).sum(-1,keepdim=True)
25    return s*(1.0+0.25*r2)
26def value_target(s): return (s*s).sum(-1,keepdim=True)
27
28class Student(nn.Module):
29    def __init__(self):
30        super().__init__(); self.body=nn.Sequential(nn.Linear(2,48),nn.Tanh(),nn.Linear(48,48),nn.Tanh())
31        self.policy=nn.Linear(48,2); self.value=nn.Linear(48,1)
32    def forward(self,x):
33        h=self.body(x); return self.policy(h),self.value(h)
34
35def train(lam_eq, lam_v, steps=900, n=256):
36    s=torch.rand(n,2,device=device)*2-1; at=target(s); vt=value_target(s)
37    model=Student().to(device); opt=torch.optim.Adam(model.parameters(),lr=3e-3)
38    for _ in range(steps):
39        idx=torch.randint(0,n,(64,),device=device); x=s[idx]; y=at[idx]; vy=vt[idx]
40        pred,val=model(x); loss=((pred-y)**2).mean()+.15*((val-vy)**2).mean()
41        xs=transforms(x); pg,vg=model(xs.reshape(-1,2)); pg=pg.reshape(4,-1,2); vg=vg.reshape(4,-1,1)
42        loss=loss+lam_eq*((pg-transforms(pred))**2).mean()+lam_v*((vg-val.unsqueeze(0))**2).mean()
43        opt.zero_grad(); loss.backward(); opt.step()
44    return model
45
46@torch.no_grad()
47def metrics(model,n=4096):
48    s=torch.rand(n,2,device=device)*2-1; xs=transforms(s); p,_=model(s)
49    pg,v=model(xs.reshape(-1,2)); pg=pg.reshape(4,n,2); v=v.reshape(4,n,1)
50    eq=((pg-transforms(p)).norm(dim=-1)/(p.norm(dim=-1).clamp_min(1e-6))).mean().item()
51    val_inv=(v-v[0:1]).abs().mean().item()
52    canon=((p-target(s))**2).mean().item(); transformed=((pg-transforms(target(s))).pow(2).mean()).item()
53    return dict(eq_error=eq,value_invariance=val_inv,canonical_mse=canon,transformed_mse=transformed)
54
55def main():
56    x=torch.tensor([[.2,-.7],[.8,.1]],device=device); tx=transforms(x)
57    core_action=(transforms(target(x))-target(tx)).abs().max().item()
58    core_value=(value_target(tx)-value_target(x).unsqueeze(0)).abs().max().item()
59    eq_lams=[0.0,0.03,0.1,0.3,1.0]; eq_results=[]
60    for l in eq_lams:
61        torch.manual_seed(SEED+int(l*1000)); eq_results.append({'lambda_eq':l,**metrics(train(l,0.0))})
62    v_lams=[0.0,0.03,0.1,0.3]; v_results=[]
63    for l in v_lams:
64        torch.manual_seed(SEED+100+int(l*1000)); v_results.append({'lambda_v':l,**metrics(train(0.3,l))})
65    baseline=eq_results[0]; idea=eq_results[-1]
66    eq_noninc=all(eq_results[i+1]['eq_error'] <= eq_results[i]['eq_error']*1.12 for i in range(4))
67    fivefold=baseline['eq_error']/max(idea['eq_error'],1e-12)
68    v_noninc=all(v_results[i+1]['value_invariance'] <= v_results[i]['value_invariance']*1.15 for i in range(3))
69    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}
70    Path('results.json').write_text(json.dumps(out,indent=2)); print(json.dumps(out,indent=2))
71if __name__=='__main__': main()