Skip to content
NotebookLive

The logit lens in PyTorch: a hands-on GPT-2 tutorial

Watch GPT-2 build its prediction layer by layer. A runnable walkthrough that decodes the residual stream after every block, shows why ln_final matters, and plots top-k and rank heatmaps.

Published
Length
15 min read
Topics
  • Interpretability
  • Logit lens
  • PyTorch
  • TransformerLens
  • GPT-2

The idea

GPT-2 is a stack of blocks. Each token carries a vector, the residual stream (768 numbers in gpt2-small), and every block reads it and adds its result back. Normally only the last vector becomes a prediction. The logit lens applies the same vector-to-word-scores step to every intermediate vector, so you can see when the model figures out the next token.

What it covers

  • Setup: load GPT-2 with its original weights and poke at W_E and W_U
  • Tokens and one forward pass, with honest labels for byte fragments
  • The residual stream: stack every layer, and why raw vectors can't be compared
  • Normalization and the first lens: plain norm vs ln_final
  • Comparing lenses by agreement and rank
  • A reusable run_example(text) pipeline
  • Top-k and rank heatmaps, including Nepali text in Devanagari

Try it without the notebook

The same lens, extended with trajectories and direct logit attribution, is the logit lens lab in attnlab. It reproduces this notebook number for number.

Credits

The logit lens in PyTorch stands on open-source work and published research. Thank you to everyone behind it.

Built with

Background

More from Sarvabhaum

Web appLive

attnlab

A set of labs, walked in order, that run real TransformerLens models behind a web page. Type a prompt, pick a model, and see how it is tokenized, where every head attends, and when the model knows its answer.

  • Mechanistic interpretability
  • TransformerLens
  • Attention
InteractiveLive

A transformer block, one matrix at a time

Follow four tokens through a single transformer block from the residual stream's point of view. Every matrix is shown with real numbers, computed live, so you can trace one token's row from embedding to next-token probabilities.

  • Transformers
  • Residual stream
  • Attention