import json, math, random from pathlib import Path import numpy as np import torch import torch.nn as nn import torch.nn.functional as F SEED = 17 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.zeros(1, device=device) except Exception: device = torch.device('cpu') # Finite monitor: two modes, each with a distinct terminal winning state. # The environment state is fully observable, while q is retained only for labels. N = 9 MODES = 2 A = 3 # left, stay, right START = 4 GOALS = [0, 8] D = 20 # Backward reachability distances on each modal transition graph. def backward_distances(): ds = np.full((MODES, N), np.inf, dtype=np.float32) for mode, goal in enumerate(GOALS): ds[mode, goal] = 0 changed = True while changed: changed = False for q in range(N): nxt = set() for action in range(A): if action == 0: z = max(0, q-1) elif action == 1: z = q else: z = min(N-1, q+1) nxt.add(z) # Existential backward reachability: q enters when an enabled edge # reaches the prior set. if q != goal and any(np.isfinite(ds[mode, z]) for z in nxt): val = 1 + min(ds[mode, z] for z in nxt if np.isfinite(ds[mode, z])) if not np.isfinite(ds[mode, q]) or val < ds[mode, q]: ds[mode, q] = val; changed = True return ds DIST = backward_distances() class ChainEnv: def __init__(self, mode, max_steps=20): self.mode = mode; self.max_steps = max_steps; self.reset() def reset(self): self.q = START; self.t = 0; return self.obs() def obs(self): x = np.zeros(N, dtype=np.float32); x[self.q] = 1.0 return x def step(self, action): old = self.q if action == 0: self.q = max(0, self.q-1) elif action == 2: self.q = min(N-1, self.q+1) self.t += 1 done = self.q == GOALS[self.mode] or self.t >= self.max_steps reward = 1.0 if self.q == GOALS[self.mode] else 0.0 return self.obs(), reward, done, old, self.q def check_math(): # Exact labels should be shortest step counts and every greedy finite-distance # transition should decrement distance by one (or hit the target). assertions = [] assertions.append(bool(np.all(DIST[0] == np.arange(N)))) assertions.append(bool(np.all(DIST[1] == np.arange(N)[::-1]))) violations = 0 for mode in range(MODES): for q in range(N): if DIST[mode,q] > 0: vals=[] for a in range(A): z=max(0,q-1) if a==0 else (q if a==1 else min(N-1,q+1)) if DIST[mode,z] < DIST[mode,q]: vals.append(DIST[mode,z]) if not vals or min(vals) != DIST[mode,q]-1: violations += 1 # Formula's hinge is exactly zero for ideal one-step-decreasing predictions. pred = torch.tensor(DIST, dtype=torch.float32) prog = [] for mode in range(MODES): for q in range(N): if DIST[mode,q] > 0: z = max(0,q-1) if mode == 0 else min(N-1,q+1) prog.append(max(0.0, float(pred[mode,z]-pred[mode,q]+1))) return {'distance_table': DIST.tolist(), 'shortest_path_assertions': assertions, 'greedy_progress_violations': violations, 'ideal_progress_hinge_sum': sum(prog), 'math_pass': all(assertions) and violations == 0 and sum(prog) == 0} class Net(nn.Module): def __init__(self, aux): super().__init__(); self.aux=aux self.body=nn.Sequential(nn.Linear(N,32),nn.Tanh()) self.pi=nn.Linear(32,A); self.v=nn.Linear(32,1) self.h=nn.Linear(32,MODES) if aux else None def forward(self,x): z=self.body(x); return self.pi(z), self.v(z).squeeze(-1), self.h(z) if self.h else None def run(aux, seed, updates=80): torch.manual_seed(seed); np.random.seed(seed); random.seed(seed) envs=[ChainEnv(0),ChainEnv(1)] net=Net(aux).to(device); opt=torch.optim.Adam(net.parameters(),lr=3e-3) successes=[]; losses=[]; aux_err=[] for upd in range(updates): obs=[]; acts=[]; rews=[]; vals=[]; dones=[]; modes=[]; qs=[]; nextobs=[] for e in envs: x=e.reset(); mode=e.mode for _ in range(20): xt=torch.tensor(x,dtype=torch.float32,device=device).unsqueeze(0) logits,v,_=net(xt); dist=torch.distributions.Categorical(logits=logits) a=int(dist.sample().item()); y,r,done,q0,q1=e.step(a) 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 if done: break X=torch.tensor(np.asarray(obs),dtype=torch.float32,device=device) Y=torch.tensor(np.asarray(nextobs),dtype=torch.float32,device=device) act=torch.tensor(acts,dtype=torch.long,device=device); R=torch.tensor(rews,dtype=torch.float32,device=device) Vold=torch.tensor(vals,dtype=torch.float32,device=device) logits,V,H=net(X); logp=F.log_softmax(logits,dim=-1).gather(1,act[:,None]).squeeze(1) # Monte-Carlo return within collected short episodes, reset at terminal transitions. ret=[]; g=0.0 for r,d in zip(rews[::-1],dones[::-1]): if d: g=0.0 g=float(r)+0.97*g; ret.append(g) returns=torch.tensor(ret[::-1],dtype=torch.float32,device=device) adv=(returns-V.detach()); adv=(adv-adv.mean())/(adv.std()+1e-6) 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() total=rl; dl= torch.tensor(0.,device=device); pl=torch.tensor(0.,device=device) if aux: Hn=net(Y)[2] labels=torch.tensor(np.clip([DIST[:, q] for q in qs], 0, D), dtype=torch.float32, device=device) dl=F.huber_loss(H,labels) # Require progress only for finite positive distances; all states here are finite. mask=labels>0 pl=F.relu(Hn-H+1.0)[mask].mean() if mask.any() else dl*0 total=rl+0.10*dl+0.05*pl aux_err.append(float((H.detach()-labels).abs().mean().item())) opt.zero_grad(); total.backward(); torch.nn.utils.clip_grad_norm_(net.parameters(),1.0); opt.step() losses.append(float(rl.item())) successes.append(float(sum(rews))) # Evaluation with greedy policy from both modes. eval_success=[]; eval_steps=[] with torch.no_grad(): for mode in range(MODES): for _ in range(20): e=ChainEnv(mode); x=e.reset(); got=False for t in range(20): lg,_,_=net(torch.tensor(x,dtype=torch.float32,device=device).unsqueeze(0)); a=int(lg.argmax(-1).item()); x,r,d,_,_=e.step(a) if r: got=True; eval_steps.append(t+1); break eval_success.append(got) return {'train_reward_last40':float(np.mean(successes[-40:])), 'eval_success':float(np.mean(eval_success)), 'eval_steps_success_only':float(np.mean(eval_steps)) if eval_steps else None, 'critic_loss_last40':float(np.mean(losses[-40:])), 'distance_abs_error_last40':float(np.mean(aux_err[-40:])) if aux_err else None} def main(): check=check_math() # Three fixed seeds provide a small reproducibility check. base=[run(False,s) for s in [101,202]] idea=[run(True,s) for s in [101,202]] def avg(rows,key): vals=[r[key] for r in rows if r[key] is not None]; return float(np.mean(vals)) result={'device':str(device),'math_check':check, 'baseline_runs':base,'idea_runs':idea, 'baseline':{k:avg(base,k) for k in ['train_reward_last40','eval_success','eval_steps_success_only','critic_loss_last40']}, 'idea':{k:avg(idea,k) for k in ['train_reward_last40','eval_success','eval_steps_success_only','critic_loss_last40','distance_abs_error_last40']}} Path('results.json').write_text(json.dumps(result,indent=2)) print(json.dumps(result,indent=2)) if __name__=='__main__': main()