Upload omniem_model.py with huggingface_hub
Browse files- 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]
|