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 mxnet as d2l
import math
import random
import time
from mxnet import autograd, gluon, init, np, npx
from mxnet.gluon import nn
npx.set_np()
class RNNScratch(d2l.Module):
    """The RNN model implemented from scratch."""
    def __init__(self, num_inputs, num_hiddens, sigma=0.01):
        super().__init__()
        self.save_hyperparameters()
        self.W_xh = d2l.randn(num_inputs, num_hiddens) * sigma
        self.W_hh = d2l.randn(
            num_hiddens, num_hiddens) * sigma
        self.b_h = d2l.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 forward(self, inputs, state=None):
    if state is None:
        # Initial state with shape: (batch_size, num_hiddens)
        state = d2l.zeros((inputs.shape[1], self.num_hiddens),
                          ctx=inputs.ctx)
    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 = d2l.randn(5, 3), d2l.tensor([0, 2], dtype=d2l.int64)
(d2l.matmul(npx.one_hot(ids, 5), W) == W[ids]).all()
array(True)

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):
        super().__init__()
        self.save_hyperparameters()
        self.init_params()

    def init_params(self):
        self.W_e = d2l.randn(self.vocab_size, self.rnn.num_inputs)
        self.W_hq = d2l.randn(
            self.rnn.num_hiddens, self.vocab_size) * self.rnn.sigma
        self.b_q = d2l.zeros(self.vocab_size)
        for param in self.get_scratch_params():
            param.attach_grad()

    def training_step(self, batch):
        l = self.loss(self(*batch[:-1]), batch[-1])
        self.plot('ppl', d2l.exp(l), train=True)
        return l

    def validation_step(self, batch):
        l = self.loss(self(*batch[:-1]), batch[-1])
        self.plot('ppl', d2l.exp(l), train=False)

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.int64))
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, model):
    params = model.parameters()
    if not isinstance(params, list):
        params = [p.data() for p in params.values()]
    norm = math.sqrt(sum((p.grad ** 2).sum() for p in params))
    if norm > grad_clip_val:
        for param in params:
            param.grad[:] *= grad_clip_val / norm

(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

ppl_scratch = float(model.board.data['val_ppl'][-1].y)
print(f'validation perplexity {ppl_scratch:.1f}')
validation perplexity 89.9

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.34 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):
    outputs, state = tok.encode(prefix), None
    for i in range(len(outputs) - 1):  # Warm up on the prefix
        X = d2l.tensor([[outputs[i]]], ctx=device)
        _, state = self.rnn(self.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]]], ctx=device)
        rnn_outputs, state = self.rnn(self.embedding(X), state)
        logits = d2l.numpy(self.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, d2l.try_gpu())
'the time traveller. I had a fainting. I was not a little\ndemed. I was in the\nsharised. I was not a little\ndemed. I was in the\nsharised.'

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,
                        d2l.try_gpu(), temperature=T, rng=random.Random(0)))
the time travellerers of mening dery; to you towards of the presentlys com muchntself to about deepaccentousily ad
the time traveller. I had a pursued the Time Machine. I seemed to my mind and to letallicie. I had come out

Doing better = decoding strategies, later in this chapter.

Concise: the framework layer

Same interface as RNNScratch, one fused call:

class RNN(d2l.Module):
    """The RNN model implemented with high-level APIs."""
    def __init__(self, num_inputs, num_hiddens):
        super().__init__()
        self.save_hyperparameters()
        self.rnn = gluon.rnn.RNN(num_hiddens)

    def forward(self, inputs, H=None):
        if H is None:
            H, = self.rnn.begin_state(inputs.shape[1], ctx=inputs.ctx)
        outputs, (H, ) = self.rnn(inputs, (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_params(self):
        self.emb = nn.Embedding(self.vocab_size, self.rnn.num_inputs,
                                weight_initializer=init.Normal(1))
        self.linear = nn.Dense(self.vocab_size, flatten=False)
        self.initialize()

    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 haseringivertainward&ped swmentP�Yaaseter theseanped sat if one'

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

ppl_concise = float(model.board.data['val_ppl'][-1].y)
pred = model.predict('the time traveller', 30, data.tokenizer, d2l.try_gpu())
print(f'perplexity {ppl_concise:.1f}, {pred!r}')
perplexity 109.8, 'the time traveller forgotten in his hand and\nside that dawned\ninstant and in a patigention of'

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     199.9     89.9
 concise      68.1    109.8
  • 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.