Co-Prime Virtual-Aperture Attention / coprime_attention_experiment.py
Mechanism confirmed, baseline not beaten
1import json, math, random, time
2from pathlib import Path
3import numpy as np
4import torch
5import torch.nn as nn
6import torch.nn.functional as F
7
8SEED = 3030
9random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
10
11def offsets(m1, m2):
12 # stage t uses multiples of M2; stage r uses multiples of M1
13 return list(range(0, m1 * m2, m2)), list(range(0, m1 * m2, m1))
14
15def verify_math():
16 m1, m2 = 3, 4
17 dt, dr = offsets(m1, m2)
18 virtual = sorted({a+b for a in dt for b in dr})
19 # A path from i to i-(a+b) exists whenever both intermediate positions exist.
20 L = 32
21 reachable = set()
22 for i in range(L):
23 for a in dt:
24 j = i-a
25 if j < 0: continue
26 for b in dr:
27 k = j-b
28 if k >= 0: reachable.add(i-k)
29 noncop_dt, noncop_dr = offsets(2, 4)
30 noncop_virtual = sorted({a+b for a in noncop_dt for b in noncop_dr})
31 return {
32 'coprime': {'gcd': math.gcd(m1,m2), 'Dt': dt, 'Dr': dr,
33 'physical_union_count': len(set(dt+dr)),
34 'physical_formula': m1+m2-1,
35 'pair_count': len(dt)*len(dr),
36 'virtual_offsets': virtual, 'virtual_count': len(virtual),
37 'graph_reachable_offsets': sorted(reachable)},
38 'non_coprime_control': {'gcd': 2, 'Dt': noncop_dt, 'Dr': noncop_dr,
39 'virtual_offsets': noncop_virtual,
40 'virtual_count': len(noncop_virtual)},
41 'claim_holds': (len(set(dt+dr)) == m1+m2-1 and set(virtual) == reachable
42 and len(virtual) > len(noncop_virtual))
43 }
44
45class SparseAttention(nn.Module):
46 def __init__(self, d_model, offsets_, vocab):
47 super().__init__(); self.d = d_model; self.offsets = offsets_
48 self.q = nn.Linear(d_model,d_model); self.k = nn.Linear(d_model,d_model)
49 self.v = nn.Linear(d_model,d_model); self.rel = nn.Parameter(torch.zeros(len(offsets_)))
50 def forward(self, x):
51 # x[:, i] attends to x[:, i-delta], hence is causal.
52 B,L,D = x.shape; q=self.q(x); k=self.k(x); v=self.v(x)
53 scores=[]; vals=[]
54 neg = torch.finfo(x.dtype).min
55 for p,delta in enumerate(self.offsets):
56 if delta == 0: src=x
57 else:
58 src=torch.cat([torch.zeros(B,delta,D,device=x.device,dtype=x.dtype), x[:,:L-delta]], 1)
59 kk=self.k(src) if delta else k
60 vv=self.v(src) if delta else v
61 scores.append((q*kk).sum(-1)/math.sqrt(D) + self.rel[p])
62 vals.append(vv)
63 # One scalar score per offset at each query position.
64 s=torch.stack(scores, -1)
65 valid=torch.stack([torch.arange(L,device=x.device)>=d for d in self.offsets],-1)
66 s=s.masked_fill(~valid[None,:,:], neg)
67 w=F.softmax(s,-1)
68 return sum(w[...,p:p+1]*vals[p] for p in range(len(vals)))
69
70class CPA(nn.Module):
71 def __init__(self, vocab, d=32, m1=3, m2=4):
72 super().__init__(); dt,dr=offsets(m1,m2); self.dt=dt; self.dr=dr
73 self.emb=nn.Embedding(vocab,d); self.a=SparseAttention(d,dt,vocab); self.b=SparseAttention(d,dr,vocab)
74 self.out=nn.Linear(d,vocab)
75 def forward(self,x):
76 h=self.emb(x); h=h+self.a(h); h=h+self.b(h); return self.out(h)
77
78class SingleSparse(nn.Module):
79 def __init__(self,vocab,d=32, offsets_=None):
80 super().__init__(); self.emb=nn.Embedding(vocab,d); self.a=SparseAttention(d,offsets_,vocab); self.out=nn.Linear(d,vocab)
81 def forward(self,x):
82 h=self.emb(x); return self.out(h+self.a(h))
83
84class Dense(nn.Module):
85 def __init__(self,vocab,d=32):
86 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)
87 def forward(self,x):
88 h=self.emb(x); q,k,v=self.q(h),self.k(h),self.v(h); L=x.shape[1]
89 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)
90 return self.out(h+(z.softmax(-1)@v))
91
92def batch(vocab=19,L=24,B=64,lag=12,device='cpu'):
93 x=torch.randint(vocab,(B,L),device=device); y=torch.zeros_like(x); y[:,lag:]=x[:,:-lag]; return x,y
94
95def train(model, steps, device, vocab=19, L=24, lag=12):
96 model.to(device); opt=torch.optim.AdamW(model.parameters(),lr=3e-3); t=time.perf_counter(); last=0
97 for _ in range(steps):
98 x,y=batch(vocab,L,64,lag,device); logits=model(x); loss=F.cross_entropy(logits[:,lag:].reshape(-1,vocab),y[:,lag:].reshape(-1))
99 opt.zero_grad(); loss.backward(); opt.step(); last=float(loss)
100 elapsed=time.perf_counter()-t
101 with torch.no_grad():
102 x,y=batch(vocab,L,256,lag,device); z=model(x); acc=(z[:,lag:].argmax(-1)==y[:,lag:]).float().mean().item()
103 return {'final_loss':last,'accuracy':acc,'seconds':elapsed}
104
105def main():
106 math_check=verify_math(); device='cuda' if torch.cuda.is_available() else 'cpu'
107 try:
108 torch.zeros(1,device=device)
109 except Exception: device='cpu'
110 vocab,L,lag=19,24,17; dt,dr=offsets(3,4)
111 # Six physical offset families is the CPA union; the control has seven offsets.
112 # Both sparse controls therefore have comparable one-stage edge counts, while CPA has two stages.
113 configs={'cpa': lambda: CPA(vocab,32,3,4),
114 'single_sparse': lambda: SingleSparse(vocab,32,list(range(7))),
115 'dense': lambda: Dense(vocab,32)}
116 results={}
117 for name,make in configs.items():
118 random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
119 try: results[name]=train(make(),180,device,vocab,L,lag)
120 except Exception as e:
121 if device!='cpu':
122 device='cpu'; random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
123 results[name]=train(make(),180,device,vocab,L,lag)
124 else: raise
125 # attention score edges per token (excluding boundary clipping), and dense equivalent.
126 results['cost_summary']={'cpa_edges_per_token_interior':len(dt)+len(dr),
127 'single_sparse_edges_per_token':7,'dense_edges_per_token_at_L':L/2,
128 'cpa_virtual_count':len(math_check['coprime']['virtual_offsets']),
129 'cpa_fraction_virtual_offsets_gt_8':sum(v>8 for v in math_check['coprime']['virtual_offsets'])/len(math_check['coprime']['virtual_offsets'])}
130 out={'seed':SEED,'device':device,'math':math_check,'training':results}
131 Path('results.json').write_text(json.dumps(out,indent=2))
132 print(json.dumps(out,indent=2))
133if __name__=='__main__': main()