class Encoder(nnx.Module):
"""The base encoder interface for the encoder-decoder architecture."""
# Later there can be additional arguments (e.g., length excluding padding)
def __call__(self, X, *args):
raise NotImplementedErrorRead one sequence, emit a different one: translation, speech-to-text, summarization. Variable lengths, no positional alignment.
The decoder is a conditional language model: P(y_{t'} \mid y_{<t'}, \mathbf{c}).
The state is the only channel. Swap the encoder and the same shape becomes Whisper (speech), image captioning, or a multimodal front-end.
Run encoder, build state, run decoder. Subclass the classifier for the training loop and loss:
class EncoderDecoder(d2l.Classifier):
"""The base class for the encoder-decoder architecture."""
def __init__(self, encoder, decoder):
super().__init__()
self.encoder = encoder
self.decoder = decoder
def forward(self, enc_X, dec_X, *args):
enc_all_outputs = self.encoder(enc_X, *args)
dec_state = self.decoder.init_state(enc_all_outputs, *args)
# Return decoder output only
return self.decoder(dec_X, dec_state)[0]English-French sentence pairs. Tokenized with one shared byte-level BPE (4k vocab): the languages share the alphabet and many words, and BPE never hits an out-of-vocabulary word.
class MTFraEng(d2l.DataModule):
"""The English-French dataset, tokenized with a shared byte-level BPE."""
def __init__(self, batch_size, num_steps=20, num_train=1024, num_val=128,
vocab_size=4000):
super().__init__()
self.save_hyperparameters()
pairs = self._pairs(self._preprocess(self._download()))
# Train ONE shared byte-level BPE over both languages.
m = min(5000, len(pairs))
self.tokenizer = d2l.BPETokenizer(
vocab_size, pattern=d2l.BPETokenizer.GPT2_PATTERN)
self.tokenizer.train('\n'.join([p[0] for p in pairs[:m]] +
[p[1] for p in pairs[:m]]))
# src_vocab / tgt_vocab both refer to the shared tokenizer.
self.src_vocab = self.tgt_vocab = self.tokenizer
pairs = pairs[:num_train + num_val]
self.src_sents = [p[0] for p in pairs]
self.tgt_sents = [p[1] for p in pairs]
self.arrays = self._encode(pairs)
def _download(self):
d2l.extract(d2l.download(
d2l.DATA_URL + 'fra-eng.zip', self.root,
'94646ad1522d915e7b0f9296181140edcf86a4f5'))
with open(self.root + '/fra-eng/fra.txt', encoding='utf-8') as f:
return f.read()
def _preprocess(self, text):
# Normalize spaces, lowercase, and put a space before punctuation.
text = text.replace(' ', ' ').replace('\xa0', ' ').lower()
no_space = lambda c, prev: c in ',.!?' and prev != ' '
return ''.join([' ' + c if i > 0 and no_space(c, text[i - 1]) else c
for i, c in enumerate(text)])
def _pairs(self, text):
return [ln.split('\t') for ln in text.split('\n') if ln.count('\t') == 1]Pad/truncate to num_steps; append <eos>; targets get a <bos> prefix. Decoder input = target shifted right, label = target shifted left.
@d2l.add_to_class(MTFraEng)
def _encode(self, pairs):
tok, t = self.tokenizer, self.num_steps
def row(sent, is_tgt):
ids = tok.encode(sent)
ids = (ids[:t - 1] + [tok.eos] if len(ids) >= t else
ids + [tok.eos] + [tok.pad] * (t - 1 - len(ids)))
return [tok.bos] + ids if is_tgt else ids
src = d2l.tensor([row(s, False) for s, _ in pairs])
tgt = d2l.tensor([row(t, True) for _, t in pairs])
valid_len = d2l.reduce_sum(d2l.astype(src != tok.pad, d2l.int32), 1)
return src, tgt[:, :-1], valid_len, tgt[:, 1:]
@d2l.add_to_class(MTFraEng)
def build(self, src_sentences, tgt_sentences):
return self._encode([(self._preprocess(s), self._preprocess(t))
for s, t in zip(src_sentences, tgt_sentences)])
@d2l.add_to_class(MTFraEng)
def get_dataloader(self, train):
idx = slice(0, self.num_train) if train else slice(self.num_train, None)
return self.get_tensorloader(self.arrays, train, idx)data = MTFraEng(batch_size=3)
src, dec_in, src_valid_len, label = next(iter(data.train_dataloader()))
print('source:', d2l.astype(src, d2l.int32))
print('decoder input:', d2l.astype(dec_in, d2l.int32))
print('source valid length:', d2l.astype(src_valid_len, d2l.int32))
print('shared vocabulary size:', len(data.tokenizer))source: [[ 524 735 256 4002 4000 4000 4000 4000 4000 4000 4000 4000 4000 4000
4000 4000 4000 4000 4000 4000]
[ 641 1346 111 466 262 4002 4000 4000 4000 4000 4000 4000 4000 4000
4000 4000 4000 4000 4000 4000]
[ 623 262 4002 4000 4000 4000 4000 4000 4000 4000 4000 4000 4000 4000
4000 4000 4000 4000 4000 4000]]
...
[4001 1118 338 1830 371 3663 262 4002 4000 4000 4000 4000 4000 4000
4000 4000 4000 4000 4000 4000]
[4001 740 45 507 262 4002 4000 4000 4000 4000 4000 4000 4000 4000
4000 4000 4000 4000 4000 4000]]
source valid length: [4 6 3]
shared vocabulary size: 4003
Which is why a small num_steps and a masked loss pay off:
Embed source tokens, run a multilayer GRU, return per-step states and the final state (the context \mathbf{c}):
class Seq2SeqEncoder(d2l.Encoder):
"""The RNN encoder for sequence-to-sequence learning."""
def __init__(self, vocab_size, embed_size, num_hiddens, num_layers,
dropout=0, rngs=None):
rngs = nnx.Rngs(params=0, dropout=1, carry=2) if rngs is None else rngs
self.embedding = nnx.Embed(vocab_size, embed_size, rngs=rngs)
self.rnn = d2l.GRU(embed_size, num_hiddens, num_layers, dropout,
rngs=rngs)
def __call__(self, X, *args):
embs = self.embedding(d2l.astype(d2l.transpose(X), d2l.int64))
outputs, state = self.rnn(embs)
# outputs: (num_steps, batch_size, num_hiddens)
# state: (num_layers, batch_size, num_hiddens)
return outputs, statevocab_size, embed_size, num_hiddens, num_layers = 10, 8, 16, 2
batch_size, num_steps = 4, 9
encoder = Seq2SeqEncoder(vocab_size, embed_size, num_hiddens, num_layers)
X = d2l.zeros((batch_size, num_steps))
enc_outputs, enc_state = encoder(X)
d2l.check_shape(enc_outputs, (num_steps, batch_size, num_hiddens))\mathbf{s}_{t'} = g(y_{t'-1}, \mathbf{c}, \mathbf{s}_{t'-1})
Initialize from the encoder state; concatenate \mathbf{c} onto the input at every step; project to vocab logits:
class Seq2SeqDecoder(d2l.Decoder):
"""The RNN decoder for sequence-to-sequence learning."""
def __init__(self, vocab_size, embed_size, num_hiddens, num_layers,
dropout=0, rngs=None):
rngs = nnx.Rngs(params=2, dropout=3, carry=4) if rngs is None else rngs
self.embedding = nnx.Embed(vocab_size, embed_size, rngs=rngs)
self.rnn = d2l.GRU(embed_size + num_hiddens, num_hiddens,
num_layers, dropout, rngs=rngs)
self.dense = nnx.Linear(num_hiddens, vocab_size, rngs=rngs)
def init_state(self, enc_all_outputs, *args):
return enc_all_outputs
def __call__(self, X, state):
embs = self.embedding(d2l.astype(d2l.transpose(X), d2l.int64))
enc_output, hidden_state = state
context = enc_output[-1]
context = jnp.tile(context, (embs.shape[0], 1, 1))
embs_and_context = d2l.concat((embs, context), -1)
outputs, hidden_state = self.rnn(embs_and_context, hidden_state)
outputs = d2l.swapaxes(self.dense(outputs), 0, 1)
return outputs, [enc_output, hidden_state]class Seq2Seq(d2l.EncoderDecoder):
"""The RNN encoder-decoder for sequence-to-sequence learning."""
def __init__(self, encoder, decoder, tgt_pad, lr):
super().__init__(encoder, decoder)
self.tgt_pad, self.lr = tgt_pad, lr
def validation_step(self, batch):
return self.loss(self(*batch[:-1]), batch[-1])
def configure_optimizers(self):
return optax.adam(learning_rate=self.lr)<pad> predictions must not count. Mask them out and average over real tokens only:
\mathcal{L} = \frac{\sum_{b,t} \mathbf{1}\{y_{b,t} \ne \texttt{<pad>}\}\, \ell(\hat{\mathbf{y}}_{b,t}, y_{b,t})}{\sum_{b,t} \mathbf{1}\{y_{b,t} \ne \texttt{<pad>}\}}
2-layer GRU, width 256, dropout 0.2, Adam lr 0.005, clip 1, 30 epochs:
data = d2l.MTFraEng(batch_size=128)
embed_size, num_hiddens, num_layers, dropout = 256, 256, 2, 0.2
encoder = Seq2SeqEncoder(
len(data.tokenizer), embed_size, num_hiddens, num_layers, dropout)
decoder = Seq2SeqDecoder(
len(data.tokenizer), embed_size, num_hiddens, num_layers, dropout)
model = Seq2Seq(encoder, decoder, tgt_pad=data.tokenizer.pad, lr=0.005)
trainer = d2l.Trainer(max_epochs=30, gradient_clip_val=1, num_gpus=1)
trainer.fit(model, data)predict_step decodes a whole minibatch at once (feed <bos>, take argmax, feed back, stop at <eos>):
@d2l.add_to_class(d2l.EncoderDecoder)
def predict_step(self, batch, num_steps, save_attention_weights=False):
model = nnx.view(self, deterministic=True, use_running_average=True,
raise_if_not_found=False)
src, tgt, src_valid_len, _ = batch
enc_all_outputs = model.encoder(src, src_valid_len)
enc_attention_weights = (getattr(model.encoder, 'attention_weights', [])
if save_attention_weights else [])
dec_state = model.decoder.init_state(enc_all_outputs, src_valid_len)
outputs, attention_weights = [d2l.expand_dims(tgt[:, 0], 1), ], []
for _ in range(num_steps):
Y, dec_state = model.decoder(outputs[-1], dec_state)
outputs.append(d2l.argmax(Y, 2))
if save_attention_weights:
attention_weights.append(model.decoder.attention_weights)
return d2l.concat(outputs[1:], 1), (attention_weights,
enc_attention_weights)engs = ['i lost .', "i'm calm .", "i'm home ."]
fras = ["j'ai perdu .", 'je suis calme .', 'je suis chez moi .']
preds, _ = model.predict_step(data.build(engs, fras), data.num_steps)
translations = [to_text(p, data.tokenizer) for p in preds]
for eng, out in zip(engs, translations):
print(f'{eng} => {out!r}')i lost . => "j'ai perdu ."
i'm calm . => 'je suis calme .'
i'm home . => 'je suis timide .'
The 8.7 toolkit plugs straight in: wrap the decoder in a step_fn (source fixed, target prefix varies), then d2l.beam_search. Greedy is beam size 1; score it vs k = 2, 5 with chrF:
def make_step_fn(model, src): # src: padded source token ids
m = nnx.view(model, deterministic=True, use_running_average=True,
raise_if_not_found=False)
enc = m.encoder(d2l.tensor([src]))
def step_fn(tgt_ids):
state = m.decoder.init_state(enc)
Y, _ = m.decoder(d2l.tensor([tgt_ids]), state)
return d2l.numpy(Y)[0, -1]
return step_fntok = data.tokenizer
for eng, fra in zip(engs, fras):
src = [int(i) for i in d2l.numpy(data.build([eng], [eng])[0][0])]
step_fn = make_step_fn(model, src)
row = []
for k in (1, 2, 5):
ids = d2l.beam_search(step_fn, [tok.bos], data.num_steps,
beam_size=k, eos_id=tok.eos)[0][1]
out = tok.decode([i for i in ids[1:] if i != tok.eos])
row.append(f'k={k}: {chrf(out, data._preprocess(fra)):.2f}')
print(f'{eng:<12} ' + ' '.join(row))i lost . k=1: 1.00 k=2: 1.00 k=5: 1.00
i'm calm . k=1: 1.00 k=2: 1.00 k=5: 1.00
i'm home . k=1: 0.34 k=2: 0.34 k=5: 0.34
Character n-gram F-score (Popovic, 2015): tokenization-free, partial credit for near-misses. The WMT default; BLEU demoted to a remark.
def chrf(pred, label, n=6, beta=2):
"""chrF (Popovic, 2015): character n-gram F-score."""
def ngrams(s, k):
s = s.replace(' ', '')
return collections.Counter(s[i:i + k] for i in range(len(s) - k + 1))
prec = rec = 0.0
for k in range(1, n + 1):
p, r = ngrams(pred, k), ngrams(label, k)
overlap = sum((p & r).values())
if sum(p.values()) and sum(r.values()):
prec += overlap / sum(p.values())
rec += overlap / sum(r.values())
prec, rec = prec / n, rec / n
if prec + rec == 0:
return 0.0
return (1 + beta**2) * prec * rec / (beta**2 * prec + rec)Everything about the source crosses one vector \mathbf{c}. chrF falls as the source grows:
import numpy as np
buckets = collections.defaultdict(list)
for s, t, p in zip(h_src, h_tgt, h_preds):
L = len(data.tokenizer.encode(s))
buckets[L].append(chrf(to_text(p, data.tokenizer), t))
xs = sorted(L for L in buckets if len(buckets[L]) >= 8 and L <= data.num_steps)
ys = [np.mean(buckets[L]) for L in xs]
d2l.set_figsize()
d2l.plt.plot(xs, ys, 'o-')
d2l.plt.xlabel('source length (tokens)')
d2l.plt.ylabel('mean chrF')
d2l.plt.grid(True);Two escapes, both already taken: attention (look back at all encoder states — where the attention chapters began) or a better recurrent state (the state space models you built).