Encoder-Decoder Transformer in C++ (LibTorch)
The Mental Model: Read, Then Write
An encoder-decoder Transformer is the original Transformer shape from "Attention Is All You Need." It is built for sequence-to-sequence tasks: translation, summarization, question answering, and any setting where the input sequence and output sequence are different streams.
Encoder
Reads the complete source sequence once. Its self-attention is bidirectional, so every source token can use left and right context. The output is called memory.
Decoder
Writes the target sequence. Its self-attention is causal, so it cannot see future target tokens, but its cross-attention can read the full encoder memory.
The key new component compared with GPT and encoder-only models is cross-attention: queries come from the decoder stream, while keys and values come from the encoder memory.
Architecture Overview
The Three Mask Rules
| Place | Mask | Why |
|---|---|---|
| Encoder self-attention | Source padding mask only | Source tokens can read both directions, but should not read [PAD]. |
| Decoder self-attention | Causal mask + target padding mask | Target position i cannot peek at future target positions. |
| Decoder cross-attention | Source padding mask only | The decoder may read every real source token; there is no causal mask over source memory. |
1. Cross-Attention Head
Cross-attention looks like self-attention except Q comes from the decoder and K,V come from the encoder memory. That changes the score shape from square (T, T) to rectangular (T_tgt, T_src).
#include <torch/torch.h>
#include <cmath>
#include <memory>
#include <string>
#include <vector>
class CrossAttentionHead : public torch::nn::Module {
int n_embd, head_size;
torch::nn::Linear Q{nullptr}, K{nullptr}, V{nullptr};
public:
CrossAttentionHead(int n_embd, int head_size)
: n_embd(n_embd), head_size(head_size) {
Q = register_module("Q", torch::nn::Linear(
torch::nn::LinearOptions(n_embd, head_size).bias(false)));
K = register_module("K", torch::nn::Linear(
torch::nn::LinearOptions(n_embd, head_size).bias(false)));
V = register_module("V", torch::nn::Linear(
torch::nn::LinearOptions(n_embd, head_size).bias(false)));
}
torch::Tensor forward(torch::Tensor decoder_x,
torch::Tensor memory,
torch::Tensor src_mask = torch::Tensor()) {
auto q = Q(decoder_x); // (B, T_tgt, head_size)
auto k = K(memory); // (B, T_src, head_size)
auto v = V(memory); // (B, T_src, head_size)
auto scores = torch::matmul(q, k.transpose(-2, -1));
scores = scores / std::sqrt((double)head_size); // (B, T_tgt, T_src)
if (src_mask.defined()) {
auto keep = src_mask.unsqueeze(1).to(torch::kBool); // (B, 1, T_src)
scores = scores.masked_fill(keep.logical_not(), -1e9);
}
auto weights = torch::softmax(scores, -1);
return torch::matmul(weights, v); // (B, T_tgt, head_size)
}
};
2. EncoderBlock and DecoderBlock
The encoder block is the same as the encoder-only post: bidirectional self-attention plus a feed-forward network. The decoder block has three sublayers: masked self-attention, cross-attention, and feed-forward.
class MultiHeadCrossAttention : public torch::nn::Module {
int num_heads, head_size, n_embd;
std::vector<std::shared_ptr<CrossAttentionHead>> heads;
torch::nn::Linear projection{nullptr};
public:
MultiHeadCrossAttention(int num_heads, int head_size, int n_embd)
: num_heads(num_heads), head_size(head_size), n_embd(n_embd) {
for (int i = 0; i < num_heads; i++) {
auto head = std::make_shared<CrossAttentionHead>(n_embd, head_size);
register_module("cross_head_" + std::to_string(i), head);
heads.push_back(head);
}
projection = register_module("projection",
torch::nn::Linear(head_size * num_heads, n_embd));
}
torch::Tensor forward(torch::Tensor decoder_x, torch::Tensor memory,
torch::Tensor src_mask = torch::Tensor()) {
std::vector<torch::Tensor> outputs;
for (int i = 0; i < num_heads; i++) {
outputs.push_back(heads[i]->forward(decoder_x, memory, src_mask));
}
return projection(torch::cat(outputs, -1));
}
};
class Seq2SeqEncoderBlock : public torch::nn::Module {
std::shared_ptr<MultiHeadSelfAttention> self_attn;
std::shared_ptr<FeedForward> ff;
torch::nn::LayerNorm ln1{nullptr}, ln2{nullptr};
public:
Seq2SeqEncoderBlock(int num_heads, int head_size, int n_embd, int d_ff) {
self_attn = std::make_shared<MultiHeadSelfAttention>(
num_heads, head_size, n_embd);
ff = std::make_shared<FeedForward>(n_embd, d_ff);
register_module("self_attn", self_attn);
register_module("ff", ff);
ln1 = register_module("ln1",
torch::nn::LayerNorm(torch::nn::LayerNormOptions({n_embd})));
ln2 = register_module("ln2",
torch::nn::LayerNorm(torch::nn::LayerNormOptions({n_embd})));
}
torch::Tensor forward(torch::Tensor x,
torch::Tensor src_mask = torch::Tensor()) {
x = x + self_attn->forward(ln1(x), src_mask);
x = x + ff->forward(ln2(x));
return x;
}
};
class Seq2SeqDecoderBlock : public torch::nn::Module {
std::shared_ptr<MaskedMultiHeadSelfAttention> masked_self_attn;
std::shared_ptr<MultiHeadCrossAttention> cross_attn;
std::shared_ptr<FeedForward> ff;
torch::nn::LayerNorm ln1{nullptr}, ln2{nullptr}, ln3{nullptr};
public:
Seq2SeqDecoderBlock(int num_heads, int head_size, int n_embd, int d_ff) {
masked_self_attn = std::make_shared<MaskedMultiHeadSelfAttention>(
num_heads, head_size, n_embd);
cross_attn = std::make_shared<MultiHeadCrossAttention>(
num_heads, head_size, n_embd);
ff = std::make_shared<FeedForward>(n_embd, d_ff);
register_module("masked_self_attn", masked_self_attn);
register_module("cross_attn", cross_attn);
register_module("ff", ff);
ln1 = register_module("ln1",
torch::nn::LayerNorm(torch::nn::LayerNormOptions({n_embd})));
ln2 = register_module("ln2",
torch::nn::LayerNorm(torch::nn::LayerNormOptions({n_embd})));
ln3 = register_module("ln3",
torch::nn::LayerNorm(torch::nn::LayerNormOptions({n_embd})));
}
torch::Tensor forward(torch::Tensor x, torch::Tensor memory,
torch::Tensor tgt_mask = torch::Tensor(),
torch::Tensor src_mask = torch::Tensor()) {
x = x + masked_self_attn->forward(ln1(x), tgt_mask);
x = x + cross_attn->forward(ln2(x), memory, src_mask);
x = x + ff->forward(ln3(x));
return x;
}
};
MaskedMultiHeadSelfAttention come from? It is the GPT-style causal multi-head attention from the decoder-only post: run causal attention heads, concatenate along the last dimension, and project back to n_embd. The only extra mask it needs is target padding.
3. Full Encoder-Decoder Model
The full model has two embedding tables, two stacks, and one generator head. During training, the target input is shifted right: if the desired output is <bos> I am here <eos>, the decoder input is <bos> I am here and the labels are I am here <eos>.
class EncoderDecoderTransformer : public torch::nn::Module {
int src_vocab, tgt_vocab, n_embd, num_heads, head_size;
int d_ff, block_size, num_layers;
std::shared_ptr<PositionalEncoding> pe;
std::vector<std::shared_ptr<Seq2SeqEncoderBlock>> encoder_blocks;
std::vector<std::shared_ptr<Seq2SeqDecoderBlock>> decoder_blocks;
torch::nn::Embedding src_emb{nullptr}, tgt_emb{nullptr};
torch::nn::LayerNorm final_ln{nullptr};
torch::nn::Linear generator{nullptr};
public:
EncoderDecoderTransformer(int src_vocab, int tgt_vocab, int n_embd,
int num_heads, int head_size, int d_ff,
int block_size, int num_layers)
: src_vocab(src_vocab), tgt_vocab(tgt_vocab), n_embd(n_embd),
num_heads(num_heads), head_size(head_size), d_ff(d_ff),
block_size(block_size), num_layers(num_layers) {
pe = std::make_shared<PositionalEncoding>(n_embd);
src_emb = register_module("src_emb",
torch::nn::Embedding(src_vocab, n_embd));
tgt_emb = register_module("tgt_emb",
torch::nn::Embedding(tgt_vocab, n_embd));
final_ln = register_module("final_ln",
torch::nn::LayerNorm(torch::nn::LayerNormOptions({n_embd})));
generator = register_module("generator",
torch::nn::Linear(n_embd, tgt_vocab));
for (int layer = 0; layer < num_layers; layer++) {
auto enc = std::make_shared<Seq2SeqEncoderBlock>(
num_heads, head_size, n_embd, d_ff);
auto dec = std::make_shared<Seq2SeqDecoderBlock>(
num_heads, head_size, n_embd, d_ff);
register_module("encoder_block_" + std::to_string(layer), enc);
register_module("decoder_block_" + std::to_string(layer), dec);
encoder_blocks.push_back(enc);
decoder_blocks.push_back(dec);
}
}
torch::Tensor encode(torch::Tensor src_ids,
torch::Tensor src_mask = torch::Tensor()) {
auto x = src_emb(src_ids);
x = x + pe->forward(x);
for (int layer = 0; layer < num_layers; layer++) {
x = encoder_blocks[layer]->forward(x, src_mask);
}
return x; // memory: (B, T_src, n_embd)
}
torch::Tensor decode(torch::Tensor tgt_ids, torch::Tensor memory,
torch::Tensor tgt_mask = torch::Tensor(),
torch::Tensor src_mask = torch::Tensor()) {
auto x = tgt_emb(tgt_ids);
x = x + pe->forward(x);
for (int layer = 0; layer < num_layers; layer++) {
x = decoder_blocks[layer]->forward(x, memory, tgt_mask, src_mask);
}
return x; // (B, T_tgt, n_embd)
}
torch::Tensor forward(torch::Tensor src_ids, torch::Tensor tgt_in,
torch::Tensor src_mask = torch::Tensor(),
torch::Tensor tgt_mask = torch::Tensor()) {
auto memory = encode(src_ids, src_mask);
auto decoded = decode(tgt_in, memory, tgt_mask, src_mask);
return project(decoded); // (B, T_tgt, tgt_vocab)
}
torch::Tensor project(torch::Tensor decoded) {
decoded = final_ln(decoded);
return generator(decoded);
}
};
Interactive Animation: Encoder-Decoder Flow
The encoder runs once. The decoder consumes shifted target tokens, masks its own future, then cross-attends to the fixed encoder memory.
Shape Trace
Example: B=4, T_src=30, T_tgt=18, n_embd=128, tgt_vocab=8000.
| Step | Expression | Output shape | Meaning |
|---|---|---|---|
| 0 | src_ids | (4, 30) | Source token IDs. |
| 1 | src_emb(src_ids) + PE | (4, 30, 128) | Source vectors. |
| 2 | encode(...) | (4, 30, 128) | Encoder memory. |
| 3 | tgt_in | (4, 18) | Shifted-right target input. |
| 4 | tgt_emb(tgt_in) + PE | (4, 18, 128) | Decoder input vectors. |
| 5 | masked self-attn | (4, 18, 128) | Target-side context without future tokens. |
| 6 | cross-attn | (4, 18, 128) | Each target position reads source memory. |
| 7 | generator(decoded) | (4, 18, 8000) | Vocabulary logits for each target position. |
Training and Inference
Training with teacher forcing
Training is parallel because we already know the true target sequence. We feed the shifted target into the decoder and compare every output position against the next target token.
// target_ids: [BOS, y1, y2, y3, EOS]
auto tgt_in = target_ids.index({
torch::indexing::Slice(),
torch::indexing::Slice(0, -1)
});
auto labels = target_ids.index({
torch::indexing::Slice(),
torch::indexing::Slice(1, torch::indexing::None)
});
auto logits = model.forward(src_ids, tgt_in, src_mask, tgt_mask);
auto loss = torch::nn::functional::cross_entropy(
logits.view({-1, tgt_vocab}),
labels.reshape({-1})
);
Greedy inference loop
Inference is autoregressive. Encode the source once, then repeatedly decode the generated prefix and append one new token.
torch::Tensor greedy_generate(EncoderDecoderTransformer& model,
torch::Tensor src_ids,
torch::Tensor src_mask,
int bos_id, int eos_id, int max_new_tokens) {
auto memory = model.encode(src_ids, src_mask); // run encoder once
auto generated = torch::full({src_ids.size(0), 1}, bos_id,
torch::TensorOptions().dtype(torch::kLong).device(src_ids.device()));
for (int step = 0; step < max_new_tokens; step++) {
auto decoded = model.decode(generated, memory, torch::Tensor(), src_mask);
auto logits = model.project(decoded);
auto last_logits = logits.select(1, logits.size(1) - 1);
auto next = last_logits.argmax(-1, true); // (B, 1)
generated = torch::cat({generated, next}, 1);
// Simple batch-size-1 stop condition for teaching.
if (next.item<int64_t>() == eos_id) break;
}
return generated;
}
is_finished vector instead of using a batch-size-1 item() stop condition. For speed, cache decoder K/V states instead of recomputing the whole prefix every step.
Common Pitfalls
| Pitfall | Symptom | Fix |
|---|---|---|
| Not shifting targets | The model is trained to copy the current token instead of predicting the next token. | Use tgt_in = target[:-1] and labels = target[1:]. |
| Causal mask in cross-attention | The decoder cannot read all source tokens. | Only apply source padding mask in cross-attention. |
| Re-encoding inside every decode step | Inference is much slower than necessary. | Call encode(src) once and reuse memory. |
| Ignoring padding masks | Attention learns from padding positions. | Pass source masks to encoder and cross-attention; pass target masks to decoder self-attention. |
What Comes Next
You now have the three major Transformer families in C++: decoder-only GPT, encoder-only readers, and encoder-decoder sequence-to-sequence models. From here, the next practical upgrades are beam search, label smoothing, padding-aware losses, and key-value caching.