Sol-Milkshake / modeling_sol_milkshake.py
j0no12's picture
Release Sol Milkshake 3M Base
7aa9bc4 verified
Raw History Blame Contribute Delete
11.1 kB
"""Exact 2,990,000-parameter Sol Milkshake model implemented in MLX."""
from __future__ import annotations
import math
from collections import OrderedDict
import mlx.core as mx
import mlx.nn as nn
from mlx.utils import tree_flatten
from sol_config import DEPLOYED_PARAMS, SolConfig
def unit_norm(x):
xf=x.astype(mx.float32)
return (xf*mx.rsqrt(mx.maximum(mx.sum(xf*xf,axis=-1,keepdims=True),mx.array(1e-12)))).astype(x.dtype)
def _count(module): return sum(int(x.size) for _,x in tree_flatten(module.trainable_parameters()))
def _rope(q,k,theta):
n,d=q.shape[-2:]; inv=mx.exp(mx.arange(0,d,2,dtype=mx.float32)*(-math.log(theta)/d))
ang=mx.arange(n,dtype=mx.float32)[:,None]*inv[None,:]; cs,sn=mx.cos(ang)[None,None],mx.sin(ang)[None,None]
def rot(x):
xf=x.astype(mx.float32); a,b=xf[...,::2],xf[...,1::2]
return mx.stack((a*cs-b*sn,a*sn+b*cs),axis=-1).reshape(x.shape).astype(x.dtype)
return rot(q),rot(k)
class TiedTokens(nn.Module):
def __init__(self,c):
super().__init__(); self.weight=unit_norm(mx.random.normal((c.vocab_size,c.d_model)))
def encode(self,ids): return self.weight[ids]
def decode(self,x,scale): return x.astype(mx.float32)@unit_norm(self.weight).astype(mx.float32).T*mx.exp(scale)
class TNGram(nn.Module):
"""Causal hashed CP retrieval for orders 2--5 (exactly 573,460 parameters)."""
def __init__(self,c):
super().__init__(); r=c.tn_rank; self.token_factor=mx.random.normal((c.vocab_size,r))*.02
self.hash_tables=mx.random.normal((4,4950,r))*.02; self.order_factors=mx.ones((4,r))
self.rank_to_hidden=mx.random.normal((r,c.d_model))*.02; self.output_bias=mx.zeros((c.d_model,))
self.gates=mx.full((4,),-2.0); self.order_mix=mx.eye(4,r); self.context_controls=mx.zeros((16,))
def __call__(self,ids):
b,t=ids.shape; out=mx.broadcast_to(self.output_bias,(b,t,self.output_bias.shape[0]))
for oi,order in enumerate(range(2,6)):
h=mx.zeros_like(ids); prod=mx.ones((b,t,self.token_factor.shape[1]),dtype=mx.float32); valid=mx.arange(t)>=order-1
for lag in range(1,order):
shifted=mx.zeros((b,t),dtype=ids.dtype) if lag>=t else mx.concatenate((mx.zeros((b,lag),dtype=ids.dtype),ids[:,:-lag]),axis=1)
h=(h*(131+oi*6)+shifted)%4950; prod=prod*self.token_factor[shifted].astype(mx.float32)
rank=prod*self.hash_tables[oi,h]*self.order_factors[oi]
rank=rank+self.order_mix[oi][None,None,:]*mx.mean(mx.tanh(self.context_controls[oi*4:(oi+1)*4]))
out=out+(rank@self.rank_to_hidden)*mx.sigmoid(self.gates[oi])*valid[None,:,None]
return out
class ShakeMemory(nn.Module):
"""Completed-chunk rolling memory (exactly 51,520 parameters)."""
def __init__(self,c):
super().__init__(); w=c.memory_width; self.chunk_size=c.chunk_size; self.slots=c.memory_slots
self.summary_router=nn.Linear(c.d_model,1,bias=False); self.compress=nn.Linear(c.d_model,w,bias=False)
self.query=nn.Linear(c.d_model,w,bias=False); self.key=nn.Linear(w,w,bias=False)
self.value=nn.Linear(w,w,bias=False); self.output=nn.Linear(w,c.d_model,bias=False)
self.slot_positions=mx.random.normal((c.memory_slots,c.d_model))*.01
self.decay=mx.zeros((w,)); self.update_gate=mx.zeros((w,))
def __call__(self,x):
b,t,d=x.shape; chunks=[]; reads=[]
for start in range(0,t,self.chunk_size):
end=min(t,start+self.chunk_size); length=end-start
if chunks:
mem=mx.stack(chunks[-self.slots:],axis=1); count=mem.shape[1]
mem=mem+self.compress(self.slot_positions[:count])[None]*mx.sigmoid(self.decay)[None,None]
a=mx.softmax((self.query(x[:,start:end]).astype(mx.float32)@self.key(mem).astype(mx.float32).transpose(0,2,1))/math.sqrt(mem.shape[-1]),axis=-1)
reads.append(self.output((a@self.value(mem).astype(mx.float32)).astype(x.dtype)))
else: reads.append(mx.zeros((b,length,d),dtype=x.dtype))
if length==self.chunk_size:
weights=mx.softmax(self.summary_router(x[:,start:end]).astype(mx.float32).squeeze(-1),axis=-1)
summary=mx.sum(weights[:,:,None]*x[:,start:end].astype(mx.float32),axis=1)
chunks.append(self.compress(summary.astype(x.dtype))*mx.sigmoid(self.update_gate))
return mx.concatenate(reads,axis=1)
class Attention(nn.Module):
def __init__(self,c):
super().__init__(); d=c.d_model; self.q=nn.Linear(d,d,bias=False); self.k=nn.Linear(d,c.n_kv_heads*c.head_dim,bias=False)
self.v=nn.Linear(d,c.n_kv_heads*c.head_dim,bias=False); self.o=nn.Linear(d,d,bias=False)
self.hq=c.n_q_heads; self.hk=c.n_kv_heads; self.hd=c.head_dim; self.theta=c.rope_theta
self.value_current=mx.zeros((1,)); self.value_prior=mx.zeros((1,))
def __call__(self,x,prior=None):
b,t,_=x.shape; q=unit_norm(self.q(x).reshape(b,t,self.hq,self.hd)).transpose(0,2,1,3)
k=unit_norm(self.k(x).reshape(b,t,self.hk,self.hd)).transpose(0,2,1,3); v=self.v(x).reshape(b,t,self.hk,self.hd).transpose(0,2,1,3)
if prior is not None: v=mx.sigmoid(self.value_current)*v+mx.sigmoid(self.value_prior)*prior
q,k=_rope(q,k,self.theta); k=mx.repeat(k,self.hq//self.hk,axis=1); vr=mx.repeat(v,self.hq//self.hk,axis=1)
scores=(q.astype(mx.float32)@k.astype(mx.float32).transpose(0,1,3,2))/math.sqrt(self.hd)
pos=mx.arange(t); mask=mx.where(pos[:,None]>=pos[None,:],0.0,-1e9)
attended=(mx.softmax(scores+mask[None,None],axis=-1)@vr.astype(mx.float32)).astype(x.dtype)
vh=unit_norm(vr); attended=attended-mx.sum(attended.astype(mx.float32)*vh.astype(mx.float32),axis=-1,keepdims=True).astype(x.dtype)*vh
return self.o(attended.transpose(0,2,1,3).reshape(b,t,-1)),v
class Block(nn.Module):
def __init__(self,c):
super().__init__(); self.attention=Attention(c); self.gate=nn.Linear(c.d_model,c.ffn_hidden,bias=False)
self.up=nn.Linear(c.d_model,c.ffn_hidden,bias=False); self.down=nn.Linear(c.ffn_hidden,c.d_model,bias=False)
self.attn_step=mx.zeros((c.d_model,)); self.ffn_step=mx.zeros((c.d_model,))
def __call__(self,x,prior=None):
a,v=self.attention(unit_norm(x),prior); x=unit_norm(x+mx.sigmoid(self.attn_step)*(unit_norm(a)-x))
f=self.down(nn.silu(self.gate(x))*self.up(x)); return unit_norm(x+mx.sigmoid(self.ffn_step)*(unit_norm(f)-x)),v
class SolMilkshake(nn.Module):
def __init__(self,c=None):
super().__init__(); self.config=c or SolConfig(); self.tokens=TiedTokens(self.config); self.tn_gram=TNGram(self.config)
self.shake_memory=ShakeMemory(self.config); self.blocks=[Block(self.config) for _ in range(5)]
self.pass_embeddings=mx.zeros((3,self.config.d_model)); self.mod_routers=[nn.Linear(self.config.d_model,1) for _ in range(6)]
self.logit_scale=mx.zeros((self.config.vocab_size,)); self.misc_controls=mx.zeros((12,)); self._verify_ledger()
def _select(self,x,old,router,capacity,lam):
delta=1-mx.sum(unit_norm(x).astype(mx.float32)*unit_norm(old).astype(mx.float32),axis=-1)
score=router(x).squeeze(-1)+lam*delta
positions=mx.arange(x.shape[1])
mask=(positions%4!=3) if capacity==.75 else (positions%2==0)
return mx.broadcast_to(mask[None,:],score.shape),mx.sigmoid(score)
def __call__(self,ids):
x=self.tokens.encode(ids); x=unit_norm(x+mx.sigmoid(self.misc_controls[0])*self.tn_gram(ids))
x=unit_norm(x+mx.sigmoid(self.misc_controls[1])*self.shake_memory(x)); x,v=self.blocks[0](x); ri=0
for p in range(3):
old=x; x=unit_norm(x+self.pass_embeddings[p]); prior=None
for bi in range(1,4):
proposal,prior=self.blocks[bi](x,prior)
if p==0: x=proposal
else:
mask,route_gate=self._select(x,old,self.mod_routers[ri],.75 if p==1 else .5,self.misc_controls[2+ri]); ri+=1
routed=unit_norm(x+route_gate[:,:,None]*(proposal-x))
x=mx.where(mask[:,:,None],routed,x)
x,_=self.blocks[4](x,v); return self.tokens.decode(x,self.logit_scale)
def parameter_ledger(self):
ledger=OrderedDict(token_embedding_tied_head=int(self.tokens.weight.size),attention_total=sum(_count(b.attention)-2 for b in self.blocks),
swiglu_total=sum(_count(b.gate)+_count(b.up)+_count(b.down) for b in self.blocks),tn_gram=_count(self.tn_gram),shake_memory=_count(self.shake_memory),
ngpt_steps=sum(int(b.attn_step.size+b.ffn_step.size) for b in self.blocks),pass_embeddings=int(self.pass_embeddings.size),
mod_routers=sum(_count(r) for r in self.mod_routers),logit_scale=int(self.logit_scale.size),value_residual=10,
misc_controls=int(self.misc_controls.size),xsa=0); ledger['deployed_total']=sum(ledger.values()); return dict(ledger)
def deployed_parameter_count(self): return _count(self)
def _verify_ledger(self):
expected={"token_embedding_tied_head":393216,"attention_total":491520,"swiglu_total":1474560,"tn_gram":573460,"shake_memory":51520,"ngpt_steps":1920,"pass_embeddings":576,"mod_routers":1158,"logit_scale":2048,"value_residual":10,"misc_controls":12,"xsa":0,"deployed_total":DEPLOYED_PARAMS}
if self.parameter_ledger()!=expected or self.deployed_parameter_count()!=DEPLOYED_PARAMS: raise RuntimeError(f"parameter ledger mismatch: {self.parameter_ledger()}, actual={self.deployed_parameter_count()}")
def future_token_causality_error(model,sequence_length=35):
ids=mx.arange(sequence_length,dtype=mx.int32)[None]%model.config.vocab_size; boundary=sequence_length//2
changed=mx.concatenate((ids[:,:boundary],(ids[:,boundary:]+17)%model.config.vocab_size),axis=1)
a,b=model(ids),model(changed); mx.eval(a,b); return float(mx.max(mx.abs(a[:,:boundary]-b[:,:boundary])))
def load_model(model_dir):
"""Load the standalone MLX model and frozen tokenizer from a snapshot."""
from pathlib import Path
from tokenizers import Tokenizer
model_dir = Path(model_dir)
model = SolMilkshake(SolConfig())
model.load_weights(str(model_dir / "model.npz"))
mx.eval(model.parameters())
tokenizer = Tokenizer.from_file(str(model_dir / "tokenizer.json"))
return model, tokenizer
def generate(model, tokenizer, prompt, max_new_tokens=64, temperature=0.0, seed=0):
"""Generate a single continuation with greedy or temperature sampling."""
mx.random.seed(seed)
tokens = tokenizer.encode(prompt).ids
if not tokens:
tokens = [0]
for _ in range(max_new_tokens):
context = tokens[-model.config.max_context:]
logits = model(mx.array([context], dtype=mx.int32))[0, -1]
if temperature and temperature > 0:
next_token = int(mx.random.categorical(logits / temperature).item())
else:
next_token = int(mx.argmax(logits).item())
tokens.append(next_token)
if next_token == 1:
break
return tokenizer.decode(tokens)