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
File size: 5,254 Bytes
521b329 | 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 | """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
@dataclass
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
|