← back to the mech-interp place

2D transformers: reading the depth axis

by Alex Jerpelea · June 2026 · code on GitHub

A transformer's forward pass leaves behind a ladder of hidden states, one per layer. What if a second, small transformer read that ladder as a sequence in its own right? And what if the two were trained jointly?

The idea

Training models on other models' hidden states is not new. Natural-language autoencoders trained decoders to reconstruct text from a latent code (Bowman et al., 2016); embedding inversion trains a model to decode text back out of frozen embeddings (Morris et al., 2023); the tuned lens fits per-layer probes that translate intermediate states into predictions (Belrose et al., 2023); Patchscopes has a model read its own patched-in states in natural language (Ghandeharioun et al., 2024); Coconut feeds final hidden states back in as inputs to reason in latent space (Hao et al., 2024). A thinner line of work aggregates the layer axis specifically: ELMo learned a scalar mix over its layers (Peters et al., 2018), and DenseFormer learns depth-weighted averages of all previous layers at every layer (Pagliardini et al., 2024).

What all the reader-style work above shares is that the hidden states are frozen: the reader adapts to whatever geometry the base model already has. Here the reader is trained jointly with the backbone, from scratch. A depth-10 backbone \(H\) (a stock nanochat model) produces the 11-rung residual ladder \([x_0, h_1, \ldots, h_{10}]\); a small bidirectional reader \(V\) attends over those rungs and owns the readout: no \(h_{10}\) skip, no gate, no identity init, so \(V\) has to earn its keep against the standard top-state baseline. Because gradients flow into \(H\) through \(V\), the hope is that the backbone bends its hidden-state space to be read, keeping things in middle layers that the top layer would otherwise discard. That could be a more interesting geometry, and maybe a more efficient way to scale depth than just stacking more layers.

What six experiments said

Measured in validation bits-per-byte on the same data, the reader kept losing, and each loss localized why. A cheap \(d_V{=}128\) reader loses by +0.056 bpb, even though its depth-attention is genuinely non-degenerate (it nearly ignores the top rung and reads the middle ones). It wasn't the 128-dim readout bottleneck: rank probes bounded that story, and a full-width \(d_V{=}640\) reader still loses (+0.051 iso-FLOP, +0.021 at equal data). It wasn't the backbone's residuals "stealing the reader's job" either: removing them made the reader worse (+0.030), because the residual stream's online accumulation is what makes the rungs good in the first place. The best current account of all the negatives: the residual stream makes every rung a partial sum of one telescoping series, so \(h_{10}\) already holds the whole sum, and re-aggregating the ladder offline is largely redundant.

The one thing that worked was reader depth: going from 2 to 4 bidirectional reader blocks closes the entire gap and lands exactly on the baseline.

Validation bpb curves: baseline 0.877, 2-block WideReader 0.898, 4-block WideReader 0.877 (tie)
A taller depth-reader ties the baseline. Same data budget: top-state baseline 0.877 bpb, full-width reader with 2 blocks 0.898, with 4 blocks 0.877. The gap closes entirely. The reader was depth-limited, not width-limited.

The tie has two readings, and they aren't resolved yet: either four blocks let \(V\) extract a readout from the ladder as good as the top state (in which case more depth might beat it), or the extra capacity let \(V\) collapse into mimicking \(h_{10}\), which also lands exactly on baseline. And a tie isn't enough anyway: the 4-block reader burns ~3.2× the FLOPs per token. The open directions: push reader depth further, and change the objective (a BERT-style masked target would make the top layer task-specialized and force it to shed information that middle rungs keep, so for the first time the ladder would be non-redundant, and \(V\) would have a real job).