Encoder-Decoder Models for Sequence Transduction

Sequence transduction

Read one sequence, emit a different one: translation, speech-to-text, summarization. Variable lengths, no positional alignment.

  • Encoder compresses the source into a fixed-shape state.
  • Decoder expands that state into the target, one token at a time.

The decoder is a conditional language model: P(y_{t'} \mid y_{<t'}, \mathbf{c}).

The encoder-decoder abstraction

The state is the only channel. Swap the encoder and the same shape becomes Whisper (speech), image captioning, or a multimodal front-end.

class Encoder(nn.Module):
    """The base encoder interface for the encoder-decoder architecture."""
    def __init__(self):
        super().__init__()

    # Later there can be additional arguments (e.g., length excluding padding)
    def forward(self, X, *args):
        raise NotImplementedError
class Decoder(nn.Module):
    """The base decoder interface for the encoder-decoder architecture."""
    def __init__(self):
        super().__init__()

    def init_state(self, enc_all_outputs, *args):
        raise NotImplementedError

    def forward(self, X, state):
        raise NotImplementedError

Wiring them together

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]

The MT dataset

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]

Padding and teacher forcing

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: tensor([[ 105, 1892,  304,  256, 4002, 4000, 4000, 4000, 4000, 4000, 4000, 4000,
         4000, 4000, 4000, 4000, 4000, 4000, 4000, 4000],
        [ 105,  297, 3084,  256, 4002, 4000, 4000, 4000, 4000, 4000, 4000, 4000,
         4000, 4000, 4000, 4000, 4000, 4000, 4000, 4000],
        [ 105,  297, 1572,  256, 4002, 4000, 4000, 4000, 4000, 4000, 4000, 4000,
         4000, 4000, 4000, 4000, 4000, 4000, 4000, 4000]], dtype=torch.int32)
...
        [4001,  270,  296, 3569, 2023,  256, 4002, 4000, 4000, 4000, 4000, 4000,
         4000, 4000, 4000, 4000, 4000, 4000, 4000, 4000],
        [4001,  270,  296, 2381,  256, 4002, 4000, 4000, 4000, 4000, 4000, 4000,
         4000, 4000, 4000, 4000, 4000, 4000, 4000, 4000]], dtype=torch.int32)
source valid length: tensor([5, 5, 5], dtype=torch.int32)
shared vocabulary size: 4003

Most sentences are short

Which is why a small num_steps and a masked loss pay off:

src_ids = [data.tokenizer.encode(s) for s in data.src_sents]
tgt_ids = [data.tokenizer.encode(s) for s in data.tgt_sents]
show_list_len_pair_hist(['source', 'target'], '# tokens per sentence',
                        'count', src_ids, tgt_ids);

Encoder: embed + GRU

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):
        super().__init__()
        self.embedding = nn.Embedding(vocab_size, embed_size)
        self.rnn = d2l.GRU(embed_size, num_hiddens, num_layers, dropout)
        self.apply(init_seq2seq)

    def forward(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, state
vocab_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))

Decoder: context-conditioned GRU

\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):
        super().__init__()
        self.embedding = nn.Embedding(vocab_size, embed_size)
        self.rnn = d2l.GRU(embed_size + num_hiddens, num_hiddens,
                           num_layers, dropout)
        self.dense = nn.LazyLinear(vocab_size)
        self.apply(init_seq2seq)

    def init_state(self, enc_all_outputs, *args):
        return enc_all_outputs

    def forward(self, X, state):
        embs = self.embedding(d2l.astype(d2l.transpose(X), d2l.int64))
        enc_output, hidden_state = state
        context = enc_output[-1]
        context = context.repeat(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]

The unrolled model

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.save_hyperparameters()

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

    def configure_optimizers(self):
        return torch.optim.Adam(self.parameters(), lr=self.lr)

Masked loss

<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>}\}}

def loss(self, Y_hat, Y):
    l = super(Seq2Seq, self).loss(Y_hat, Y, averaged=False)
    mask = d2l.astype(d2l.reshape(Y, (-1,)) != self.tgt_pad, d2l.float32)
    return d2l.reduce_sum(l * mask) / d2l.reduce_sum(mask)

Training

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)

Greedy translation

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, device, num_steps,
                 save_attention_weights=False):
    self.eval()
    batch = [d2l.to(a, device) for a in batch]
    src, tgt, src_valid_len, _ = batch
    enc_all_outputs = self.encoder(src, src_valid_len)
    dec_state = self.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 = self.decoder(outputs[-1], dec_state)
        outputs.append(d2l.argmax(Y, 2))
        if save_attention_weights:
            attention_weights.append(self.decoder.attention_weights)
    return d2l.concat(outputs[1:], 1), 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), d2l.try_gpu(),
                              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 chez moi .'

Evaluation: chrF

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)
for eng, fra, out in zip(engs, fras, translations):
    print(f'{eng} => {out!r}, chrF {chrf(out, data._preprocess(fra)):.3f}')
i lost . => "j'ai perdu .", chrF 1.000
i'm calm . => 'je suis calme .', chrF 1.000
i'm home . => 'je suis chez moi .', chrF 1.000

The fixed-vector bottleneck

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).

Recap

  • Encoder-decoder: source \to fixed state \to target; decoder is a conditional LM.
  • Shared BPE vocab, teacher forcing, masked cross-entropy.
  • Decode with greedy or beam search (8.7’s toolkit, unchanged).
  • chrF over BLEU for lexical scoring.
  • The single-vector bottleneck degrades with length and motivates the rest of the book.