mlnomad commited on
Commit
d940345
·
verified ·
1 Parent(s): bdacd6b

Upload omniem_model.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. omniem_model.py +55 -0
omniem_model.py ADDED
@@ -0,0 +1,55 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+
2
+ import math, torch, torch.nn as nn, torch.nn.functional as F
3
+ LMAX=6; HEADS=6; HMLP=1536
4
+ class RMSNorm(nn.Module):
5
+ def __init__(s,d): super().__init__(); s.g=nn.Parameter(torch.ones(d))
6
+ def forward(s,x): return x*torch.rsqrt(x.pow(2).mean(-1,keepdim=True)+1e-6)*s.g
7
+ def normz(x): return x/x.norm(dim=-1,keepdim=True).clamp_min(1e-8)
8
+ def Px(v,x): return v-(v*x).sum(-1,keepdim=True)*x
9
+ def Exp(x,v):
10
+ nv=v.norm(dim=-1,keepdim=True).clamp(1e-8, math.pi-1e-3); return torch.cos(nv)*x+torch.sin(nv)*(v/nv)
11
+ def Gamma(x,y,u):
12
+ xy=(x*y).sum(-1,keepdim=True); den=(1.0+xy).clamp_min(1e-6); sm=x+y
13
+ return u-((u*sm).sum(-1,keepdim=True)/den)*sm+2*(u*x).sum(-1,keepdim=True)*y
14
+ class GOATV(nn.Module):
15
+ def __init__(s,d,h):
16
+ super().__init__(); s.h=h; s.dh=d//h
17
+ s.wv=nn.Linear(d,d,bias=False); s.wo=nn.Linear(d,d,bias=False)
18
+ s.b_raw=nn.Parameter(torch.zeros(h,1,1)); s.eps_raw=nn.Parameter(torch.zeros(h,1,1))
19
+ def forward(s,x,am):
20
+ B,T,d=x.shape; H,D=s.h,s.dh
21
+ q=x.view(B,T,H,D).transpose(1,2); k=q; v=s.wv(x).view(B,T,H,D).transpose(1,2)
22
+ b=F.softplus(s.b_raw); eps=F.softplus(s.eps_raw)
23
+ dot=torch.matmul(q,k.transpose(-1,-2)); numer=(dot+b).pow(2)
24
+ qsq=(q*q).sum(-1,keepdim=True); ksq=(k*k).sum(-1,keepdim=True).transpose(-1,-2)
25
+ dist=(qsq+ksq-2*dot).clamp_min(0.0); scores=numer/(dist+eps)
26
+ scores=scores*am.view(B,1,1,T).to(scores.dtype); w=scores/(scores.sum(-1,keepdim=True)+1e-8)
27
+ return s.wo(torch.matmul(w,v).transpose(1,2).reshape(B,T,d))
28
+ class YatMLP(nn.Module):
29
+ def __init__(s,d,Hm):
30
+ super().__init__(); s.W=nn.Parameter(torch.randn(d,Hm)/math.sqrt(d)); s.le=nn.Parameter(torch.tensor(0.0))
31
+ s.A=nn.Parameter(torch.randn(Hm,d)/math.sqrt(Hm)); s.c=nn.Parameter(torch.zeros(d))
32
+ s.b=nn.Parameter(torch.tensor(1.0)); s.ar=nn.Parameter(torch.tensor(0.0))
33
+ def forward(s,x):
34
+ eps=s.le.exp().clamp_min(1e-4); dots=x@s.W
35
+ sqd=((x*x).sum(-1,keepdim=True)-2*dots+(s.W*s.W).sum(0,keepdim=True)).clamp_min(0)
36
+ return (F.softplus(s.ar)*((dots+s.b).pow(2)/(sqd+eps)))@s.A+s.c
37
+ class Student(nn.Module):
38
+ """OmniEM-EN. Warm tokenizer/embeddings from intfloat/multilingual-e5-small; emb = that model's embeddings."""
39
+ def __init__(s,emb,d,h,Hm):
40
+ super().__init__(); s.emb=emb; s.n1=RMSNorm(d); s.attn=GOATV(d,h); s.n2=RMSNorm(d); s.mlp=YatMLP(d,Hm); s.nf=RMSNorm(d)
41
+ s.halt=nn.Sequential(nn.Linear(d,d//2),nn.GELU(),nn.Linear(d//2,1)); s.halt[-1].bias.data.fill_(-2.0)
42
+ s.alpha=nn.Parameter(torch.tensor(0.1)); s.beta_logit=nn.Parameter(torch.tensor(0.0))
43
+ def forward(s,ids,am):
44
+ h=s.emb(input_ids=ids).float(); x=normz(h)
45
+ mom=torch.zeros_like(x); xp=x; beta=torch.sigmoid(s.beta_logit); al=s.alpha
46
+ m=am.unsqueeze(-1).float(); embeds=[]; halts=[]
47
+ for _ in range(LMAX):
48
+ for sub,nrm in ((s.attn,s.n1),(s.mlp,s.n2)):
49
+ inp=nrm(x)
50
+ with torch.autocast("cuda",dtype=torch.bfloat16):
51
+ f=(sub(inp.to(torch.bfloat16),am) if sub is s.attn else sub(inp.to(torch.bfloat16)))
52
+ f=f.float(); g=Px(al*f,x); mt=Px(Gamma(xp,x,mom),x); mom=beta*mt+g; xp=x; x=Exp(x,mom)
53
+ pooled=(s.nf(x)*m).sum(1)/m.sum(1).clamp_min(1e-6)
54
+ halts.append(s.halt(pooled).squeeze(-1)); embeds.append(F.normalize(pooled,dim=-1))
55
+ return torch.stack(embeds,1), torch.stack(halts,1) # (B,LMAX,d), (B,LMAX); use depth-1 embeds[:,0]