How d2l-zh Explains the Encoder-Decoder Architecture for Sequence-to-Sequence Tasks
The d2l-zh textbook implements the encoder-decoder architecture as an abstract interface with three components—Encoder, Decoder, and EncoderDecoder—concretely realized using GRU-based recurrent networks for machine translation tasks.
The d2l-zh repository provides a comprehensive educational implementation of the encoder-decoder framework for sequence-to-sequence learning. This architecture serves as the foundation for modern neural machine translation and other variable-length sequence mapping tasks.
Motivation for Variable-Length Sequence Mapping
Machine translation exemplifies the core challenge that the encoder-decoder architecture solves: converting an English sentence (a sequence of tokens) into a French sentence of potentially different length. The textbook proposes a two-stage process where the model encodes the source sequence into a fixed-size context vector (the encoding state) and then decodes that context into the target sequence token-by-token.
Core Components of the Encoder-Decoder Interface
The architecture in chapter_recurrent-modern/encoder-decoder.md defines three abstract interfaces implemented across MXNet, PyTorch, TensorFlow, and Paddle.
The Encoder Abstraction
The Encoder receives a variable-length source tensor X and returns a tuple (output, state). The state—typically the final hidden vector of a recurrent network—becomes the initial state for the decoder.
class Encoder(nn.Block):
def __init__(self, **kwargs):
super(Encoder, self).__init__(**kwargs)
def forward(self, X, *args):
raise NotImplementedError
The Decoder Abstraction
The Decoder defines two critical methods:
init_state(enc_outputs, …)transforms the encoder's output into the decoder's initial hidden state.forward(X, state)consumes the previous target token(s)Xtogether with the current hidden state, producing the next token distribution and an updated state.
class Decoder(nn.Block):
def __init__(self, **kwargs):
super(Decoder, self).__init__(**kwargs)
def init_state(self, enc_outputs, *args):
raise NotImplementedError
def forward(self, X, state):
raise NotImplementedError
The EncoderDecoder Wrapper
The EncoderDecoder class orchestrates the two components. Its forward(enc_X, dec_X, …) method runs the encoder, initializes the decoder state, and executes the decoder on the target input.
class EncoderDecoder(nn.Block):
def __init__(self, encoder, decoder, **kwargs):
super(EncoderDecoder, self).__init__(**kwargs)
self.encoder = encoder
self.decoder = decoder
def forward(self, enc_X, dec_X, *args):
enc_outputs = self.encoder(enc_X, *args)
dec_state = self.decoder.init_state(enc_outputs, *args)
return self.decoder(dec_X, dec_state)
Concrete Implementation with GRU Networks
The chapter_recurrent-modern/seq2seq.md file provides a concrete realization using Gated Recurrent Units (GRU) for machine translation.
Seq2SeqEncoder Implementation
The encoder embeds input tokens and processes them through a multi-layer GRU. It returns the final hidden state as the context vector.
class Seq2SeqEncoder(d2l.Encoder):
def __init__(self, vocab_size, embed_size, num_hiddens,
num_layers, dropout=0, **kwargs):
super(Seq2SeqEncoder, self).__init__(**kwargs)
self.embedding = nn.Embedding(vocab_size, embed_size)
self.rnn = rnn.GRU(num_hiddens, num_layers, dropout=dropout)
def forward(self, X, *args):
X = self.embedding(X) # (batch, steps, embed)
X = X.swapaxes(0, 1) # (steps, batch, embed)
state = self.rnn.begin_state(batch_size=X.shape[1],
ctx=X.ctx)
output, state = self.rnn(X, state)
return output, state
Seq2SeqDecoder Implementation
The decoder uses the encoder's final hidden state as its initial state. It implements teacher forcing during training by feeding ground-truth tokens as inputs while generating predictions autoregressively during inference.
class Seq2SeqDecoder(d2l.Decoder):
def __init__(self, vocab_size, embed_size, num_hiddens,
num_layers, dropout=0, **kwargs):
super(Seq2SeqDecoder, self).__init__(**kwargs)
self.embedding = nn.Embedding(vocab_size, embed_size)
self.rnn = rnn.GRU(num_hiddens, num_layers, dropout=dropout)
self.dense = nn.Dense(vocab_size, flatten=False)
def init_state(self, enc_outputs, *args):
# enc_outputs[1] is the final hidden state of the encoder
return enc_outputs[1]
def forward(self, X, state):
X = self.embedding(X).swapaxes(0, 1) # (steps, batch, embed)
context = state[0][-1] # (batch, num_hiddens)
context = np.broadcast_to(context,
(X.shape[0], context.shape[0], context.shape[1]))
X_and_context = d2l.concat((X, context), dim=2)
output, state = self.rnn(X_and_context, state)
output = self.dense(output).swapaxes(0, 1) # (batch, steps, vocab)
return output, state
Training and Inference Procedures
The training loop in chapter_recurrent-modern/seq2seq.md feeds the source sequence X and a teacher-forced target input dec_input (ground-truth shifted right with a <bos> token). The loss function uses masked softmax cross-entropy to ignore padding tokens.
def train_seq2seq(net, data_iter, lr, num_epochs, tgt_vocab, device):
# Training implementation with teacher forcing
# Loss ignores padding tokens (<pad>)
pass
During inference, the predict_seq2seq function generates translations by repeatedly feeding the decoder's own most-likely token as the next input until an <eos> token appears or the maximum step count is reached.
Summary
- The encoder-decoder architecture in d2l-zh provides a generic framework for mapping variable-length input sequences to variable-length output sequences.
- The implementation defines three abstract interfaces—Encoder, Decoder, and EncoderDecoder—that are framework-agnostic and implemented for MXNet, PyTorch, TensorFlow, and Paddle.
- Concrete seq2seq models use GRU-based recurrent networks where the encoder compresses the source sentence into a context vector and the decoder generates the target sentence token-by-token using teacher forcing during training.
- The complete training and inference pipeline, including masked loss computation and autoregressive prediction, is implemented in
chapter_recurrent-modern/seq2seq.md.
Frequently Asked Questions
What is the primary purpose of the encoder-decoder architecture in d2l-zh?
The primary purpose is to handle sequence-to-sequence tasks where input and output lengths differ, such as machine translation. The architecture separates the problem into two stages: an encoder that compresses the variable-length source sequence into a fixed-size context vector, and a decoder that expands this context into a variable-length target sequence.
How does the EncoderDecoder class connect the encoder and decoder?
The EncoderDecoder class acts as a wrapper that stores both components. Its forward method first calls self.encoder(enc_X) to obtain outputs and state, then initializes the decoder's state via self.decoder.init_state(enc_outputs), and finally returns self.decoder(dec_X, dec_state). This orchestration ensures the decoder receives the proper initial context from the encoder.
Why does the decoder use teacher forcing during training?
Teacher forcing improves training stability and convergence speed by feeding the ground-truth previous token (from dec_input) rather than the model's own prediction into the decoder at each step. This prevents error accumulation during the early training stages when the model's predictions are unreliable. During inference, the model switches to autoregressive generation, feeding its own predictions back as inputs.
Where can I find the complete implementation code?
The abstract interface definitions reside in chapter_recurrent-modern/encoder-decoder.md, while the concrete GRU-based seq2seq implementation, training loop, and inference functions are located in chapter_recurrent-modern/seq2seq.md. Both files provide framework-agnostic explanations alongside MXNet, PyTorch, TensorFlow, and Paddle implementations.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →