Implementing RNN Language Models

Implementing RNN Language Models

An RNN language model on The Time Machine, over the 1,024-token BPE vocabulary, built twice: from raw tensor ops, then with the framework’s recurrent layer. Four pieces:

  1. RNN cell: the recurrence \mathbf{h}_t = \tanh(\mathbf{W}_{xh} \mathbf{x}_t + \mathbf{W}_{hh} \mathbf{h}_{t-1} + \mathbf{b}).
  2. Embedding: token ids become trainable vectors (no more one-hot).
  3. Output head: hidden state to vocab logits at every step.
  4. Gradient clipping + training + generation.

The RNN cell

Parameters: \mathbf{W}_{xh}, \mathbf{W}_{hh}, \mathbf{b}. Initialize randomly, scaled to keep activations sensible:

%matplotlib inline
from d2l import jax as d2l
from flax import nnx
import jax
from jax import numpy as jnp
import math
import random
import time
class RNNScratch(nnx.Module):
    """The RNN model implemented from scratch."""
    def __init__(self, num_inputs, num_hiddens, sigma=0.01, rngs=None):
        rngs = nnx.Rngs(0) if rngs is None else rngs
        self.num_inputs, self.num_hiddens = num_inputs, num_hiddens
        self.sigma = sigma
        self.W_xh = nnx.Param(
            rngs.params.normal((num_inputs, num_hiddens)) * sigma)
        self.W_hh = nnx.Param(
            rngs.params.normal((num_hiddens, num_hiddens)) * sigma)
        self.b_h = nnx.Param(jnp.zeros(num_hiddens))

Forward, unrolled

Walk a length-T input one step at a time, carrying the hidden state forward:

@d2l.add_to_class(RNNScratch)
def __call__(self, inputs, state=None):
    if state is None:
        # Initial state with shape: (batch_size, num_hiddens)
        state = jnp.zeros((inputs.shape[1], self.num_hiddens))
    outputs = []
    for X in inputs:  # Shape of inputs: (num_steps, batch_size, num_inputs)
        state = d2l.tanh(d2l.matmul(X, self.W_xh) +
                         d2l.matmul(state, self.W_hh) + self.b_h)
        outputs.append(state)
    return outputs, state
batch_size, num_inputs, num_hiddens, num_steps = 2, 16, 32, 100
rnn = RNNScratch(num_inputs, num_hiddens)
X = d2l.ones((num_steps, batch_size, num_inputs))
outputs, state = rnn(X)

Sanity check on output shapes:

def check_len(a, n):
    """Check the length of a list."""
    assert len(a) == n, f'list\'s length {len(a)} != expected length {n}'

def check_shape(a, shape):
    """Check the shape of a tensor."""
    assert a.shape == shape, \
            f'tensor\'s shape {a.shape} != expected shape {shape}'

check_len(outputs, num_steps)
check_shape(outputs[0], (batch_size, num_hiddens))
check_shape(state, (batch_size, num_hiddens))

Embeddings, not one-hot

With |\mathcal{V}| = 1{,}024, one-hot inputs waste a 1,024-wide multiply per step on a vector of zeros.

  • Embedding lookup: row i of a trainable \mathbf{W}_e \in \mathbb{R}^{|\mathcal{V}| \times d}.
  • Same map as one-hot \times matrix, but the rows are learned.
@d2l.add_to_class(RNNLMScratch)
def embedding(self, X):
    # Output shape: (num_steps, batch_size, num_inputs)
    return self.W_e[X.T]

The equivalence, verified:

W, ids = jax.random.normal(d2l.get_key(), (5, 3)), d2l.tensor([0, 2])
jnp.allclose(jax.nn.one_hot(ids, 5) @ W, W[ids])
Array(True, dtype=bool)

Wrapping as a language model

Embedding in, vocab-sized projection out; plot perplexity instead of loss:

class RNNLMScratch(d2l.Classifier):
    """The RNN-based language model implemented from scratch."""
    def __init__(self, rnn, vocab_size, lr=0.01, rngs=None):
        super().__init__()
        self.save_hyperparameters(ignore=['rnn', 'rngs'])
        self.rnn = rnn
        rngs = nnx.Rngs(1) if rngs is None else rngs
        self.W_e = nnx.Param(
            rngs.params.normal((vocab_size, rnn.num_inputs)))
        self.W_hq = nnx.Param(rngs.params.normal(
            (rnn.num_hiddens, vocab_size)) * rnn.sigma)
        self.b_q = nnx.Param(jnp.zeros(vocab_size))

    def training_step(self, batch):
        return self.loss(self(*batch[:-1]), batch[-1])

    def validation_step(self, batch):
        return self.loss(self(*batch[:-1]), batch[-1])

    def plot(self, key, value, train):
        # The train/val steps run inside `@nnx.jit` and only return the mean
        # loss: plotting a tracer from there would crash the board's drawing
        # thread. `Trainer.fit_epoch` instead calls this with the materialized
        # loss (outside jit), which we relabel as perplexity for parity with
        # the other tabs.
        if key == 'loss':
            key, value = 'ppl', d2l.exp(value)
        super().plot(key, value, train)

