%matplotlib inline
from d2l import torch as d2l
import math
import random
import time
import torch
from torch import nn
from torch.nn import functional as FAn 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:
Parameters: \mathbf{W}_{xh}, \mathbf{W}_{hh}, \mathbf{b}. Initialize randomly, scaled to keep activations sensible:
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 = nn.Parameter(
d2l.randn(num_inputs, num_hiddens) * sigma)
self.W_hh = nn.Parameter(
d2l.randn(num_hiddens, num_hiddens) * sigma)
self.b_h = nn.Parameter(d2l.zeros(num_hiddens))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),
device=inputs.device)
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, stateSanity 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))With |\mathcal{V}| = 1{,}024, one-hot inputs waste a 1,024-wide multiply per step on a vector of zeros.
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 = nn.Parameter(
d2l.randn(self.vocab_size, self.rnn.num_inputs))
self.W_hq = nn.Parameter(
d2l.randn(
self.rnn.num_hiddens, self.vocab_size) * self.rnn.sigma)
self.b_q = nn.Parameter(d2l.zeros(self.vocab_size))
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)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)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}.
(PyTorch/MXNet: called by fit_epoch; TF: inside the compiled step; JAX: optax.clip_by_global_norm does it inside fit.)
50k windows of 32 BPE tokens, batch 1024, 10 epochs, clip at 1. Fresh zero state per window = truncated BPTT:
validation perplexity 89.3
Val ppl ~90–100 over 1,024 tokens vs. char-level ppl ~7 over 27: not comparable. Convert to 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.
Warm up on the prefix, then feed each chosen token back in. Greedy (T=0) or temperature sampling:
@d2l.add_to_class(RNNLMScratch)
@torch.no_grad() # inference only: no autograd graph needed
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]]], device=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]]], device=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)'the time travellerularlyification, and the little people, and the little people, and the little people, and the little people, and the little people, and ...
Greedy is fluent, then circles. Sampling breaks the loop at the price of stranger choices:
the time travellerhingions of the roantound.
'So,' said the refms soible I felt other flicks of gurnod
the time travellerour to the find
its the Time Machine, the Palace of Green Porcelain.
'The Time
Doing better = decoding strategies, later in this chapter.
Same interface as RNNScratch, one fused call:
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)
self.linear = nn.LazyLinear(self.vocab_size)
def embedding(self, X):
return self.emb(X.T)
def output_layer(self, hiddens):
return d2l.swapaxes(self.linear(hiddens), 0, 1)Untrained model generates byte soup, but the wiring (tokenizer to model and back) is sound:
'it has upsolars slby al suchownousby� car oTheart wayll slTheart'
Same trainer, same data:
perplexity 105.4, "the time travellerlidence.\n\n'In that the\nface, and\nthere were\ndelight, and\nafter a"
model time (s) val ppl
scratch 20.1 89.3
concise 6.6 105.4
SimpleRNN has no fused GPU kernel; the compiled scratch loop matches it.num_steps tokens.