import json, math, random, time from pathlib import Path import numpy as np import torch import torch.nn as nn import torch.nn.functional as F SEED = 3030 random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED) def offsets(m1, m2): # stage t uses multiples of M2; stage r uses multiples of M1 return list(range(0, m1 * m2, m2)), list(range(0, m1 * m2, m1)) def verify_math(): m1, m2 = 3, 4 dt, dr = offsets(m1, m2) virtual = sorted({a+b for a in dt for b in dr}) # A path from i to i-(a+b) exists whenever both intermediate positions exist. L = 32 reachable = set() for i in range(L): for a in dt: j = i-a if j < 0: continue for b in dr: k = j-b if k >= 0: reachable.add(i-k) noncop_dt, noncop_dr = offsets(2, 4) noncop_virtual = sorted({a+b for a in noncop_dt for b in noncop_dr}) return { 'coprime': {'gcd': math.gcd(m1,m2), 'Dt': dt, 'Dr': dr, 'physical_union_count': len(set(dt+dr)), 'physical_formula': m1+m2-1, 'pair_count': len(dt)*len(dr), 'virtual_offsets': virtual, 'virtual_count': len(virtual), 'graph_reachable_offsets': sorted(reachable)}, 'non_coprime_control': {'gcd': 2, 'Dt': noncop_dt, 'Dr': noncop_dr, 'virtual_offsets': noncop_virtual, 'virtual_count': len(noncop_virtual)}, 'claim_holds': (len(set(dt+dr)) == m1+m2-1 and set(virtual) == reachable and len(virtual) > len(noncop_virtual)) } class SparseAttention(nn.Module): def __init__(self, d_model, offsets_, vocab): super().__init__(); self.d = d_model; self.offsets = offsets_ self.q = nn.Linear(d_model,d_model); self.k = nn.Linear(d_model,d_model) self.v = nn.Linear(d_model,d_model); self.rel = nn.Parameter(torch.zeros(len(offsets_))) def forward(self, x): # x[:, i] attends to x[:, i-delta], hence is causal. B,L,D = x.shape; q=self.q(x); k=self.k(x); v=self.v(x) scores=[]; vals=[] neg = torch.finfo(x.dtype).min for p,delta in enumerate(self.offsets): if delta == 0: src=x else: src=torch.cat([torch.zeros(B,delta,D,device=x.device,dtype=x.dtype), x[:,:L-delta]], 1) kk=self.k(src) if delta else k vv=self.v(src) if delta else v scores.append((q*kk).sum(-1)/math.sqrt(D) + self.rel[p]) vals.append(vv) # One scalar score per offset at each query position. s=torch.stack(scores, -1) valid=torch.stack([torch.arange(L,device=x.device)>=d for d in self.offsets],-1) s=s.masked_fill(~valid[None,:,:], neg) w=F.softmax(s,-1) return sum(w[...,p:p+1]*vals[p] for p in range(len(vals))) class CPA(nn.Module): def __init__(self, vocab, d=32, m1=3, m2=4): super().__init__(); dt,dr=offsets(m1,m2); self.dt=dt; self.dr=dr self.emb=nn.Embedding(vocab,d); self.a=SparseAttention(d,dt,vocab); self.b=SparseAttention(d,dr,vocab) self.out=nn.Linear(d,vocab) def forward(self,x): h=self.emb(x); h=h+self.a(h); h=h+self.b(h); return self.out(h) class SingleSparse(nn.Module): def __init__(self,vocab,d=32, offsets_=None): super().__init__(); self.emb=nn.Embedding(vocab,d); self.a=SparseAttention(d,offsets_,vocab); self.out=nn.Linear(d,vocab) def forward(self,x): h=self.emb(x); return self.out(h+self.a(h)) class Dense(nn.Module): def __init__(self,vocab,d=32): super().__init__(); self.emb=nn.Embedding(vocab,d); self.q=nn.Linear(d,d); self.k=nn.Linear(d,d); self.v=nn.Linear(d,d); self.out=nn.Linear(d,vocab) def forward(self,x): h=self.emb(x); q,k,v=self.q(h),self.k(h),self.v(h); L=x.shape[1] z=(q@k.transpose(-1,-2))/math.sqrt(h.shape[-1]); z=z.masked_fill(torch.triu(torch.ones(L,L,device=x.device),1).bool(),-1e9) return self.out(h+(z.softmax(-1)@v)) def batch(vocab=19,L=24,B=64,lag=12,device='cpu'): x=torch.randint(vocab,(B,L),device=device); y=torch.zeros_like(x); y[:,lag:]=x[:,:-lag]; return x,y def train(model, steps, device, vocab=19, L=24, lag=12): model.to(device); opt=torch.optim.AdamW(model.parameters(),lr=3e-3); t=time.perf_counter(); last=0 for _ in range(steps): x,y=batch(vocab,L,64,lag,device); logits=model(x); loss=F.cross_entropy(logits[:,lag:].reshape(-1,vocab),y[:,lag:].reshape(-1)) opt.zero_grad(); loss.backward(); opt.step(); last=float(loss) elapsed=time.perf_counter()-t with torch.no_grad(): x,y=batch(vocab,L,256,lag,device); z=model(x); acc=(z[:,lag:].argmax(-1)==y[:,lag:]).float().mean().item() return {'final_loss':last,'accuracy':acc,'seconds':elapsed} def main(): math_check=verify_math(); device='cuda' if torch.cuda.is_available() else 'cpu' try: torch.zeros(1,device=device) except Exception: device='cpu' vocab,L,lag=19,24,17; dt,dr=offsets(3,4) # Six physical offset families is the CPA union; the control has seven offsets. # Both sparse controls therefore have comparable one-stage edge counts, while CPA has two stages. configs={'cpa': lambda: CPA(vocab,32,3,4), 'single_sparse': lambda: SingleSparse(vocab,32,list(range(7))), 'dense': lambda: Dense(vocab,32)} results={} for name,make in configs.items(): random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED) try: results[name]=train(make(),180,device,vocab,L,lag) except Exception as e: if device!='cpu': device='cpu'; random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED) results[name]=train(make(),180,device,vocab,L,lag) else: raise # attention score edges per token (excluding boundary clipping), and dense equivalent. results['cost_summary']={'cpa_edges_per_token_interior':len(dt)+len(dr), 'single_sparse_edges_per_token':7,'dense_edges_per_token_at_L':L/2, 'cpa_virtual_count':len(math_check['coprime']['virtual_offsets']), 'cpa_fraction_virtual_offsets_gt_8':sum(v>8 for v in math_check['coprime']['virtual_offsets'])/len(math_check['coprime']['virtual_offsets'])} out={'seed':SEED,'device':device,'math':math_check,'training':results} Path('results.json').write_text(json.dumps(out,indent=2)) print(json.dumps(out,indent=2)) if __name__=='__main__': main()