Text Generation
MLX
English
sol_milkshake
causal-lm
decoder-only
small-language-model
recurrent-depth
ngpt
research
Instructions to use solintellegence/Sol-Milkshake with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- MLX
How to use solintellegence/Sol-Milkshake with MLX:
# Make sure mlx-lm is installed # pip install --upgrade mlx-lm # if on a CUDA device, also pip install mlx[cuda] # Generate text with mlx-lm from mlx_lm import load, generate model, tokenizer = load("solintellegence/Sol-Milkshake") prompt = "Once upon a time in" text = generate(model, tokenizer, prompt=prompt, verbose=True) - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- LM Studio
- MLX LM
How to use solintellegence/Sol-Milkshake with MLX LM:
Generate or start a chat session
# Install MLX LM uv tool install mlx-lm # Generate some text mlx_lm.generate --model "solintellegence/Sol-Milkshake" --prompt "Once upon a time"
- Atomic Chat
Download modeling_sol_milkshake.py from solintellegence/Sol-Milkshake: direct link, hf CLI and curl.
- Browser
- Download file 11.1 kB
-
https://huggingface.co/solintellegence/Sol-Milkshake/resolve/main/modeling_sol_milkshake.py
- Command line
-
hf download hf://solintellegence/Sol-Milkshake/modeling_sol_milkshake.py
-
curl -L -o modeling_sol_milkshake.py https://huggingface.co/solintellegence/Sol-Milkshake/resolve/main/modeling_sol_milkshake.py
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) | |