Encoder-Decoder Transformers: Sequence to Sequence, Variable Output Length
vision
transformers
A decoder that emits digits until it emits an end token, trained by teacher forcing and scored on whether the whole sequence is right.
Author
Rosh Beed
Published
June 25, 2026
This model reads a whole number off an image, one digit at a time, and decides for itself how many there were.
That is a larger step than it looks. A classifier produces a fixed-size answer: ten scores, pick the biggest. A sequence has a length the model has to choose, which means something has to say where it stops.
The encoder is unchanged from the classifier. Everything new is on the right-hand side.
What a Decoder Adds
A start token. The decoder generates one position at a time, each conditioned on what came before. At the first step there is no “before”, so you feed it a token that means begin.
A causal mask. During training the decoder sees the whole target sequence at once, for speed. But position 2 must not see position 3. Otherwise the model learns to read the answer it is being asked to predict.
At inference the future genuinely is not there, so a model trained that way produces nothing useful. The mask stops it: each position sees only itself and everything to its left.
Cross-attention. The decoder needs to look at the image. Self-attention lets the output positions look at each other; cross-attention lets each output position query the encoder’s patch tokens and pull in what it needs.
And then an end token, so the model can say it’s done. In a task whose answers vary in length, that is what makes the length the model’s decision rather than a hyperparameter.
The canvas here is always the same width, but between one and three digits are drawn on it and the rest is left blank. So the end token carries real weight: the model has to decide how many digits it saw, and a read counts as correct only if the digits are right and it stops in the right place. A three-digit answer to a two-digit image is wrong however good the digits are.
Show the code
# Setup, inlined rather than imported so this notebook runs on its own.# Keep it folded; nothing below it depends on anything outside this file.import torch.nn as nnimport visualtorchimport matplotlib.pyplot as plt# --- chart styling -------------------------------------------------------# Categorical slots of a CVD-validated palette: blue, orange, aqua, purple.COLOURS = ["#2a78d6", "#eb6834", "#1baf7a", "#8a63d2"]MUTED, GRID, AXIS ="#5b6570", "#e6e6e3", "#d5d5d1"def style_axes(ax, xlabel=None, ylabel=None, grid="y"):"""Strip an axes back to the ink that carries information."""if xlabel: ax.set_xlabel(xlabel, color=MUTED, fontsize=9)if ylabel: ax.set_ylabel(ylabel, color=MUTED, fontsize=9)if grid: ax.grid(axis=grid, color=GRID, linewidth=0.8) ax.set_axisbelow(True)for side in ("top", "right"): ax.spines[side].set_visible(False)for side in ("left", "bottom"): ax.spines[side].set_color(AXIS) ax.tick_params(colors=MUTED, labelsize=9, length=0)return axdef figure(width=7.0, height=4.2, **kw): fig, ax = plt.subplots(figsize=(width, height), **kw)return fig, ax# --- visualtorch, configured ---------------------------------------------# Dropout is hidden: nn.TransformerEncoderLayer exposes its internal Dropout as a# traceable leaf, every model here builds it with dropout=0.0, and drawing a no-op# layer says it is part of the architecture when it is not. Uncoloured it also took# visualtorch's default orange, near enough to MultiheadAttention's to be confusing.COLOUR_MAP = { nn.Linear: {"fill": COLOURS[0]}, nn.MultiheadAttention: {"fill": COLOURS[1]}, nn.Embedding: {"fill": COLOURS[3]}, nn.LayerNorm: {"fill": "#aeb6bf"}, nn.GELU: {"fill": COLOURS[2]}, nn.ReLU: {"fill": COLOURS[2]}, nn.Flatten: {"fill": "#aeb6bf"}, nn.Dropout: {"fill": "#e6e6e3"},}_COMMON =dict(color_map=COLOUR_MAP, connector_fill="#c3c9d0", background_fill="white", font_color=MUTED, legend=True, show_dimension=True, type_ignore=[nn.Dropout])def diagram(model, input_shape, style="graph", **overrides):"""Render `model` as a PIL image in this site's colours. `graph` draws real neurons and suits a short stack; `flow` draws volumetric blocks whose size tracks the layer's and stays readable on a deep one. `flow` reports a 3-D activation as (1, 1, width), losing the sequence length, so its shape labels are off. `level_gap=1` keeps a residual connection drawn close to the blocks it skips. """if style =="graph": settings =dict(node_size=24, layer_spacing=110, node_spacing=8, ellipsize_after=5, **_COMMON)else: settings =dict(spacing=26, scale_xy=2.4, max_xy=280, one_dim_orientation="y", level_gap=1, **_COMMON) settings["show_dimension"] =False settings.update(overrides)return visualtorch.render(model, input_shape=input_shape, style=style, **settings)import matplotlib.pyplot as pltimport numpy as npimport torch# Re-running this page must reproduce it, so a result that shifts between runs would let# the prose and the output disagree. Torch's multithreaded CPU reductions add floats# in whatever order the threads finish in; over a training loop that compounds into a# different model. One thread makes the run reproducible.torch.set_num_threads(1)L =3labels = ["<start>", "digit 1", "digit 2", "digit 3"]mask = np.tril(np.ones((L +1, L +1)))fig, ax = plt.subplots(figsize=(4.4, 4.0))ax.imshow(mask, cmap="Blues", vmin=0, vmax=1.7)ax.set_xticks(range(L +1), labels, rotation=45, ha="right")ax.set_yticks(range(L +1), labels)ax.set_xlabel("can attend to", color=MUTED, fontsize=9)ax.set_ylabel("generating", color=MUTED, fontsize=9)ax.tick_params(colors=MUTED, labelsize=9, length=0)for s in ax.spines.values(): s.set_visible(False)fig.tight_layout()
Figure 1: The causal mask. A filled cell means that output position is allowed to attend to that one; the blank upper triangle is the future, hidden.
Teacher Forcing
There’s a subtlety in how this gets trained that took me a while to appreciate.
During training the decoder is fed the true previous digits at every position. That’s called teacher forcing, and it’s what lets the whole sequence be computed in one pass instead of three.
At inference there are no true previous digits. The model is fed its own previous outputs. So if it gets digit 1 wrong, digit 2 is being predicted from a prefix that never occurred during training.
So the training loss systematically flatters the model. The honest measure is generating the whole sequence the way you actually would.
That is the number reported below. Guessing has to get both the digits and the length right, which is why chance sits near one in thirty rather than one in a thousand.
Show the code
import torch.nn as nnfrom huggingface_hub import hf_hub_downloadDEVICE = torch.device("mps"if torch.backends.mps.is_available() else"cpu")REVISION ="6bbb59474efaf5edfb62408686909305dbd45335"data = np.load(hf_hub_download("roshbeed/ai-residency-blog-data", "mnist/mnist-full.npz", repo_type="dataset", revision=REVISION))x_train = torch.from_numpy(data["x_train"]).float() /255.0y_train = torch.from_numpy(data["y_train"]).long()x_test = torch.from_numpy(data["x_test"]).float() /255.0y_test = torch.from_numpy(data["y_test"]).long()START, END, PAD =10, 11, 12# three entries alongside the ten digitsVOCAB =13def compose(images, labels, n, seed):"""One to three digits on a 28x84 canvas, left-aligned, the rest left blank. The canvas is always the same size. How much of it is occupied is not, which is the whole point: the model has to decide how many digits to emit. """ rng = np.random.default_rng(seed) counts = torch.from_numpy(rng.integers(1, L +1, n)) pick = torch.from_numpy(rng.integers(0, len(images), (n, L))) present = torch.arange(L).unsqueeze(0) < counts.unsqueeze(1) tiles = images[pick] * present.unsqueeze(-1).unsqueeze(-1) canvas = torch.cat([tiles[:, slot] for slot inrange(L)], dim=2) digits = torch.where(present, labels[pick], torch.full_like(labels[pick], PAD))return canvas, digits, countsdef teacher_pair(digits, counts):"""`<start> d1 d2` in, `d1 d2 <end>` out, padded to the same width.""" rows =len(digits) given = torch.cat([torch.full((rows, 1), START), digits], dim=1) wanted = torch.full((rows, L +1), PAD, dtype=torch.long) wanted[:, :L] = digits wanted[torch.arange(rows), counts] = ENDreturn given, wantedX, Y, counts_train = compose(x_train, y_train, 60_000, seed=0)X_test, Y_test, counts_test = compose(x_test, y_test, 10_000, seed=1)decoder_input, decoder_target = teacher_pair(Y, counts_train)_, target_test = teacher_pair(Y_test, counts_test)X, Y = X.to(DEVICE), Y.to(DEVICE)X_test, Y_test = X_test.to(DEVICE), Y_test.to(DEVICE)decoder_input, decoder_target = decoder_input.to(DEVICE), decoder_target.to(DEVICE)target_test, counts_test = target_test.to(DEVICE), counts_test.to(DEVICE)spread = torch.bincount(counts_train, minlength=L +1)[1:]chance =sum((spread[k -1] /len(X)).item() *10**-k for k inrange(1, L +1))print(f"{len(X):,} composites of shape {tuple(X.shape[1:])}")print(f"digits per image: "+", ".join(f"{k}: {spread[k-1]:,}"for k inrange(1, L +1)))print(f"guessing digits and length uniformly: {chance:.4f} exact-sequence accuracy")
60,000 composites of shape (28, 84)
digits per image: 1: 19,935, 2: 20,172, 3: 19,893
guessing digits and length uniformly: 0.0369 exact-sequence accuracy
class Reader(nn.Module):def__init__(self):super().__init__()self.project = nn.Linear(PATCH**2, DIM)self.image_positions = nn.Parameter(torch.randn(1, PATCHES, DIM) *0.02)self.encoder = nn.ModuleList([EncoderBlock() for _ inrange(LAYERS)])self.embed = nn.Embedding(VOCAB, DIM) # ten digits, <start>, <end>, <pad>self.token_positions = nn.Parameter(torch.randn(1, L +1, DIM) *0.02)self.decoder = nn.ModuleList([DecoderBlock() for _ inrange(LAYERS)])self.out = nn.Linear(DIM, VOCAB)def encode(self, images): h =self.project(patchify(images)) +self.image_positionsfor block inself.encoder: h = block(h)return hdef decode(self, memory, tokens, want_weights=False): h =self.embed(tokens) +self.token_positions[:, :tokens.shape[1]] mask = torch.triu(torch.full((tokens.shape[1],) *2, float("-inf"), device=tokens.device), diagonal=1) weights =Nonefor i, block inenumerate(self.decoder): h, w = block(h, memory, mask, want_weights and i == LAYERS -1)if w isnotNone: weights = wreturnself.out(h), weightsdef forward(self, images, tokens):returnself.decode(self.encode(images), tokens)[0]
One in twenty-seven is the bar, and it is that high because a guess only has to get one digit right a third of the time. Anything above it is the model reading rather than guessing.
Both halves together. The image enters top left. The tokens enter bottom left through their embedding. The two streams meet in the decoder.
Show the code
diagram(Reader(), input_shape=((1, 28, 84), (1, L +1)), style="flow", input_dtype=(torch.float32, torch.long))
Figure 2: The encoder-decoder. Two inputs, two stacks, joined by the decoder’s cross-attention.
So the model below is built with two ways to run it: one for training, and one for measuring, which generates a digit at a time from the model’s own output the way inference has to.
Show the code
torch.manual_seed(0)model = Reader().to(DEVICE)optimiser = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=0.01)generator = torch.Generator().manual_seed(0)@torch.no_grad()def generate():"""One token at a time from the model's own output, as at inference.""" tokens = torch.full((len(X_test), 1), START, device=DEVICE) memory = model.encode(X_test)for _ inrange(L +1): # three digits, and then <end> nxt = model.decode(memory, tokens)[0][:, -1].argmax(-1, keepdim=True) tokens = torch.cat([tokens, nxt], dim=1)return tokens[:, 1:]def exact_sequence():"""Correct means the right digits, and <end> in the right place. A three-digit answer to a two-digit image is wrong however good the digits are, so length is scored, not assumed. """ out = generate() wanted = target_test scored = wanted != PAD # nothing after <end> is scoredreturn ((out == wanted) |~scored).all(1).float().mean().item()print(f"{sum(p.numel() for p in model.parameters()):,} parameters")
11,091,981 parameters
Fifteen epochs, reporting after each how often a held-out image is read exactly right: every digit correct, and the end token in the right place.
96.3% of held-out images read exactly right, against 3.7% for guessing: every digit correct, and the end token in the right place.
The column is bumpier than a training loss would be, and that is the measure rather than the model: one wrong digit early in a sequence takes the whole sequence with it, so a small change in per-digit accuracy moves this number three times as far.
Cross-Attention Alignment
Nobody told this model the digits run left to right.
The encoder produces 48 patch tokens in a 4×12 grid. Which patches does each output position attend to?
Figure 3: Cross-attention from each output position onto the image patches. The alignment is learned, not specified.
Putting numbers on that: the share of each output position’s attention landing in each third of the image.
Show the code
print("share of each output's attention falling on each third of the image\n")print(f"{'':10}{'left':>8}{'middle':>8}{'right':>8}")for position inrange(L): by_column = attention[position].sum(0) thirds = [by_column[i:i +4].sum() / by_column.sum() for i in (0, 4, 8)]print(f"digit {position +1:<4} "+"".join(f"{t:>8.2f}"for t in thirds))
share of each output's attention falling on each third of the image
left middle right
digit 1 0.86 0.12 0.01
digit 2 0.27 0.68 0.05
digit 3 0.00 0.01 0.99
Each output position learns to look at its own third of the image. The loss never mentions position. It only ever says the first digit is a 7. The alignment between output order and image geometry is something the model works out because it’s the only way to get the answer right.
The same mechanism makes translation work. There the alignment is between words in two languages rather than positions in an image. It’s why cross-attention replaced the fixed-size context vector that encoder-decoders used before it.
Conclusion
A decoder is not a bigger classifier. It is a different contract.
The model commits to one token
It then conditions on its own commitment
Errors compound down the sequence
Teacher forcing hides all of that during training
Which is why the number that matters has to be generated the slow way.
One consequence of generating a token at a time is that you can watch it happen. The service streams each digit the moment the decoder emits it, so the page shows the number being read rather than appearing at once. That is not a UI flourish, it is what the architecture does, made visible.
Where this goes
The same encoder-decoder points at something that is not an image at all.
Swap the patch projection for a convolution over a waveform and the decoder now writes music notation. Nothing in the middle changed. That is the same point the encoder diagram made, and it is what fine-tuning Whisper is built on.