Text Generation
MLX
Safetensors
English
pretraining
from-scratch
small-language-model
post-training
silicon
Instructions to use OpenSML/OpenSML-150M with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- MLX
How to use OpenSML/OpenSML-150M 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("OpenSML/OpenSML-150M") 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 OpenSML/OpenSML-150M with MLX LM:
Generate or start a chat session
# Install MLX LM uv tool install mlx-lm # Generate some text mlx_lm.generate --model "OpenSML/OpenSML-150M" --prompt "Once upon a time"
- Atomic Chat
File size: 10,552 Bytes
8662ab2 | 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 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 | #!/usr/bin/env python3
"""MLX-only transformer model for small causal language modeling."""
from __future__ import annotations
import math
from dataclasses import dataclass
from typing import List, Optional, Tuple
import mlx.core as mx
import mlx.nn as nn
import mlx.nn.losses as losses
@dataclass
class TransformerConfig:
vocab_size: int
max_seq_len: int = 1024
d_model: int = 768
n_heads: int = 12
n_kv_heads: Optional[int] = None
n_layers: int = 12
mlp_ratio: float = 4.0
mlp_multiple_of: int = 256
rope_base: float = 10000.0
bias: bool = False
qk_norm: bool = True
attention_impl: str = "fast"
loss_dtype: str = "float32"
ce_impl: str = "reference"
ffn_impl: str = "reference"
def __post_init__(self) -> None:
if self.vocab_size <= 0:
raise ValueError("vocab_size must be > 0")
if self.max_seq_len <= 0:
raise ValueError("max_seq_len must be > 0")
if self.d_model <= 0:
raise ValueError("d_model must be > 0")
if self.n_heads <= 0 or self.d_model % self.n_heads != 0:
raise ValueError("n_heads must divide d_model")
if self.n_kv_heads is None:
self.n_kv_heads = self.n_heads
if self.n_kv_heads <= 0 or self.n_heads % self.n_kv_heads != 0:
raise ValueError("n_kv_heads must be positive and divide n_heads")
if self.n_layers <= 0:
raise ValueError("n_layers must be > 0")
if self.mlp_ratio <= 1.0:
raise ValueError("mlp_ratio must be > 1.0")
if self.mlp_multiple_of <= 0:
raise ValueError("mlp_multiple_of must be > 0")
if self.attention_impl not in {"fast", "vanilla"}:
raise ValueError("attention_impl must be one of: fast, vanilla")
if self.loss_dtype not in {"float32", "model"}:
raise ValueError("loss_dtype must be one of: float32, model")
if self.ce_impl not in {"reference", "metal"}:
raise ValueError("ce_impl must be reference or metal")
if self.ce_impl == "metal" and self.loss_dtype != "float32":
raise ValueError("Metal CE uses FP32 reductions")
if self.ffn_impl not in {"reference", "packed-metal"}:
raise ValueError("ffn_impl must be reference or packed-metal")
def _tree_leaves(tree):
if isinstance(tree, dict):
for v in tree.values():
yield from _tree_leaves(v)
elif isinstance(tree, (list, tuple)):
for v in tree:
yield from _tree_leaves(v)
else:
yield tree
def count_parameters(model: nn.Module) -> int:
total = 0
for leaf in _tree_leaves(model.parameters()):
if isinstance(leaf, mx.array):
n = 1
for d in leaf.shape:
n *= int(d)
total += n
return total
class CausalSelfAttention(nn.Module):
def __init__(self, cfg: TransformerConfig):
super().__init__()
self.n_heads = cfg.n_heads
self.n_kv_heads = cfg.n_kv_heads
self.head_dim = cfg.d_model // cfg.n_heads
self.q_dim = cfg.n_heads * self.head_dim
self.kv_dim = cfg.n_kv_heads * self.head_dim
self.max_seq_len = cfg.max_seq_len
self.qkv_proj = nn.Linear(cfg.d_model, self.q_dim + 2 * self.kv_dim, bias=cfg.bias)
self.proj = nn.Linear(cfg.d_model, cfg.d_model, bias=cfg.bias)
self.q_norm = nn.RMSNorm(self.head_dim) if cfg.qk_norm else None
self.k_norm = nn.RMSNorm(self.head_dim) if cfg.qk_norm else None
self.rope = nn.RoPE(self.head_dim, base=cfg.rope_base)
self.attention_impl = cfg.attention_impl
def __call__(
self,
x: mx.array,
cache: Optional[Tuple[mx.array, mx.array]] = None,
) -> Tuple[mx.array, Tuple[mx.array, mx.array]]:
bsz, seqlen, d_model = x.shape
qkv = self.qkv_proj(x)
q, k, v = mx.split(qkv, [self.q_dim, self.q_dim + self.kv_dim], axis=-1)
def split_heads(t: mx.array, n_heads: int) -> mx.array:
t = t.reshape(bsz, seqlen, n_heads, self.head_dim)
return t.transpose(0, 2, 1, 3)
q = split_heads(q, self.n_heads)
k = split_heads(k, self.n_kv_heads)
v = split_heads(v, self.n_kv_heads)
if self.q_norm is not None:
q = self.q_norm(q)
k = self.k_norm(k)
cache_offset = 0
if cache is not None and cache[0] is not None:
cache_offset = int(cache[0].shape[2])
q = self.rope(q, offset=cache_offset)
k = self.rope(k, offset=cache_offset)
if cache is not None:
k_cache, v_cache = cache
if k_cache is not None and v_cache is not None:
k = mx.concatenate([k_cache, k], axis=2)
v = mx.concatenate([v_cache, v], axis=2)
q_len = q.shape[2]
k_len = k.shape[2]
if q_len > self.max_seq_len or k_len > self.max_seq_len:
raise ValueError(
f"Sequence length {q_len}/{k_len} exceeds max_seq_len={self.max_seq_len}"
)
scale = 1.0 / math.sqrt(self.head_dim)
# Cache compact KV heads, not the expanded heads used by vanilla GQA.
next_cache = (k, v)
if self.attention_impl == "fast":
attn = mx.fast.scaled_dot_product_attention(
q,
k,
v,
scale=scale,
mask="causal",
)
else:
repeats = self.n_heads // self.n_kv_heads
if repeats > 1:
k = mx.repeat(k, repeats, axis=1)
v = mx.repeat(v, repeats, axis=1)
full_mask = nn.MultiHeadAttention.create_additive_causal_mask(k_len)
mask = full_mask[k_len - q_len : k_len, :k_len]
mask = mx.expand_dims(mx.expand_dims(mask, axis=0), axis=0).astype(q.dtype)
scores = mx.matmul(q, k.transpose(0, 1, 3, 2)) * scale
scores = scores + mask
probs = mx.softmax(scores.astype(mx.float32), axis=-1).astype(v.dtype)
attn = mx.matmul(probs, v)
out = attn.transpose(0, 2, 1, 3).reshape(bsz, q_len, d_model)
out = self.proj(out)
return out, next_cache
class FeedForward(nn.Module):
def __init__(self, cfg: TransformerConfig):
super().__init__()
# SwiGLU uses three projections, so 2/3 keeps parameter count comparable
# to a conventional 4x two-projection MLP.
hidden_dim = int((2.0 / 3.0) * cfg.d_model * cfg.mlp_ratio)
hidden_dim = cfg.mlp_multiple_of * math.ceil(hidden_dim / cfg.mlp_multiple_of)
self.hidden_dim = hidden_dim
self.ffn_impl = cfg.ffn_impl
self.gate_up = nn.Linear(cfg.d_model, 2 * hidden_dim, bias=cfg.bias)
self.down = nn.Linear(hidden_dim, cfg.d_model, bias=cfg.bias)
def __call__(self, x: mx.array) -> mx.array:
gate_up = self.gate_up(x)
if self.ffn_impl == "packed-metal":
try:
from .fast_swiglu import packed_swiglu
except ImportError:
from fast_swiglu import packed_swiglu
return self.down(packed_swiglu(gate_up))
gate, up = mx.split(gate_up, [self.hidden_dim], axis=-1)
return self.down(nn.silu(gate) * up)
class TransformerBlock(nn.Module):
def __init__(self, cfg: TransformerConfig):
super().__init__()
self.norm1 = nn.RMSNorm(cfg.d_model)
self.attn = CausalSelfAttention(cfg)
self.norm2 = nn.RMSNorm(cfg.d_model)
self.ffn = FeedForward(cfg)
def __call__(
self,
x: mx.array,
cache: Optional[Tuple[mx.array, mx.array]] = None,
) -> Tuple[mx.array, Tuple[mx.array, mx.array]]:
attn_out, new_cache = self.attn(self.norm1(x), cache=cache)
x = x + attn_out
x = x + self.ffn(self.norm2(x))
return x, new_cache
class TransformerLM(nn.Module):
def __init__(self, cfg: TransformerConfig):
super().__init__()
self.cfg = cfg
self.embed = nn.Embedding(cfg.vocab_size, cfg.d_model)
self.blocks = [TransformerBlock(cfg) for _ in range(cfg.n_layers)]
# MLX >= 0.32.2 supplies the memory-reduced Metal RMSNorm backward
# for every nn.RMSNorm here, including block and Q/K norms. Keep the
# native primitive so compilation/autodiff select that implementation.
self.norm = nn.RMSNorm(cfg.d_model)
def logits(
self,
input_ids: mx.array,
caches: Optional[List[Tuple[mx.array, mx.array]]] = None,
) -> Tuple[mx.array, List[Tuple[mx.array, mx.array]]]:
x = self.embed(input_ids)
if caches is None:
caches = [None] * len(self.blocks)
new_caches = []
for block, block_cache in zip(self.blocks, caches):
x, next_cache = block(x, cache=block_cache)
new_caches.append(next_cache)
x = self.norm(x)
weight = self.embed.weight
logits = mx.matmul(x, weight.transpose(1, 0))
return logits, new_caches
def __call__(
self,
input_ids: mx.array,
targets: Optional[mx.array] = None,
ignore_index: int = -100,
) -> dict:
logits, _ = self.logits(input_ids, caches=None)
out = {"logits": logits}
if targets is None:
return out
loss_logits = logits.astype(mx.float32) if self.cfg.loss_dtype == "float32" else logits
mask = targets != ignore_index
safe_targets = mx.where(mask, targets, mx.zeros_like(targets))
if self.cfg.ce_impl == "metal":
try:
from .fast_loss import cross_entropy
except ImportError:
from fast_loss import cross_entropy
per_token = cross_entropy(logits, safe_targets)
else:
per_token = losses.cross_entropy(loss_logits, safe_targets, reduction="none")
mask = mask.astype(per_token.dtype)
denom = mx.maximum(mask.sum(), mx.array(1.0, dtype=per_token.dtype))
loss = (per_token * mask).sum() / denom
out["loss"] = loss
out["token_loss"] = per_token
return out
def step(
self,
token_ids: mx.array,
caches: Optional[List[Tuple[mx.array, mx.array]]] = None,
) -> Tuple[mx.array, List[Tuple[mx.array, mx.array]]]:
# token_ids shape: [B, 1]
logits, new_caches = self.logits(token_ids, caches=caches)
return logits[:, -1, :], new_caches
|