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
Download native_model.py from OpenSML/OpenSML-150M: direct link, hf CLI and curl.
- Browser
- Download file 10.6 kB
-
https://huggingface.co/OpenSML/OpenSML-150M/resolve/main/native_model.py
- Command line
-
hf download hf://OpenSML/OpenSML-150M/native_model.py
-
curl -L -o native_model.py https://huggingface.co/OpenSML/OpenSML-150M/resolve/main/native_model.py
10.6 kB
| #!/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 | |
| 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 | |