← All Posts

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

Source token IDs src embedding + PE EncoderBlock x N Bidirectional MHA FeedForward memory = encoder output Shifted target IDs tgt embedding + PE DecoderBlock x N Masked MHA Cross-Attention FeedForward vocab logits
Pre-norm implementation. The original paper's diagram places LayerNorm after each residual addition. This C++ series uses pre-norm for consistency with the GPT implementation: normalize before attention or FFN, then add the residual.

The Three Mask Rules

PlaceMaskWhy
Encoder self-attentionSource padding mask onlySource tokens can read both directions, but should not read [PAD].
Decoder self-attentionCausal mask + target padding maskTarget position i cannot peek at future target positions.
Decoder cross-attentionSource padding mask onlyThe 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;
    }
};
Where does 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.

Ready
Encoder row Decoder row Source IDs "I like cats" src_emb + PE (B, T_src, d) EncoderBlock x N bidirectional attention memory: (B, T_src, d) Shifted target IDs "<bos> j'aime" tgt_emb + PE (B, T_tgt, d) Masked self-attention decoder cannot see future target tokens Cross-attention Q from decoder, K/V from memory FFN + generator (B, T_tgt, tgt_vocab)
Click Step to start with source token IDs.

Shape Trace

Example: B=4, T_src=30, T_tgt=18, n_embd=128, tgt_vocab=8000.

StepExpressionOutput shapeMeaning
0src_ids(4, 30)Source token IDs.
1src_emb(src_ids) + PE(4, 30, 128)Source vectors.
2encode(...)(4, 30, 128)Encoder memory.
3tgt_in(4, 18)Shifted-right target input.
4tgt_emb(tgt_in) + PE(4, 18, 128)Decoder input vectors.
5masked self-attn(4, 18, 128)Target-side context without future tokens.
6cross-attn(4, 18, 128)Each target position reads source memory.
7generator(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;
}
Production note. For real batched inference, track an 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

PitfallSymptomFix
Not shifting targetsThe 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-attentionThe decoder cannot read all source tokens.Only apply source padding mask in cross-attention.
Re-encoding inside every decode stepInference is much slower than necessary.Call encode(src) once and reuse memory.
Ignoring padding masksAttention 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.

Previous full model Encoder-Only Transformer