import json, math, random import numpy as np import torch import torch.nn as nn import torch.nn.functional as F def seed_all(s=7): random.seed(s); np.random.seed(s); torch.manual_seed(s) if torch.cuda.is_available(): torch.cuda.manual_seed_all(s) def get_device(): if torch.cuda.is_available(): try: torch.empty(1, device='cuda') return torch.device('cuda') except Exception: pass return torch.device('cpu') class Profile(nn.Module): def __init__(self, hidden=16): super().__init__() self.net = nn.Sequential(nn.Linear(2, hidden), nn.Tanh(), nn.Linear(hidden, 1)) nn.init.zeros_(self.net[-1].weight) nn.init.constant_(self.net[-1].bias, 0.0) def forward(self, s, q): z = self.net(torch.stack((s, q), -1)).squeeze(-1) return F.softplus(z + 2.0) + 1e-4 class MetricAttention(nn.Module): def __init__(self, d, constrained=True): super().__init__() self.log_h = nn.Parameter(torch.zeros(d)) self.b = nn.Parameter(torch.randn(d) * .08) self.profile = Profile() self.constrained = constrained def metric(self, y, need_derivatives=False): H = F.softplus(self.log_h) + 1e-4 alpha = torch.sqrt((H * y * y).sum(-1) + 1e-8) beta = (self.b * y).sum(-1) s = beta / (alpha + 1e-8) q = (self.b * self.b / H).sum() phi = self.profile(s, q.expand_as(s)) if not need_derivatives: return alpha * phi, None # derivatives of scalar profile with respect to s ss = s.detach().requires_grad_(True) qq = q.detach().expand_as(ss) pp = self.profile(ss, qq) p1 = torch.autograd.grad(pp.sum(), ss, create_graph=True)[0] p2 = torch.autograd.grad(p1.sum(), ss, create_graph=True)[0] g = pp - ss*p1 + (q.detach() - ss*ss)*p2 return alpha * phi, (g, phi, s, q, H) def forward(self, Q, K, V, return_aux=False): y = Q[:, :, None, :] - K[:, None, :, :] cost, aux = self.metric(y, need_derivatives=return_aux) logits = -(cost * cost) / math.sqrt(Q.shape[-1]) a = F.softmax(logits, -1) out = a @ V return out, a, aux class DotAttention(nn.Module): def __init__(self, d): super().__init__(); self.scale = d ** -0.5 def forward(self, Q, K, V, return_aux=False): a = F.softmax(Q @ K.transpose(-1, -2) * self.scale, -1) out = a @ V return out, a, None class TinyModel(nn.Module): def __init__(self, vocab=32, d=24, constrained=True, kind='metric'): super().__init__(); self.emb=nn.Embedding(vocab,d); self.attn=DotAttention(d) if kind=='dot' else MetricAttention(d,constrained) self.ff=nn.Sequential(nn.Linear(d,d*2),nn.GELU(),nn.Linear(d*2,d)); self.norm=nn.LayerNorm(d); self.head=nn.Linear(d,vocab) def forward(self,x, return_aux=False): z=self.emb(x); o,a,aux=self.attn(z,z,z,return_aux); z=self.norm(z+o); z=self.norm(z+self.ff(z)); logits=self.head(z) return logits, a, aux def convexity_check(dev): # Directly verify g for phi=1+c*s: g=1 and inspect numerical Hessian of F^2/2. d=3; c=.25; b=torch.tensor([.35,-.2,.1],device=dev); H=torch.tensor([1.2,.8,1.5],device=dev) vals=[]; mine=[] for _ in range(30): y=torch.randn(d,device=dev); y.requires_grad_() alpha=torch.sqrt((H*y*y).sum()); s=(b*y).sum()/alpha; q=(b*b/H).sum() phi=1+c*s; g=phi-s*c+(q-s*s)*0.0; vals.append(float(g)) f=.5*(alpha*phi)**2 grad=torch.autograd.grad(f,y,create_graph=True)[0] rows=[torch.autograd.grad(grad[i],y,retain_graph=True)[0] for i in range(d)] mine.append(float(torch.linalg.eigvalsh(torch.stack(rows)).min())) # Compare with an intentionally nonconvex quadratic profile. bad=[] for _ in range(20): y=torch.randn(d,device=dev); y.requires_grad_(); alpha=torch.sqrt((H*y*y).sum()); s=(b*y).sum()/alpha f=.5*(alpha*(1-20*s*s))**2; gr=torch.autograd.grad(f,y,create_graph=True)[0] hs=[torch.autograd.grad(gr[i],y,retain_graph=True)[0] for i in range(d)] bad.append(float(torch.linalg.eigvalsh(torch.stack(hs)).min())) return {'linear_profile_g_min':min(vals),'linear_profile_g_max':max(vals),'linear_profile_hessian_min':min(mine),'bad_profile_hessian_min':min(bad)} def run(kind, constrained, dev, seed=7, steps=180): seed_all(seed); vocab=32; L=12; B=64 # deterministic modular next-token task with local context base=torch.arange(L,device=dev)[None,:].repeat(B,1) x=(base + torch.randint(0,vocab,(B,1),device=dev)) % vocab y=(x+1)%vocab model=TinyModel(vocab=vocab,kind=kind,constrained=constrained).to(dev) opt=torch.optim.AdamW(model.parameters(),lr=3e-3) losses=[]; neg=[]; ent=[]; spikes=[] for step in range(steps): # refresh offsets, keeping the same simple task x=(base + torch.randint(0,vocab,(B,1),device=dev)) % vocab; y=(x+1)%vocab logits,a,aux=model(x,return_aux=True) loss=F.cross_entropy(logits.reshape(-1,vocab),y.reshape(-1)) barrier=torch.tensor(0.,device=dev); frac=0. if aux is not None and constrained: g=aux[0]; barrier=F.softplus(.05-g).square().mean(); loss=loss+.03*barrier; frac=float((g<0).float().mean()) opt.zero_grad(); loss.backward(); gn=float(torch.nn.utils.clip_grad_norm_(model.parameters(),10.0)); opt.step() losses.append(float(loss.detach())); neg.append(frac); spikes.append(gn) p=a.detach().clamp_min(1e-8); ent.append(float((-(p*p.log()).sum(-1).mean()).cpu())) with torch.no_grad(): logits,_,aux=model(x,return_aux=False); val=float(F.cross_entropy(logits.reshape(-1,vocab),y.reshape(-1))) if aux is not None: neg_final=float((aux[0]<0).float().mean()) else: neg_final=0. return {'final_loss':val,'mean_last20':float(np.mean(losses[-20:])),'loss_std_last50':float(np.std(losses[-50:])),'max_grad':max(spikes),'final_negative_g_fraction':neg_final,'mean_attention_entropy':float(np.mean(ent[-20:]))} def main(): seed_all(7); dev=get_device(); checks=convexity_check(dev) results={'device':str(dev),'math_check':checks} results['baseline_dot']=run('dot',False,dev,7) results['unconstrained_metric']=run('metric',False,dev,7) results['constrained_metric']=run('metric',True,dev,7) with open('results.json','w') as f: json.dump(results,f,indent=2) print(json.dumps(results,indent=2)) if __name__=='__main__': main()