%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, stateCheck the 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 98.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.38 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 traveller. I had the same. I was not beenred, and the same. I was
not beenred, and the same. I was not beenred, and the same. I was not
beenred, and the same. I was not been'
Greedy is fluent, then circles. Sampling breaks the loop at the price of stranger choices:
the time travellerions move it was sensle of the Lil of an with the L Space
this curesivity of the little peoplepeove
the time traveller of the Time Traveller. 'I wonderful of the great
diffto the Morlocks, and I could see the panelight
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>oratoryV part eyes� blackhed le mind�pp downultJ Man black seemed were8'
Same trainer, same data:
perplexity 101.1, "the time traveller.\nThe Time Traveller'ser--
aitsie.\n\n'Incendly, and\nwere was a f"
model time (s) val ppl
scratch 16.9 98.3
concise 6.5 101.1
SimpleRNN has no fused GPU kernel; the compiled scratch loop matches it.num_steps tokens.