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
File size: 11,100 Bytes
7aa9bc4 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 | """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)
|