class GPT(nnx.Module):
"""Decoder-only transformer language model built from configurable
blocks."""
class CausalAttention(nnx.Module):
"""Multi-head causal self-attention, optionally rotary."""
def __init__(self, num_hiddens, num_heads, bias=False, rope=False,
rngs=None):
rngs = nnx.Rngs(0) if rngs is None else rngs
self.num_heads, self.rope = num_heads, rope
self.W_qkv = nnx.Linear(num_hiddens, 3 * num_hiddens,
use_bias=bias, rngs=rngs)
self.W_o = nnx.Linear(num_hiddens, num_hiddens, use_bias=bias,
rngs=rngs)
def _rope(self, x):
# x: (batch, num_steps, num_heads, head_dim)
d = x.shape[-1]
pos = jnp.arange(x.shape[1], dtype=jnp.float32)
inv_freq = 10000.0 ** (-jnp.arange(0, d, 2) / d)
theta = pos[:, None] * inv_freq[None, :]
cos = jnp.cos(theta)[:, None, :] # broadcast over heads
sin = jnp.sin(theta)[:, None, :]
x1, x2 = x[..., 0::2], x[..., 1::2]
return jnp.stack([x1 * cos - x2 * sin,
x1 * sin + x2 * cos], -1).reshape(x.shape)
def __call__(self, X, *_):
B, T, D = X.shape
q, k, v = jnp.split(self.W_qkv(X), 3, axis=-1)
q, k, v = (u.reshape(B, T, self.num_heads, -1)
for u in (q, k, v))
if self.rope:
q, k = self._rope(q), self._rope(k)
Y = jax.nn.dot_product_attention(q, k, v, is_causal=True)
return self.W_o(Y.reshape(B, T, D)), None
def __init__(self, vocab_size, num_hiddens=256, num_heads=8, num_blks=6,
max_len=1024, pos='rope', norm='rms', act='swiglu',
pre_norm=True, bias=False, dropout=0, rngs=None):
rngs = nnx.Rngs(0) if rngs is None else rngs
self.pos, self.max_len = pos, max_len
init = nnx.initializers.normal(0.02)
self.token_emb = nnx.Embed(vocab_size, num_hiddens,
embedding_init=init, rngs=rngs)
if pos == 'learned':
self.pos_emb = nnx.Embed(max_len, num_hiddens,
embedding_init=init, rngs=rngs)
attn = lambda rngs: self.CausalAttention(
num_hiddens, num_heads, bias, rope=(pos == 'rope'), rngs=rngs)
self.blks = nnx.List([
d2l.TransformerBlock(num_hiddens, num_heads, dropout, norm, act,
pre_norm, bias, attn_factory=attn,
rngs=rngs)
for _ in range(num_blks)])
self.norm = (nnx.RMSNorm if norm == 'rms'
else nnx.LayerNorm)(num_hiddens, rngs=rngs)
def __call__(self, X):
H = self.token_emb(X)
if self.pos == 'learned':
H = H + self.pos_emb(jnp.arange(X.shape[1]))
for blk in self.blks:
H = blk(H)
return self.token_emb.attend(self.norm(H))