class TinyLM(nnx.Module):
"""A small decoder-only transformer language model."""
def __init__(self, vocab_size, d_model=128, num_heads=2, num_blks=2,
max_len=64, rngs=None):
rngs = nnx.Rngs(0) if rngs is None else rngs
self.num_heads = num_heads
self.token_emb = nnx.Embed(vocab_size, d_model, rngs=rngs)
self.pos_emb = nnx.Embed(max_len, d_model, rngs=rngs)
self.blks = nnx.List([nnx.Dict(
norm1=nnx.LayerNorm(d_model, rngs=rngs),
qkv=nnx.Linear(d_model, 3 * d_model, rngs=rngs),
proj=nnx.Linear(d_model, d_model, rngs=rngs),
norm2=nnx.LayerNorm(d_model, rngs=rngs),
mlp1=nnx.Linear(d_model, 4 * d_model, rngs=rngs),
mlp2=nnx.Linear(4 * d_model, d_model, rngs=rngs))
for _ in range(num_blks)])
self.norm = nnx.LayerNorm(d_model, rngs=rngs)
self.head = nnx.Linear(d_model, vocab_size, rngs=rngs)
def attention(self, blk, X):
B, T, D = X.shape
q, k, v = jnp.split(blk['qkv'](X), 3, axis=-1)
q, k, v = (u.reshape(B, T, self.num_heads, -1) for u in (q, k, v))
Y = jax.nn.dot_product_attention(q, k, v, is_causal=True)
return blk['proj'](Y.reshape(B, T, D))
def __call__(self, X):
H = self.token_emb(X) + self.pos_emb(jnp.arange(X.shape[1]))
for blk in self.blks:
H = H + self.attention(blk, blk['norm1'](H))
H = H + blk['mlp2'](jax.nn.gelu(blk['mlp1'](blk['norm2'](H))))
return self.head(self.norm(H))