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