Backward-Reachability Distance Head / experiment.py
Mechanism failed
1import json, math, random
2from pathlib import Path
3import numpy as np
4import torch
5import torch.nn as nn
6import torch.nn.functional as F
7
8SEED = 17
9random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
10torch.set_num_threads(4)
11try:
12 device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
13 if device.type == 'cuda':
14 torch.zeros(1, device=device)
15except Exception:
16 device = torch.device('cpu')
17
18# Finite monitor: two modes, each with a distinct terminal winning state.
19# The environment state is fully observable, while q is retained only for labels.
20N = 9
21MODES = 2
22A = 3 # left, stay, right
23START = 4
24GOALS = [0, 8]
25D = 20
26
27# Backward reachability distances on each modal transition graph.
28def backward_distances():
29 ds = np.full((MODES, N), np.inf, dtype=np.float32)
30 for mode, goal in enumerate(GOALS):
31 ds[mode, goal] = 0
32 changed = True
33 while changed:
34 changed = False
35 for q in range(N):
36 nxt = set()
37 for action in range(A):
38 if action == 0: z = max(0, q-1)
39 elif action == 1: z = q
40 else: z = min(N-1, q+1)
41 nxt.add(z)
42 # Existential backward reachability: q enters when an enabled edge
43 # reaches the prior set.
44 if q != goal and any(np.isfinite(ds[mode, z]) for z in nxt):
45 val = 1 + min(ds[mode, z] for z in nxt if np.isfinite(ds[mode, z]))
46 if not np.isfinite(ds[mode, q]) or val < ds[mode, q]:
47 ds[mode, q] = val; changed = True
48 return ds
49
50DIST = backward_distances()
51
52class ChainEnv:
53 def __init__(self, mode, max_steps=20):
54 self.mode = mode; self.max_steps = max_steps; self.reset()
55 def reset(self):
56 self.q = START; self.t = 0; return self.obs()
57 def obs(self):
58 x = np.zeros(N, dtype=np.float32); x[self.q] = 1.0
59 return x
60 def step(self, action):
61 old = self.q
62 if action == 0: self.q = max(0, self.q-1)
63 elif action == 2: self.q = min(N-1, self.q+1)
64 self.t += 1
65 done = self.q == GOALS[self.mode] or self.t >= self.max_steps
66 reward = 1.0 if self.q == GOALS[self.mode] else 0.0
67 return self.obs(), reward, done, old, self.q
68
69def check_math():
70 # Exact labels should be shortest step counts and every greedy finite-distance
71 # transition should decrement distance by one (or hit the target).
72 assertions = []
73 assertions.append(bool(np.all(DIST[0] == np.arange(N))))
74 assertions.append(bool(np.all(DIST[1] == np.arange(N)[::-1])))
75 violations = 0
76 for mode in range(MODES):
77 for q in range(N):
78 if DIST[mode,q] > 0:
79 vals=[]
80 for a in range(A):
81 z=max(0,q-1) if a==0 else (q if a==1 else min(N-1,q+1))
82 if DIST[mode,z] < DIST[mode,q]: vals.append(DIST[mode,z])
83 if not vals or min(vals) != DIST[mode,q]-1: violations += 1
84 # Formula's hinge is exactly zero for ideal one-step-decreasing predictions.
85 pred = torch.tensor(DIST, dtype=torch.float32)
86 prog = []
87 for mode in range(MODES):
88 for q in range(N):
89 if DIST[mode,q] > 0:
90 z = max(0,q-1) if mode == 0 else min(N-1,q+1)
91 prog.append(max(0.0, float(pred[mode,z]-pred[mode,q]+1)))
92 return {'distance_table': DIST.tolist(), 'shortest_path_assertions': assertions,
93 'greedy_progress_violations': violations, 'ideal_progress_hinge_sum': sum(prog),
94 'math_pass': all(assertions) and violations == 0 and sum(prog) == 0}
95
96class Net(nn.Module):
97 def __init__(self, aux):
98 super().__init__(); self.aux=aux
99 self.body=nn.Sequential(nn.Linear(N,32),nn.Tanh())
100 self.pi=nn.Linear(32,A); self.v=nn.Linear(32,1)
101 self.h=nn.Linear(32,MODES) if aux else None
102 def forward(self,x):
103 z=self.body(x); return self.pi(z), self.v(z).squeeze(-1), self.h(z) if self.h else None
104
105def run(aux, seed, updates=80):
106 torch.manual_seed(seed); np.random.seed(seed); random.seed(seed)
107 envs=[ChainEnv(0),ChainEnv(1)]
108 net=Net(aux).to(device); opt=torch.optim.Adam(net.parameters(),lr=3e-3)
109 successes=[]; losses=[]; aux_err=[]
110 for upd in range(updates):
111 obs=[]; acts=[]; rews=[]; vals=[]; dones=[]; modes=[]; qs=[]; nextobs=[]
112 for e in envs:
113 x=e.reset(); mode=e.mode
114 for _ in range(20):
115 xt=torch.tensor(x,dtype=torch.float32,device=device).unsqueeze(0)
116 logits,v,_=net(xt); dist=torch.distributions.Categorical(logits=logits)
117 a=int(dist.sample().item()); y,r,done,q0,q1=e.step(a)
118 obs.append(x); acts.append(a); rews.append(r); vals.append(float(v.item())); dones.append(done); modes.append(mode); qs.append(q0); nextobs.append(y); x=y
119 if done: break
120 X=torch.tensor(np.asarray(obs),dtype=torch.float32,device=device)
121 Y=torch.tensor(np.asarray(nextobs),dtype=torch.float32,device=device)
122 act=torch.tensor(acts,dtype=torch.long,device=device); R=torch.tensor(rews,dtype=torch.float32,device=device)
123 Vold=torch.tensor(vals,dtype=torch.float32,device=device)
124 logits,V,H=net(X); logp=F.log_softmax(logits,dim=-1).gather(1,act[:,None]).squeeze(1)
125 # Monte-Carlo return within collected short episodes, reset at terminal transitions.
126 ret=[]; g=0.0
127 for r,d in zip(rews[::-1],dones[::-1]):
128 if d: g=0.0
129 g=float(r)+0.97*g; ret.append(g)
130 returns=torch.tensor(ret[::-1],dtype=torch.float32,device=device)
131 adv=(returns-V.detach()); adv=(adv-adv.mean())/(adv.std()+1e-6)
132 rl=-(logp*adv).mean()+0.5*F.mse_loss(V,returns)-0.01*(F.softmax(logits,-1)*F.log_softmax(logits,-1)).sum(-1).mean()
133 total=rl; dl= torch.tensor(0.,device=device); pl=torch.tensor(0.,device=device)
134 if aux:
135 Hn=net(Y)[2]
136 labels=torch.tensor(np.clip([DIST[:, q] for q in qs], 0, D), dtype=torch.float32, device=device)
137 dl=F.huber_loss(H,labels)
138 # Require progress only for finite positive distances; all states here are finite.
139 mask=labels>0
140 pl=F.relu(Hn-H+1.0)[mask].mean() if mask.any() else dl*0
141 total=rl+0.10*dl+0.05*pl
142 aux_err.append(float((H.detach()-labels).abs().mean().item()))
143 opt.zero_grad(); total.backward(); torch.nn.utils.clip_grad_norm_(net.parameters(),1.0); opt.step()
144 losses.append(float(rl.item()))
145 successes.append(float(sum(rews)))
146 # Evaluation with greedy policy from both modes.
147 eval_success=[]; eval_steps=[]
148 with torch.no_grad():
149 for mode in range(MODES):
150 for _ in range(20):
151 e=ChainEnv(mode); x=e.reset(); got=False
152 for t in range(20):
153 lg,_,_=net(torch.tensor(x,dtype=torch.float32,device=device).unsqueeze(0)); a=int(lg.argmax(-1).item()); x,r,d,_,_=e.step(a)
154 if r: got=True; eval_steps.append(t+1); break
155 eval_success.append(got)
156 return {'train_reward_last40':float(np.mean(successes[-40:])), 'eval_success':float(np.mean(eval_success)),
157 'eval_steps_success_only':float(np.mean(eval_steps)) if eval_steps else None,
158 'critic_loss_last40':float(np.mean(losses[-40:])),
159 'distance_abs_error_last40':float(np.mean(aux_err[-40:])) if aux_err else None}
160
161def main():
162 check=check_math()
163 # Three fixed seeds provide a small reproducibility check.
164 base=[run(False,s) for s in [101,202]]
165 idea=[run(True,s) for s in [101,202]]
166 def avg(rows,key):
167 vals=[r[key] for r in rows if r[key] is not None]; return float(np.mean(vals))
168 result={'device':str(device),'math_check':check,
169 'baseline_runs':base,'idea_runs':idea,
170 'baseline':{k:avg(base,k) for k in ['train_reward_last40','eval_success','eval_steps_success_only','critic_loss_last40']},
171 'idea':{k:avg(idea,k) for k in ['train_reward_last40','eval_success','eval_steps_success_only','critic_loss_last40','distance_abs_error_last40']}}
172 Path('results.json').write_text(json.dumps(result,indent=2))
173 print(json.dumps(result,indent=2))
174if __name__=='__main__': main()