import json, math, random import numpy as np import torch import torch.nn as nn import torch.nn.functional as F SEED=1669 random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED) torch.set_num_threads(4) try: device="cuda" if torch.cuda.is_available() else "cpu" if device=="cuda": torch.zeros(1,device="cuda") except Exception: device="cpu" class Schedule: def __init__(self,n,lam=0.,rho=0.,s=1): self.n,self.lam,self.rho,self.s=n,lam,rho,s assert lam+rho<=1 def sample(self,batch,generator=None): u=torch.rand(batch,generator=generator); branch=torch.zeros(batch,dtype=torch.long) low=(u>=1-self.rho-self.lam)&(u<1-self.rho); full=u>=1-self.rho branch[low]=1; branch[full]=2 vis=torch.full((batch,),self.n-1,dtype=torch.long) if low.any(): vis[low]=torch.randint(0,self.s+1,(int(low.sum()),),generator=generator) vis[full]=0 rank=torch.rand(batch,self.n,generator=generator).argsort(1) visible=rank < vis[:,None] return ~visible,branch,vis def analytic(self): return {"v_min":0 if self.lam+self.rho>0 else self.n-1, "pi_s":self.rho+self.lam if self.s>=1 else self.rho+self.lam/(self.s+1), "full_atom":self.rho} def schedule_check(): rows=[] for lam,rho in [(0,0),(.1,0),(.1,.01),(.25,.01)]: sc=Schedule(8,lam,rho,1); g=torch.Generator().manual_seed(SEED+7) _,b,v=sc.sample(100000,g) rows.append({"lambda":lam,"rho":rho,"pred_pi1":sc.analytic()["pi_s"],"obs_pi1":float((v<=1).float().mean()),"pred_full":rho,"obs_full":float((b==2).float().mean()),"v_min":int(v.min())}) return rows def anchor_check(): p=.8; eta=1.; steps=100; rows=[] for rho in [.01,.03,.1,.2]: b=0. for _ in range(steps): b-=eta*rho*(1/(1+math.exp(-b))-p) q=1/(1+math.exp(-b)); obs=abs(q-p); pred=.3*(1-eta*rho*p*(1-p))**steps rows.append({"rho":rho,"observed_error":obs,"linear_prediction":pred,"ratio":obs/pred}) return rows class TinyMaskedPredictor(nn.Module): def __init__(self,n=8): super().__init__(); self.n=n; self.emb=nn.Embedding(3,12); self.pos=nn.Embedding(n,12) self.net=nn.Sequential(nn.Linear(n*12,40),nn.Tanh(),nn.Linear(40,n*2)) def forward(self,x): z=self.emb(x.long())+self.pos(torch.arange(self.n,device=x.device))[None] return self.net(z.reshape(x.shape[0],-1)).reshape(x.shape[0],self.n,2) def make_data(num,n=8,p=.8): mode=(torch.rand(num)