Output projection

Project every hidden state through the shared head, then a shape smoke test: (batch, steps) ids in, (batch, steps, vocab) logits out:

@d2l.add_to_class(RNNLMScratch)
def output_layer(self, rnn_outputs):
    outputs = [d2l.matmul(H, self.W_hq) + self.b_q for H in rnn_outputs]
    return d2l.stack(outputs, 1)

@d2l.add_to_class(RNNLMScratch)
def forward(self, X, state=None):
    embs = self.embedding(X)
    rnn_outputs, _ = self.rnn(embs, state)
    return self.output_layer(rnn_outputs)
model = RNNLMScratch(rnn, vocab_size=1024)
outputs = model(d2l.ones((batch_size, num_steps), dtype=d2l.int32))
check_shape(outputs, (batch_size, num_steps, 1024))

Gradient clipping

Backprop through T steps multiplies T Jacobians, one explosion-prone product. Clip the gradient onto a ball of radius \theta before each update:

\mathbf{g} \leftarrow \min\!\left(1, \frac{\theta}{\|\mathbf{g}\|}\right)\mathbf{g}.

@d2l.add_to_class(d2l.Trainer)
def clip_gradients(self, grad_clip_val, grads):
    grad_leaves, _ = jax.tree_util.tree_flatten(grads)
    norm = jnp.sqrt(sum(jnp.vdot(x, x) for x in grad_leaves))
    clip = lambda grad: jnp.where(norm < grad_clip_val,
                                  grad, grad * (grad_clip_val / norm))
    return jax.tree_util.tree_map(clip, grads)

(PyTorch/MXNet: called by fit_epoch; TF: inside the compiled step; JAX: optax.clip_by_global_norm does it inside fit.)

Training

50k windows of 32 BPE tokens, batch 1024, 10 epochs, clip at 1. Fresh zero state per window = truncated BPTT:

data = d2l.TimeMachine(batch_size=1024, num_steps=32,
                       num_train=50000, num_val=5000)
rnn = RNNScratch(num_inputs=64, num_hiddens=128)
model = RNNLMScratch(rnn, vocab_size=len(data.vocab), lr=4)
trainer = d2l.Trainer(max_epochs=10, gradient_clip_val=1, num_gpus=1)
model.board.yscale = 'log'  # perplexity spans orders of magnitude
t0 = time.time()
trainer.fit(model, data)
t_scratch = time.time() - t0

Reading the perplexity

total_loss = num_tokens = 0
for X_val, y_val in data.val_dataloader():
    losses = model.loss(model(X_val), y_val, averaged=False)
    total_loss += float(losses.sum())
    num_tokens += losses.size
ppl_scratch = math.exp(total_loss / num_tokens)
print(f'validation perplexity {ppl_scratch:.1f}')
validation perplexity 88.3

Val ppl ~90–100 over 1,024 tokens vs. char-level ppl ~7 over 27: not comparable. Convert to bits per byte:

ids = d2l.numpy(data.X[data.num_train:data.num_train+data.num_val, 0]).tolist()
bytes_per_token = len(data.tokenizer.decode(ids).encode('utf-8')) / len(ids)
print(f'{bytes_per_token:.2f} bytes/token, '
      f'{math.log2(ppl_scratch) / bytes_per_token:.2f} bits per byte')
2.78 bytes/token, 2.33 bits per byte

~2.4 bpb beats the char-trigram baseline’s 2.68 bpb: the “worse” perplexity is the better language model.

Generating text

Warm up on the prefix, then feed each chosen token back in. Greedy (T=0) or temperature sampling:

@d2l.add_to_class(RNNLMScratch)
def predict(self, prefix, num_tokens, tok, device=None, temperature=0.0,
            rng=None):
    model = nnx.view(self, deterministic=True, use_running_average=True,
                     raise_if_not_found=False)
    outputs, state = tok.encode(prefix), None
    for i in range(len(outputs) - 1):  # Warm up on the prefix
        X = d2l.tensor([[outputs[i]]])
        _, state = model.rnn(model.embedding(X), state)
    rng = random.Random() if rng is None else rng
    for _ in range(num_tokens):  # Generate num_tokens continuation tokens
        X = d2l.tensor([[outputs[-1]]])
        rnn_outputs, state = model.rnn(model.embedding(X), state)
        logits = d2l.numpy(model.output_layer(rnn_outputs))[0, 0]
        if temperature == 0:
            outputs.append(int(logits.argmax()))
        else:
            weights = [math.exp(l) for l in
                       (logits - logits.max()) / temperature]
            outputs.append(rng.choices(range(len(weights)), weights)[0])
    return tok.decode(outputs)

