Co-Prime Virtual-Aperture Attention / coprime_attention_experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  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()