import json, math, time from pathlib import Path import numpy as np import torch from torch import nn SEED = 640 np.random.seed(SEED) torch.manual_seed(SEED) torch.set_num_threads(4) device = "cuda" if torch.cuda.is_available() else "cpu" try: if device == "cuda": torch.cuda.set_device(0) torch.zeros(1, device=device) except Exception: device = "cpu" BETA, TAU, KTRAIN, KEVAL = 0.92, 1.4, 8, 256 def logmeanexp(x, dim=-1): return torch.logsumexp(x, dim=dim) - math.log(x.shape[dim]) def transition(s, a, noise): return torch.tanh(0.72*s + 0.35*a + noise) def reward(s, a): return 1.0 - 0.55*s.square() - 0.12*a.square() def mlp(out=1): return nn.Sequential(nn.Linear(2, 48), nn.Tanh(), nn.Linear(48, 48), nn.Tanh(), nn.Linear(48, out)).to(device) def sample_batch(n, k): s = (torch.rand(n, 1, device=device)*2-1) a = (torch.rand(n, 1, device=device)*2-1) noise = torch.randn(n, k, 1, device=device)*0.38 sn = transition(s[:,None,:], a[:,None,:], noise) return s, a, reward(s,a), sn @torch.no_grad() def high_k_target(vtarget, s, a, k=KEVAL): noise = torch.randn(s.shape[0], k, 1, device=device)*0.38 sn = transition(s[:,None,:], a[:,None,:], noise) inp = torch.cat([sn.reshape(-1,1), a[:,None,:].expand(-1,k,-1).reshape(-1,1)], 1) vals = vtarget(inp).reshape(s.shape[0], k) return -logmeanexp(-TAU*vals, 1)/TAU def run(seed, mode, steps=1200): torch.manual_seed(seed) v = mlp(); vt = mlp(); vt.load_state_dict(v.state_dict()) m = mlp() if mode == 'aux' else None opt = torch.optim.Adam(list(v.parameters()) + ([] if m is None else list(m.parameters())), lr=1e-3) records=[]; t0=time.perf_counter() for step in range(steps): s,a,r,sn = sample_batch(64,KTRAIN) if mode == 'direct': with torch.no_grad(): inp=torch.cat([sn.reshape(-1,1), a[:,None,:].expand(-1,KTRAIN,-1).reshape(-1,1)],1) vv=vt(inp).reshape(64,KTRAIN) ce=-logmeanexp(-TAU*vv,1) target=r+BETA*ce pred=v(torch.cat([s,a],1)).squeeze(1) loss=(pred-target).square().mean() else: with torch.no_grad(): inp=torch.cat([sn.reshape(-1,1), a[:,None,:].expand(-1,KTRAIN,-1).reshape(-1,1)],1) vv=vt(inp).reshape(64,KTRAIN) ce=-logmeanexp(-TAU*vv,1) ma=m(torch.cat([s,a],1)).squeeze(1) pred=v(torch.cat([s,a],1)).squeeze(1) bell=(r+BETA*ma-pred).square().mean() cel=(ma-ce).square().mean() loss=bell+0.7*cel opt.zero_grad(); loss.backward(); torch.nn.utils.clip_grad_norm_(list(v.parameters())+([] if m is None else list(m.parameters())), 10); opt.step() with torch.no_grad(): for p,q in zip(vt.parameters(),v.parameters()): p.mul_(0.97).add_(q,alpha=0.03) if step in (steps//2, steps-1): s0,a0,_,_=sample_batch(256,1) true=high_k_target(vt,s0,a0) with torch.no_grad(): if mode=='direct': est=high_k_target(vt,s0,a0,KTRAIN) # direct target estimator distribution auxerr=float('nan') predv=v(torch.cat([s0,a0],1)).squeeze(1) # Bellman residual using high-K CE br=(reward(s0,a0)+BETA*true-predv) else: ma=m(torch.cat([s0,a0],1)).squeeze(1) auxerr=(ma-true).abs().mean().item() br=(reward(s0,a0)+BETA*ma-v(torch.cat([s0,a0],1)).squeeze(1)) records.append({'step':step+1,'ce_abs_error':auxerr,'bellman_rmse':float(br.square().mean().sqrt()),'loss':float(loss)}) # estimator variance at fixed states, over repeated K samples with torch.no_grad(): s,a,_,_=sample_batch(128,1) vals=[] for _ in range(30): vals.append(high_k_target(vt,s,a,KTRAIN).cpu().numpy()) estvar=float(np.mean(np.var(np.stack(vals),axis=0))) true=high_k_target(vt,s,a,KEVAL) vpred=v(torch.cat([s,a],1)).squeeze(1) value_rmse=float((vpred- (reward(s,a)+BETA*true)).square().mean().sqrt()) if m is not None: ce_rmse=float((m(torch.cat([s,a],1)).squeeze(1)-true).square().mean().sqrt()) else: ce_rmse=float('nan') return {'mode':mode,'seed':seed,'device':device,'seconds':time.perf_counter()-t0,'estimator_variance':estvar,'value_target_rmse':value_rmse,'ce_rmse':ce_rmse,'checkpoints':records} def math_check(): # Exact equality for constant continuation values, and convergence of sample CE. x=torch.tensor([[1.,2.,3.]]) got=float((-logmeanexp(-TAU*x,1)/TAU).item()) expected=float(-math.log(np.mean(np.exp(-TAU*x.numpy())))/TAU) rng=np.random.default_rng(SEED) errs=[] for k in (2,8,32,128): z=rng.normal(size=(1000*k,)) mx0=np.max(-TAU*z) exact=-(mx0+math.log(np.mean(np.exp(-TAU*z-mx0))))/TAU e=[] for i in range(1000): q=z[i*k:(i+1)*k] mx=np.max(-TAU*q) e.append((-(mx+math.log(np.mean(np.exp(-TAU*q-mx))))/TAU)-exact) errs.append({'K':k,'abs_error_mean':float(np.mean(np.abs(e))),'bias':float(np.mean(e))}) return {'logmeanexp_absolute_difference':abs(got-expected),'sample_scaling':errs} if __name__=='__main__': out={'math_check':math_check(),'runs':[]} for seed in (11,22,33): for mode in ('direct','aux'): out['runs'].append(run(seed,mode)) Path('results.json').write_text(json.dumps(out,indent=2)) print(json.dumps(out,indent=2))