Greedy vs. temperature

model.predict('the time traveller', 50, data.tokenizer)
"the time traveller.\n\n'In the Time Traveller, and\nsossibly the Time Traveller.\n\n'In the Time Traveller, and\nsossibly the Time Traveller.\n\n'In the ...

Greedy is fluent, then circles. Sampling breaks the loop at the price of stranger choices:

for T in (1.0, 0.5):
    print(model.predict('the time traveller', 30, data.tokenizer,
                        temperature=T, rng=random.Random(0)))
the time traveller soly in a little recognive out of pastentre Neverildren of explain genepecies than
the time traveller.

'Soverny, but the half-dimension of the probleiss, and I had come

Doing better = decoding strategies, later in this chapter.

Concise: the framework layer

Same interface as RNNScratch, one fused call:

class RNN(nnx.Module):
    """The RNN model implemented with high-level APIs."""
    def __init__(self, num_inputs, num_hiddens, rngs=None):
        rngs = nnx.Rngs(0) if rngs is None else rngs
        self.num_inputs, self.num_hiddens = num_inputs, num_hiddens
        self.rnn = nnx.RNN(
            nnx.SimpleCell(num_inputs, num_hiddens, rngs=rngs),
            time_major=True, return_carry=True, rngs=rngs)

    def __call__(self, inputs, H=None):
        H, outputs = self.rnn(inputs, initial_carry=H)
        return outputs, H

The LM wrapper is inherited: swap in framework embedding and dense layers:

class RNNLM(d2l.RNNLMScratch):
    """The RNN-based language model implemented with high-level APIs."""
    def __init__(self, rnn, vocab_size, lr=0.01, rngs=None):
        d2l.Classifier.__init__(self)
        self.save_hyperparameters(ignore=['rnn', 'rngs'])
        self.rnn = rnn
        rngs = nnx.Rngs(2) if rngs is None else rngs
        self.emb = nnx.Embed(vocab_size, rnn.num_inputs,
                             embedding_init=nnx.initializers.normal(1.0),
                             rngs=rngs)
        self.linear = nnx.Linear(rnn.num_hiddens, vocab_size, rngs=rngs)

    def embedding(self, X):
        return self.emb(X.T)

    def output_layer(self, hiddens):
        return d2l.swapaxes(self.linear(hiddens), 0, 1)

Sanity check, then train

Untrained model generates byte soup, but the wiring (tokenizer to model and back) is sound:

rnn = RNN(num_inputs=64, num_hiddens=128)
model = RNNLM(rnn, vocab_size=len(data.vocab), lr=4)
model.predict('it has', 20, data.tokenizer)
'it has m\x1d e5asause sornessare ex( lab� Thenhing�imes now strange'

Same trainer, same data:

trainer = d2l.Trainer(max_epochs=10, gradient_clip_val=1, num_gpus=1)
model.board.yscale = 'log'
t0 = time.time()
trainer.fit(model, data)
t_concise = time.time() - t0

total_loss = num_tokens = 0
for X_val, y_val in data.val_dataloader():
    losses = model.loss(model(X_val), y_val, averaged=False)
    total_loss += float(losses.sum())
    num_tokens += losses.size
ppl_concise = math.exp(total_loss / num_tokens)
pred = model.predict('the time traveller', 30, data.tokenizer)
print(f'perplexity {ppl_concise:.1f}, {pred!r}')
perplexity 107.5, "the time traveller, and\nthemathem of\nthem, and,\nthe silent of of\nthe Time Traveller. 'You"

Scratch vs. concise, measured

print(f'{"model":>8} {"time (s)":>9} {"val ppl":>8}')
for name, t, p in [('scratch', t_scratch, ppl_scratch),
                   ('concise', t_concise, ppl_concise)]:
    print(f'{name:>8} {t:>9.1f} {p:>8.1f}')
   model  time (s)  val ppl
 scratch      21.0     88.3
 concise      15.0    107.5
  • PyTorch/MXNet: fused kernel wins severalfold; per-step launch overhead dominates at this size.
  • JAX: both versions JIT-compile; the gap is small.
  • TF: SimpleRNN has no fused GPU kernel; the compiled scratch loop matches it.

Recap

  • BPE-token RNN LM: embedding + hand-rolled cell + shared head + cross-entropy.
  • Gradient clipping is mandatory for stable RNN training.
  • Fresh state per window = truncated BPTT through num_steps tokens.
  • Compare models across tokenizers by bits per byte, never perplexity.
  • Greedy decoding loops; temperature trades repetition for noise; decoding gets its own section.
  • The same scaffold takes any cell (LSTM, GRU): only the recurrence changes. Coming next.