Instructions to use ntedvs/irex with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- MLX
How to use ntedvs/irex 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("ntedvs/irex") 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 ntedvs/irex with MLX LM:
Generate or start a chat session
# Install MLX LM uv tool install mlx-lm # Generate some text mlx_lm.generate --model "ntedvs/irex" --prompt "Once upon a time"
- Atomic Chat
Download model.py from ntedvs/irex: direct link, hf CLI and curl.
- Browser
- Download file 5.25 kB
-
https://huggingface.co/ntedvs/irex/resolve/main/model.py
- Command line
-
hf download hf://ntedvs/irex/model.py
-
curl -L -o model.py https://huggingface.co/ntedvs/irex/resolve/main/model.py
5.25 kB
| """Tiny GPT in MLX: pre-norm, RMSNorm, RoPE, SwiGLU, tied embeddings. | |
| RoPE takes explicit per-row positions so batched decoding can left-pad. | |
| """ | |
| import math | |
| from dataclasses import dataclass, asdict | |
| import mlx.core as mx | |
| import mlx.nn as nn | |
| class Config: | |
| vocab: int = 4355 | |
| dim: int = 384 | |
| layers: int = 6 | |
| heads: int = 6 | |
| ffn_mult: float = 8 / 3 | |
| max_len: int = 384 | |
| dropout: float = 0.0 | |
| rope_base: float = 10000.0 | |
| compute: str = "bfloat16" # matmul dtype; master weights stay fp32 | |
| prefix: bool = False # prefix-LM: bidirectional attention over the english prompt | |
| def dict(self): | |
| return asdict(self) | |
| class Linear(nn.Module): | |
| """fp32 master weight, matmul in compute dtype (mixed precision).""" | |
| def __init__(self, i, o, dt): | |
| super().__init__() | |
| self.weight = mx.random.uniform(-i ** -0.5, i ** -0.5, (o, i)) | |
| self.dt = dt | |
| def __call__(self, x): | |
| return x.astype(self.dt) @ self.weight.astype(self.dt).T | |
| class RMSNorm(nn.Module): | |
| def __init__(self, d): | |
| super().__init__() | |
| self.weight = mx.ones((d,)) | |
| def __call__(self, x): | |
| return mx.fast.rms_norm(x, self.weight.astype(x.dtype), 1e-5) | |
| def rope(x, pos, base): | |
| # x: (B, H, T, D) pos: (B, T) or None (= 0..T-1, fused kernel) | |
| if pos is None: | |
| return mx.fast.rope(x, x.shape[-1], traditional=False, base=base, scale=1.0, offset=0) | |
| d = x.shape[-1] // 2 | |
| inv = mx.exp(-math.log(base) * mx.arange(d, dtype=mx.float32) / d) | |
| ang = pos[:, None, :, None].astype(mx.float32) * inv # B,1,T,d | |
| cos, sin = mx.cos(ang).astype(x.dtype), mx.sin(ang).astype(x.dtype) | |
| x1, x2 = x[..., :d], x[..., d:] | |
| return mx.concatenate([x1 * cos - x2 * sin, x1 * sin + x2 * cos], axis=-1) | |
| class Attention(nn.Module): | |
| def __init__(self, c: Config): | |
| super().__init__() | |
| self.h, self.hd, self.base = c.heads, c.dim // c.heads, c.rope_base | |
| dt = getattr(mx, c.compute) | |
| self.qkv = Linear(c.dim, 3 * c.dim, dt) | |
| self.out = Linear(c.dim, c.dim, dt) | |
| def __call__(self, x, pos, mask, cache=None): | |
| B, T, _ = x.shape | |
| q, k, v = mx.split(self.qkv(x), 3, axis=-1) | |
| q, k, v = (t.reshape(B, T, self.h, self.hd).transpose(0, 2, 1, 3) for t in (q, k, v)) | |
| q, k = rope(q, pos, self.base), rope(k, pos, self.base) | |
| if cache is not None: | |
| if cache[0] is not None: | |
| k = mx.concatenate([cache[0], k], axis=2) | |
| v = mx.concatenate([cache[1], v], axis=2) | |
| cache[0], cache[1] = k, v | |
| if isinstance(mask, mx.array): | |
| mask = mask.astype(q.dtype) | |
| o = mx.fast.scaled_dot_product_attention(q, k, v, scale=self.hd ** -0.5, mask=mask) | |
| return self.out(o.transpose(0, 2, 1, 3).reshape(B, T, -1)) | |
| class MLP(nn.Module): | |
| def __init__(self, c: Config): | |
| super().__init__() | |
| h = int(c.dim * c.ffn_mult / 64 + 0.999) * 64 | |
| dt = getattr(mx, c.compute) | |
| self.gate, self.up, self.down = Linear(c.dim, h, dt), Linear(c.dim, h, dt), Linear(h, c.dim, dt) | |
| def __call__(self, x): | |
| return self.down(nn.silu(self.gate(x)) * self.up(x)) | |
| class Block(nn.Module): | |
| def __init__(self, c: Config): | |
| super().__init__() | |
| self.n1, self.n2 = RMSNorm(c.dim), RMSNorm(c.dim) | |
| self.attn, self.mlp = Attention(c), MLP(c) | |
| self.drop = nn.Dropout(c.dropout) | |
| def __call__(self, x, pos, mask, cache=None): | |
| x = x + self.drop(self.attn(self.n1(x), pos, mask, cache)) | |
| return x + self.drop(self.mlp(self.n2(x))) | |
| class GPT(nn.Module): | |
| def __init__(self, c: Config): | |
| super().__init__() | |
| self.c = c | |
| self.emb = nn.Embedding(c.vocab, c.dim) | |
| self.blocks = [Block(c) for _ in range(c.layers)] | |
| self.norm = RMSNorm(c.dim) | |
| self.dt = getattr(mx, c.compute) | |
| self.drop = nn.Dropout(c.dropout) | |
| # GPT-2 style init: small embeddings, residual projections scaled by depth | |
| self.emb.weight = mx.random.normal(self.emb.weight.shape) * 0.02 | |
| for b in self.blocks: | |
| for lin in (b.attn.out, b.mlp.down): | |
| lin.weight = lin.weight * (1 / math.sqrt(2 * c.layers)) | |
| def __call__(self, x, pos=None, mask="causal", cache=None): | |
| B, T = x.shape | |
| h = self.drop(self.emb(x).astype(self.dt)) | |
| for i, b in enumerate(self.blocks): | |
| h = b(h, pos, mask, None if cache is None else cache[i]) | |
| h = self.norm(h) | |
| return h @ self.emb.weight.astype(self.dt).T | |
| def n_params(self): | |
| from mlx.utils import tree_flatten | |
| return sum(v.size for _, v in tree_flatten(self.parameters())) | |
| def quantize(model: GPT, bits: int = 8, group_size: int = 64): | |
| """Swap every Linear for an MLX QuantizedLinear (inference only).""" | |
| def q(lin): | |
| ref = nn.Linear(lin.weight.shape[1], lin.weight.shape[0], bias=False) | |
| ref.weight = lin.weight | |
| return nn.QuantizedLinear.from_linear(ref, group_size=group_size, bits=bits) | |
| for b in model.blocks: | |
| b.attn.qkv, b.attn.out = q(b.attn.qkv), q(b.attn.out) | |
| b.mlp.gate, b.mlp.up, b.mlp.down = q(b.mlp.gate), q(b.mlp.up), q(b.mlp.down) | |
| return model | |