<?xml version="1.0" encoding="utf-8"?>
<feed xmlns="http://www.w3.org/2005/Atom"><title>Dom 🆚 Machine</title><link href="https://weirdmachine.wtf/" rel="alternate"/><link href="https://weirdmachine.wtf/feeds/all.atom.xml" rel="self"/><id>https://weirdmachine.wtf/</id><updated>2026-10-04T00:00:00-04:00</updated><entry><title>Two Routes Through a Transformer</title><link href="https://weirdmachine.wtf/posts/two-routes-through-a-transformer/" rel="alternate"/><published>2026-10-04T00:00:00-04:00</published><updated>2026-10-04T00:00:00-04:00</updated><author><name>Dominic Wang</name></author><id>tag:weirdmachine.wtf,2026-10-04:/posts/two-routes-through-a-transformer/</id><summary type="html">&lt;p&gt;The groundwork I needed before A Mathematical Framework for Transformer Circuits and its walkthrough video made sense. What a one-layer attention-only model actually computes.&lt;/p&gt;</summary><content type="html">&lt;p&gt;&lt;a href="https://transformer-circuits.pub/2021/framework/index.html"&gt;A Mathematical Framework for Transformer Circuits&lt;/a&gt; has a &lt;a href="https://www.youtube.com/watch?v=KV5gbOmHbjU"&gt;walkthrough video&lt;/a&gt; to go with it. I could not follow either until I had worked out what a one-layer attention-only model actually computes, so this post is that groundwork: the path decomposition end to end, with every number in the example derived rather than asserted. Architecture background is in &lt;a href="https://weirdmachine.wtf/posts/understanding-transformers/"&gt;Understanding Transformers&lt;/a&gt;.&lt;/p&gt;
&lt;h2 id="path-decomposition-of-transformer-output"&gt;Path decomposition of transformer output&lt;/h2&gt;
&lt;h3 id="two-routes"&gt;Two routes&lt;/h3&gt;
&lt;p&gt;One-layer attention-only transformer. Tokens in, logits out at each position. One sequence runs through the whole note:&lt;/p&gt;
&lt;div class="table-wrap"&gt;&lt;table&gt;
&lt;thead&gt;
&lt;tr&gt;
&lt;th&gt;pos&lt;/th&gt;
&lt;th&gt;1&lt;/th&gt;
&lt;th&gt;2&lt;/th&gt;
&lt;th&gt;3&lt;/th&gt;
&lt;th&gt;4&lt;/th&gt;
&lt;th&gt;5&lt;/th&gt;
&lt;th&gt;6&lt;/th&gt;
&lt;th&gt;7&lt;/th&gt;
&lt;th&gt;8&lt;/th&gt;
&lt;th&gt;9&lt;/th&gt;
&lt;/tr&gt;
&lt;/thead&gt;
&lt;tbody&gt;
&lt;tr&gt;
&lt;td&gt;token&lt;/td&gt;
&lt;td&gt;&lt;code&gt;the&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;&lt;code&gt;cat&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;&lt;code&gt;sat&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;&lt;code&gt;on&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;&lt;code&gt;the&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;&lt;code&gt;mat&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;&lt;code&gt;.&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;&lt;code&gt;the&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;&lt;code&gt;cat&lt;/code&gt;&lt;/td&gt;
&lt;/tr&gt;
&lt;/tbody&gt;
&lt;/table&gt;&lt;/div&gt;
&lt;p&gt;Two routes run from a token to a logit. Both are below, over the first three positions. The next two sections take one each:&lt;/p&gt;
&lt;div class="language-text highlight"&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;          pos 1       pos 2       pos 3
token:    &amp;quot;the&amp;quot;       &amp;quot;cat&amp;quot;       &amp;quot;sat&amp;quot;
            |           |           |
          embed       embed       embed
            |           |           |
            |           |           &amp;#39;---------------.
            |           |                           |
            |           |      route 1 (direct):    |
            |           |      rides the residual   |
            |           |      stream, untouched    |
            |           |                           |
            &amp;#39;-----------+-----&amp;gt; head h -------------+
                                                    |
              route 2 (attention):                  |
              the head reads back over              |
              earlier positions, taking             |
              some % from each, and adds            |
              what it grabbed                       |
                                                    v
                                                 unembed
                                                    |
                                                    v
                                            logits at pos 3
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;p&gt;Position 3 can read positions 1, 2 and 3, never forward. That restriction is the &lt;strong&gt;causal mask&lt;/strong&gt;.&lt;/p&gt;
&lt;p&gt;Three simplifications throughout: no MLPs, LayerNorm folded into adjacent weights, and no positional embeddings.&lt;/p&gt;
&lt;h3 id="route-1-the-tokens-own-embedding"&gt;Route 1: the token's own embedding&lt;/h3&gt;
&lt;p&gt;&lt;div class="mermaid"&gt;flowchart TB
    t1["pos 1&amp;lt;br/&amp;gt;the"] --&amp;gt; e1["embed"]
    t2["pos 2&amp;lt;br/&amp;gt;cat"] --&amp;gt; e2["embed"]
    t3["pos 3&amp;lt;br/&amp;gt;sat"] --&amp;gt; e3["embed"]

    e3 ==&amp;gt;|"route 1 (direct)&amp;lt;br/&amp;gt;rides the residual stream, untouched"| u["unembed"]
    u --&amp;gt; out["logits at pos 3"]

    classDef dim stroke-dasharray: 4 4
    class t1,t2,e1,e2 dim&lt;/div&gt;
&lt;em&gt;Dashed = present in the sequence, unused by this route.&lt;/em&gt;&lt;/p&gt;
&lt;p&gt;&lt;code&gt;"sat"&lt;/code&gt; is embedded, rides the residual stream, hits the unembed untouched. Alone that makes a &lt;strong&gt;bigram table&lt;/strong&gt;: a lookup keyed on one previous token.&lt;/p&gt;
&lt;div class="table-wrap"&gt;&lt;table&gt;
&lt;thead&gt;
&lt;tr&gt;
&lt;th&gt;&lt;/th&gt;
&lt;th&gt;&lt;code&gt;the&lt;/code&gt;&lt;/th&gt;
&lt;th&gt;&lt;code&gt;cat&lt;/code&gt;&lt;/th&gt;
&lt;th&gt;&lt;code&gt;sat&lt;/code&gt;&lt;/th&gt;
&lt;th&gt;&lt;code&gt;on&lt;/code&gt;&lt;/th&gt;
&lt;th&gt;&lt;code&gt;mat&lt;/code&gt;&lt;/th&gt;
&lt;th&gt;&lt;code&gt;.&lt;/code&gt;&lt;/th&gt;
&lt;/tr&gt;
&lt;/thead&gt;
&lt;tbody&gt;
&lt;tr&gt;
&lt;td&gt;&lt;strong&gt;&lt;code&gt;the&lt;/code&gt;&lt;/strong&gt;&lt;/td&gt;
&lt;td&gt;0.0&lt;/td&gt;
&lt;td&gt;3.7&lt;/td&gt;
&lt;td&gt;0.2&lt;/td&gt;
&lt;td&gt;0.1&lt;/td&gt;
&lt;td&gt;4.1&lt;/td&gt;
&lt;td&gt;0.0&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;strong&gt;&lt;code&gt;cat&lt;/code&gt;&lt;/strong&gt;&lt;/td&gt;
&lt;td&gt;0.3&lt;/td&gt;
&lt;td&gt;0.0&lt;/td&gt;
&lt;td&gt;3.9&lt;/td&gt;
&lt;td&gt;0.4&lt;/td&gt;
&lt;td&gt;0.1&lt;/td&gt;
&lt;td&gt;1.2&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;strong&gt;&lt;code&gt;sat&lt;/code&gt;&lt;/strong&gt;&lt;/td&gt;
&lt;td&gt;1.1&lt;/td&gt;
&lt;td&gt;0.1&lt;/td&gt;
&lt;td&gt;0.0&lt;/td&gt;
&lt;td&gt;4.2&lt;/td&gt;
&lt;td&gt;0.2&lt;/td&gt;
&lt;td&gt;0.9&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;strong&gt;&lt;code&gt;on&lt;/code&gt;&lt;/strong&gt;&lt;/td&gt;
&lt;td&gt;4.4&lt;/td&gt;
&lt;td&gt;0.2&lt;/td&gt;
&lt;td&gt;0.0&lt;/td&gt;
&lt;td&gt;0.0&lt;/td&gt;
&lt;td&gt;0.6&lt;/td&gt;
&lt;td&gt;0.1&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;strong&gt;&lt;code&gt;mat&lt;/code&gt;&lt;/strong&gt;&lt;/td&gt;
&lt;td&gt;0.5&lt;/td&gt;
&lt;td&gt;0.1&lt;/td&gt;
&lt;td&gt;0.1&lt;/td&gt;
&lt;td&gt;0.2&lt;/td&gt;
&lt;td&gt;0.0&lt;/td&gt;
&lt;td&gt;3.8&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;strong&gt;&lt;code&gt;.&lt;/code&gt;&lt;/strong&gt;&lt;/td&gt;
&lt;td&gt;3.5&lt;/td&gt;
&lt;td&gt;0.2&lt;/td&gt;
&lt;td&gt;0.0&lt;/td&gt;
&lt;td&gt;0.1&lt;/td&gt;
&lt;td&gt;0.1&lt;/td&gt;
&lt;td&gt;0.0&lt;/td&gt;
&lt;/tr&gt;
&lt;/tbody&gt;
&lt;/table&gt;&lt;/div&gt;
&lt;p&gt;&lt;strong&gt;Route 1: the bigram table.&lt;/strong&gt; &lt;em&gt;Row = the token you just saw, column = a candidate next token, cell = score. Six tokens here; the real table is &lt;code&gt;n_vocab x n_vocab&lt;/code&gt;.&lt;/em&gt;&lt;/p&gt;
&lt;p&gt;Every table in this post is hand-built to make the mechanism visible. None of these numbers come off a trained model. What is derived is everything downstream of them: the attention percentages, the logit shifts and the totals all follow from these tables by arithmetic you can check.&lt;/p&gt;
&lt;p&gt;Walk the first three positions through it:&lt;/p&gt;
&lt;ul&gt;
&lt;li&gt;&lt;strong&gt;&lt;code&gt;"the"&lt;/code&gt;&lt;/strong&gt; -&amp;gt; read row &lt;code&gt;the&lt;/code&gt; -&amp;gt; &lt;code&gt;mat&lt;/code&gt; 4.1 beats &lt;code&gt;cat&lt;/code&gt; 3.7. Predicts &lt;code&gt;mat&lt;/code&gt;. Our sequence has &lt;code&gt;cat&lt;/code&gt;, so this one is wrong.&lt;/li&gt;
&lt;li&gt;&lt;strong&gt;&lt;code&gt;"the cat"&lt;/code&gt;&lt;/strong&gt; -&amp;gt; read row &lt;code&gt;cat&lt;/code&gt; -&amp;gt; &lt;code&gt;sat&lt;/code&gt; scores 3.9. Predicts &lt;code&gt;sat&lt;/code&gt;. Correct.&lt;/li&gt;
&lt;li&gt;&lt;strong&gt;&lt;code&gt;"the cat sat"&lt;/code&gt;&lt;/strong&gt; -&amp;gt; read row &lt;code&gt;sat&lt;/code&gt; -&amp;gt; &lt;code&gt;on&lt;/code&gt; scores 4.2. Predicts &lt;code&gt;on&lt;/code&gt;. Correct.&lt;/li&gt;
&lt;/ul&gt;
&lt;p&gt;The last step read row &lt;code&gt;sat&lt;/code&gt; and nothing else. There is no row for &lt;code&gt;"the cat sat"&lt;/code&gt;, so &lt;code&gt;"the cat"&lt;/code&gt; is gone. By position 3 the model remembers only the token under it.&lt;/p&gt;
&lt;p&gt;Reaching back for what it forgot is route 2's job.&lt;/p&gt;
&lt;h3 id="route-2-what-the-head-reaches-back-for"&gt;Route 2: what the head reaches back for&lt;/h3&gt;
&lt;div class="mermaid"&gt;flowchart TB
    t1["pos 1&amp;lt;br/&amp;gt;the"] --&amp;gt; e1["embed"]
    t2["pos 2&amp;lt;br/&amp;gt;cat"] --&amp;gt; e2["embed"]
    t3["pos 3&amp;lt;br/&amp;gt;sat"] --&amp;gt; e3["embed"]

    e1 --&amp;gt; h["head h&amp;lt;br/&amp;gt;&amp;lt;i&amp;gt;takes some % from each&amp;lt;/i&amp;gt;"]
    e2 --&amp;gt; h
    e3 --&amp;gt; h

    h ==&amp;gt;|"route 2 (attention)&amp;lt;br/&amp;gt;adds what it grabbed"| u["unembed"]
    u --&amp;gt; out["logits at pos 3"]&lt;/div&gt;
&lt;p&gt;The head reads back over earlier positions, takes a percentage of each, and adds what it grabbed to the residual stream, which carries it to the unembed.&lt;/p&gt;
&lt;p&gt;Stand at position 8, on the third &lt;code&gt;the&lt;/code&gt;. Route 1 reads row &lt;code&gt;the&lt;/code&gt; and answers &lt;code&gt;mat&lt;/code&gt; over &lt;code&gt;cat&lt;/code&gt;, the same mistake as position 1. It cannot know this sentence is about a cat.&lt;/p&gt;
&lt;p&gt;Route 2 takes two steps: &lt;strong&gt;pick where to look&lt;/strong&gt;, then &lt;strong&gt;use what is there&lt;/strong&gt;.&lt;/p&gt;
&lt;p&gt;&lt;strong&gt;Step 1: where to look.&lt;/strong&gt; Every position scores every position it may see. A score is how well a source matches what the destination wants. With no positional embeddings a position carries nothing but its token, so a score depends only on the two tokens involved, the one you are standing on and the one you are looking at. &lt;code&gt;x&lt;/code&gt; marks a cell the &lt;strong&gt;causal mask&lt;/strong&gt; blocks:&lt;/p&gt;
&lt;div class="table-wrap"&gt;&lt;table&gt;
&lt;thead&gt;
&lt;tr&gt;
&lt;th&gt;to  from&lt;/th&gt;
&lt;th&gt;1 &lt;code&gt;the&lt;/code&gt;&lt;/th&gt;
&lt;th&gt;2 &lt;code&gt;cat&lt;/code&gt;&lt;/th&gt;
&lt;th&gt;3 &lt;code&gt;sat&lt;/code&gt;&lt;/th&gt;
&lt;th&gt;4 &lt;code&gt;on&lt;/code&gt;&lt;/th&gt;
&lt;th&gt;5 &lt;code&gt;the&lt;/code&gt;&lt;/th&gt;
&lt;th&gt;6 &lt;code&gt;mat&lt;/code&gt;&lt;/th&gt;
&lt;th&gt;7 &lt;code&gt;.&lt;/code&gt;&lt;/th&gt;
&lt;th&gt;8 &lt;code&gt;the&lt;/code&gt;&lt;/th&gt;
&lt;th&gt;9 &lt;code&gt;cat&lt;/code&gt;&lt;/th&gt;
&lt;/tr&gt;
&lt;/thead&gt;
&lt;tbody&gt;
&lt;tr&gt;
&lt;td&gt;1 &lt;code&gt;the&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;0.4&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;2 &lt;code&gt;cat&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;0.5&lt;/td&gt;
&lt;td&gt;0.3&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;3 &lt;code&gt;sat&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;0.6&lt;/td&gt;
&lt;td&gt;1.2&lt;/td&gt;
&lt;td&gt;0.3&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;4 &lt;code&gt;on&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;0.7&lt;/td&gt;
&lt;td&gt;1.0&lt;/td&gt;
&lt;td&gt;0.5&lt;/td&gt;
&lt;td&gt;0.3&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;5 &lt;code&gt;the&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;0.4&lt;/td&gt;
&lt;td&gt;3.4&lt;/td&gt;
&lt;td&gt;0.7&lt;/td&gt;
&lt;td&gt;0.6&lt;/td&gt;
&lt;td&gt;0.4&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;6 &lt;code&gt;mat&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;0.5&lt;/td&gt;
&lt;td&gt;0.9&lt;/td&gt;
&lt;td&gt;0.4&lt;/td&gt;
&lt;td&gt;0.6&lt;/td&gt;
&lt;td&gt;0.5&lt;/td&gt;
&lt;td&gt;0.3&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;7 &lt;code&gt;.&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;0.4&lt;/td&gt;
&lt;td&gt;0.8&lt;/td&gt;
&lt;td&gt;0.5&lt;/td&gt;
&lt;td&gt;0.4&lt;/td&gt;
&lt;td&gt;0.4&lt;/td&gt;
&lt;td&gt;0.9&lt;/td&gt;
&lt;td&gt;0.3&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;strong&gt;8 &lt;code&gt;the&lt;/code&gt;&lt;/strong&gt;&lt;/td&gt;
&lt;td&gt;&lt;strong&gt;0.4&lt;/strong&gt;&lt;/td&gt;
&lt;td&gt;&lt;strong&gt;3.4&lt;/strong&gt;&lt;/td&gt;
&lt;td&gt;&lt;strong&gt;0.7&lt;/strong&gt;&lt;/td&gt;
&lt;td&gt;&lt;strong&gt;0.6&lt;/strong&gt;&lt;/td&gt;
&lt;td&gt;&lt;strong&gt;0.4&lt;/strong&gt;&lt;/td&gt;
&lt;td&gt;&lt;strong&gt;2.3&lt;/strong&gt;&lt;/td&gt;
&lt;td&gt;&lt;strong&gt;0.5&lt;/strong&gt;&lt;/td&gt;
&lt;td&gt;&lt;strong&gt;0.4&lt;/strong&gt;&lt;/td&gt;
&lt;td&gt;&lt;strong&gt;x&lt;/strong&gt;&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;9 &lt;code&gt;cat&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;0.5&lt;/td&gt;
&lt;td&gt;0.3&lt;/td&gt;
&lt;td&gt;0.6&lt;/td&gt;
&lt;td&gt;0.4&lt;/td&gt;
&lt;td&gt;0.5&lt;/td&gt;
&lt;td&gt;0.8&lt;/td&gt;
&lt;td&gt;0.3&lt;/td&gt;
&lt;td&gt;0.5&lt;/td&gt;
&lt;td&gt;0.3&lt;/td&gt;
&lt;/tr&gt;
&lt;/tbody&gt;
&lt;/table&gt;&lt;/div&gt;
&lt;p&gt;&lt;strong&gt;Route 2, step 1: the score matrix.&lt;/strong&gt; &lt;em&gt;Row = where you are standing, column = where you are looking, cell = the score.&lt;/em&gt;&lt;/p&gt;
&lt;p&gt;Lower triangular: the causal mask drawn out. Rows 5 and 8 are both &lt;code&gt;the&lt;/code&gt; and both score &lt;code&gt;cat&lt;/code&gt; highest, because a &lt;code&gt;the&lt;/code&gt; wants a noun. Row 1 is a &lt;code&gt;the&lt;/code&gt; with nothing behind it.&lt;/p&gt;
&lt;p&gt;Row 8 is ours. Its &lt;code&gt;x&lt;/code&gt; in column 9 hides &lt;code&gt;cat&lt;/code&gt;, the token position 8 must predict. Unmasked it reads 3.4, takes most of the attention, and the model copies the answer instead of inferring it.&lt;/p&gt;
&lt;p&gt;Now take row 8 out of the matrix:&lt;/p&gt;
&lt;div class="language-text highlight"&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;  0.4   3.4   0.7   0.6   0.4   2.3   0.5   0.4    x
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;p&gt;These go into softmax. Real attention divides by &lt;span class="arithmatex"&gt;\(\sqrt{d_\text{head}}\)&lt;/span&gt; first, which changes the numbers but not the shape. Exponentiate each, divide by the total:&lt;/p&gt;
&lt;div class="table-wrap"&gt;&lt;table&gt;
&lt;thead&gt;
&lt;tr&gt;
&lt;th&gt;pos&lt;/th&gt;
&lt;th&gt;token&lt;/th&gt;
&lt;th&gt;score&lt;/th&gt;
&lt;th&gt;exp&lt;/th&gt;
&lt;th&gt;attention&lt;/th&gt;
&lt;/tr&gt;
&lt;/thead&gt;
&lt;tbody&gt;
&lt;tr&gt;
&lt;td&gt;1&lt;/td&gt;
&lt;td&gt;&lt;code&gt;the&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;0.4&lt;/td&gt;
&lt;td&gt;1.5&lt;/td&gt;
&lt;td&gt;1.5 / 50 = 3%&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;2&lt;/td&gt;
&lt;td&gt;&lt;code&gt;cat&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;&lt;strong&gt;3.4&lt;/strong&gt;&lt;/td&gt;
&lt;td&gt;&lt;strong&gt;30.0&lt;/strong&gt;&lt;/td&gt;
&lt;td&gt;30 / 50 = &lt;strong&gt;60%&lt;/strong&gt;&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;3&lt;/td&gt;
&lt;td&gt;&lt;code&gt;sat&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;0.7&lt;/td&gt;
&lt;td&gt;2.0&lt;/td&gt;
&lt;td&gt;2.0 / 50 = 4%&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;4&lt;/td&gt;
&lt;td&gt;&lt;code&gt;on&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;0.6&lt;/td&gt;
&lt;td&gt;1.8&lt;/td&gt;
&lt;td&gt;1.8 / 50 = 4%&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;5&lt;/td&gt;
&lt;td&gt;&lt;code&gt;the&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;0.4&lt;/td&gt;
&lt;td&gt;1.5&lt;/td&gt;
&lt;td&gt;1.5 / 50 = 3%&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;6&lt;/td&gt;
&lt;td&gt;&lt;code&gt;mat&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;&lt;strong&gt;2.3&lt;/strong&gt;&lt;/td&gt;
&lt;td&gt;&lt;strong&gt;10.0&lt;/strong&gt;&lt;/td&gt;
&lt;td&gt;10 / 50 = &lt;strong&gt;20%&lt;/strong&gt;&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;7&lt;/td&gt;
&lt;td&gt;&lt;code&gt;.&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;0.5&lt;/td&gt;
&lt;td&gt;1.7&lt;/td&gt;
&lt;td&gt;1.7 / 50 = 3%&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;8&lt;/td&gt;
&lt;td&gt;&lt;code&gt;the&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;0.4&lt;/td&gt;
&lt;td&gt;1.5&lt;/td&gt;
&lt;td&gt;1.5 / 50 = 3%&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;9&lt;/td&gt;
&lt;td&gt;&lt;code&gt;cat&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;&lt;strong&gt;0.0&lt;/strong&gt;&lt;/td&gt;
&lt;td&gt;0 / 50 = &lt;strong&gt;0%&lt;/strong&gt;&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;/td&gt;
&lt;td&gt;&lt;/td&gt;
&lt;td&gt;&lt;/td&gt;
&lt;td&gt;&lt;strong&gt;50.0&lt;/strong&gt;&lt;/td&gt;
&lt;td&gt;&lt;strong&gt;100%&lt;/strong&gt;&lt;/td&gt;
&lt;/tr&gt;
&lt;/tbody&gt;
&lt;/table&gt;&lt;/div&gt;
&lt;p&gt;&lt;strong&gt;Route 2, step 1: row 8 softmaxed.&lt;/strong&gt; &lt;em&gt;Score to &lt;code&gt;exp&lt;/code&gt; to share of the total.&lt;/em&gt;&lt;/p&gt;
&lt;p&gt;Exponentiating decides it:&lt;/p&gt;
&lt;ul&gt;
&lt;li&gt;&lt;strong&gt;&lt;code&gt;cat&lt;/code&gt; at position 2&lt;/strong&gt; -&amp;gt; scores 3.4 -&amp;gt; &lt;code&gt;exp&lt;/code&gt; 30.0 -&amp;gt; 30.0 / 50.0 = &lt;strong&gt;60%&lt;/strong&gt;&lt;/li&gt;
&lt;li&gt;&lt;strong&gt;&lt;code&gt;mat&lt;/code&gt; at position 6&lt;/strong&gt; -&amp;gt; scores 2.3 -&amp;gt; &lt;code&gt;exp&lt;/code&gt; 10.0 -&amp;gt; 10.0 / 50.0 = &lt;strong&gt;20%&lt;/strong&gt;&lt;/li&gt;
&lt;li&gt;&lt;strong&gt;position 9&lt;/strong&gt; -&amp;gt; masked -&amp;gt; &lt;code&gt;exp&lt;/code&gt; 0.0 -&amp;gt; &lt;strong&gt;0%&lt;/strong&gt;&lt;/li&gt;
&lt;/ul&gt;
&lt;p&gt;A gap of &lt;span class="arithmatex"&gt;\(3.4 - 2.3 = 1.1\)&lt;/span&gt; in score becomes a factor of &lt;span class="arithmatex"&gt;\(30.0 / 10.0 = 3\)&lt;/span&gt; in attention. That is what exponentiating does, and it is why &lt;code&gt;cat&lt;/code&gt; ends up holding most of the head.&lt;/p&gt;
&lt;p&gt;This row is row 8 of the head's attention pattern, written &lt;span class="arithmatex"&gt;\(A\)&lt;/span&gt; here and &lt;span class="arithmatex"&gt;\(A^h\)&lt;/span&gt; once there is more than one head.&lt;/p&gt;
&lt;p&gt;&lt;strong&gt;Step 2: what it does.&lt;/strong&gt; Route 2's own lookup table, same shape as the bigram one, keyed on the position attended to rather than the token under you:&lt;/p&gt;
&lt;div class="table-wrap"&gt;&lt;table&gt;
&lt;thead&gt;
&lt;tr&gt;
&lt;th&gt;&lt;/th&gt;
&lt;th&gt;&lt;code&gt;the&lt;/code&gt;&lt;/th&gt;
&lt;th&gt;&lt;code&gt;cat&lt;/code&gt;&lt;/th&gt;
&lt;th&gt;&lt;code&gt;sat&lt;/code&gt;&lt;/th&gt;
&lt;th&gt;&lt;code&gt;on&lt;/code&gt;&lt;/th&gt;
&lt;th&gt;&lt;code&gt;mat&lt;/code&gt;&lt;/th&gt;
&lt;th&gt;&lt;code&gt;.&lt;/code&gt;&lt;/th&gt;
&lt;/tr&gt;
&lt;/thead&gt;
&lt;tbody&gt;
&lt;tr&gt;
&lt;td&gt;1 &lt;code&gt;the&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;&lt;strong&gt;0.9&lt;/strong&gt;&lt;/td&gt;
&lt;td&gt;0.1&lt;/td&gt;
&lt;td&gt;0.0&lt;/td&gt;
&lt;td&gt;0.1&lt;/td&gt;
&lt;td&gt;0.1&lt;/td&gt;
&lt;td&gt;0.0&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;2 &lt;code&gt;cat&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;0.1&lt;/td&gt;
&lt;td&gt;&lt;strong&gt;5.0&lt;/strong&gt;&lt;/td&gt;
&lt;td&gt;0.2&lt;/td&gt;
&lt;td&gt;0.0&lt;/td&gt;
&lt;td&gt;0.0&lt;/td&gt;
&lt;td&gt;0.0&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;3 &lt;code&gt;sat&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;0.2&lt;/td&gt;
&lt;td&gt;0.1&lt;/td&gt;
&lt;td&gt;&lt;strong&gt;4.1&lt;/strong&gt;&lt;/td&gt;
&lt;td&gt;0.1&lt;/td&gt;
&lt;td&gt;0.1&lt;/td&gt;
&lt;td&gt;0.0&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;4 &lt;code&gt;on&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;0.1&lt;/td&gt;
&lt;td&gt;0.0&lt;/td&gt;
&lt;td&gt;0.1&lt;/td&gt;
&lt;td&gt;&lt;strong&gt;3.4&lt;/strong&gt;&lt;/td&gt;
&lt;td&gt;0.2&lt;/td&gt;
&lt;td&gt;0.0&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;5 &lt;code&gt;the&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;&lt;strong&gt;0.9&lt;/strong&gt;&lt;/td&gt;
&lt;td&gt;0.1&lt;/td&gt;
&lt;td&gt;0.0&lt;/td&gt;
&lt;td&gt;0.1&lt;/td&gt;
&lt;td&gt;0.1&lt;/td&gt;
&lt;td&gt;0.0&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;6 &lt;code&gt;mat&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;0.1&lt;/td&gt;
&lt;td&gt;0.2&lt;/td&gt;
&lt;td&gt;0.1&lt;/td&gt;
&lt;td&gt;0.0&lt;/td&gt;
&lt;td&gt;&lt;strong&gt;2.0&lt;/strong&gt;&lt;/td&gt;
&lt;td&gt;0.1&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;7 &lt;code&gt;.&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;0.0&lt;/td&gt;
&lt;td&gt;0.1&lt;/td&gt;
&lt;td&gt;0.0&lt;/td&gt;
&lt;td&gt;0.0&lt;/td&gt;
&lt;td&gt;0.0&lt;/td&gt;
&lt;td&gt;&lt;strong&gt;1.8&lt;/strong&gt;&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;8 &lt;code&gt;the&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;&lt;strong&gt;0.9&lt;/strong&gt;&lt;/td&gt;
&lt;td&gt;0.1&lt;/td&gt;
&lt;td&gt;0.0&lt;/td&gt;
&lt;td&gt;0.1&lt;/td&gt;
&lt;td&gt;0.1&lt;/td&gt;
&lt;td&gt;0.0&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;9 &lt;code&gt;cat&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;0.1&lt;/td&gt;
&lt;td&gt;&lt;strong&gt;5.0&lt;/strong&gt;&lt;/td&gt;
&lt;td&gt;0.2&lt;/td&gt;
&lt;td&gt;0.0&lt;/td&gt;
&lt;td&gt;0.0&lt;/td&gt;
&lt;td&gt;0.0&lt;/td&gt;
&lt;/tr&gt;
&lt;/tbody&gt;
&lt;/table&gt;&lt;/div&gt;
&lt;p&gt;&lt;strong&gt;Route 2, step 2: the OV table.&lt;/strong&gt; &lt;em&gt;Row = a position, column = a candidate next token, cell = how much attending there shifts the logit.&lt;/em&gt;&lt;/p&gt;
&lt;p&gt;A row depends only on the token at that position, so 1, 5 and 8 are identical, as are 2 and 9. Position 9 has a row like any other; step 1 gave it 0%. The big diagonal is the tell: attend to &lt;code&gt;cat&lt;/code&gt;, boost &lt;code&gt;cat&lt;/code&gt;. A &lt;strong&gt;copying head&lt;/strong&gt;.&lt;/p&gt;
&lt;p&gt;Multiply step 1 by step 2:&lt;/p&gt;
&lt;ul&gt;
&lt;li&gt;&lt;strong&gt;position 2&lt;/strong&gt; -&amp;gt; 60% attention, row &lt;code&gt;2 cat&lt;/code&gt;, &lt;code&gt;cat&lt;/code&gt; scores 5.0 -&amp;gt; 0.60 x 5.0 = &lt;strong&gt;+3.0 to &lt;code&gt;cat&lt;/code&gt;&lt;/strong&gt;&lt;/li&gt;
&lt;li&gt;&lt;strong&gt;position 6&lt;/strong&gt; -&amp;gt; 20% attention, row &lt;code&gt;6 mat&lt;/code&gt;, &lt;code&gt;mat&lt;/code&gt; scores 2.0 -&amp;gt; 0.20 x 2.0 = &lt;strong&gt;+0.4 to &lt;code&gt;mat&lt;/code&gt;&lt;/strong&gt;&lt;/li&gt;
&lt;li&gt;&lt;strong&gt;everything else&lt;/strong&gt; -&amp;gt; 3 to 4% each, on rows that boost &lt;code&gt;the&lt;/code&gt;, &lt;code&gt;sat&lt;/code&gt;, &lt;code&gt;on&lt;/code&gt; and &lt;code&gt;.&lt;/code&gt; -&amp;gt; under 0.1 to either candidate, so the ranking is untouched&lt;/li&gt;
&lt;/ul&gt;
&lt;p&gt;Add the routes:&lt;/p&gt;
&lt;div class="language-text highlight"&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;            route 1   route 2    total
  &amp;quot;mat&amp;quot;       4.1       +0.4       4.5
  &amp;quot;cat&amp;quot;       3.7       +3.0       6.7   &amp;lt;- wins
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;p&gt;Route 1 answered from the token under it and got &lt;code&gt;mat&lt;/code&gt; again. Route 2 reached six positions back, found &lt;code&gt;cat&lt;/code&gt;, and flipped the prediction. That shape is a &lt;strong&gt;skip-trigram&lt;/strong&gt;: &lt;em&gt;X is in the context, so predict X&lt;/em&gt;. A bigram table cannot do it, and it is what one-layer heads mostly learn. Copying names is the classic case: &lt;code&gt;Mr Dursley ... Mr ___&lt;/code&gt; -&amp;gt; &lt;code&gt;Dursley&lt;/code&gt;.&lt;/p&gt;
&lt;p&gt;Two routes, added:&lt;/p&gt;
&lt;div class="language-text highlight"&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;logits at a position  =  route 1  +  one term per attention head
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;h3 id="the-trick"&gt;The trick&lt;/h3&gt;
&lt;p&gt;Those percentages are the &lt;strong&gt;only nonlinear part&lt;/strong&gt; of the model. Everything else is matrix multiplication.&lt;/p&gt;
&lt;p&gt;Pretend they are fixed constants. The path from input token to output logit becomes a chain of matrix multiplies, and a chain of matrices collapses into &lt;strong&gt;one&lt;/strong&gt;.&lt;/p&gt;
&lt;p&gt;That matrix is what you want: vocab in, vocab out, read straight off the weights without running the model. Route 1 gives one, each head another. The formula is the list.&lt;/p&gt;
&lt;h3 id="the-formula"&gt;The formula&lt;/h3&gt;
&lt;div class="arithmatex"&gt;\[T = \text{Id} \otimes W_U W_E + \sum_{h \in H} A^h \otimes (W_U W_{OV}^h W_E)\]&lt;/div&gt;
&lt;p&gt;Every symbol in the formula maps to something from the worked example:&lt;/p&gt;
&lt;div class="table-wrap"&gt;&lt;table&gt;
&lt;thead&gt;
&lt;tr&gt;
&lt;th&gt;Symbol&lt;/th&gt;
&lt;th&gt;Shape&lt;/th&gt;
&lt;th&gt;What it is&lt;/th&gt;
&lt;/tr&gt;
&lt;/thead&gt;
&lt;tbody&gt;
&lt;tr&gt;
&lt;td&gt;&lt;span class="arithmatex"&gt;\(T\)&lt;/span&gt;&lt;/td&gt;
&lt;td&gt;&lt;/td&gt;
&lt;td&gt;the whole model: tokens in, logits out&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;span class="arithmatex"&gt;\(W_E\)&lt;/span&gt;&lt;/td&gt;
&lt;td&gt;&lt;code&gt;[d_model, n_vocab]&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;embed&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;span class="arithmatex"&gt;\(W_U\)&lt;/span&gt;&lt;/td&gt;
&lt;td&gt;&lt;code&gt;[n_vocab, d_model]&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;unembed&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;span class="arithmatex"&gt;\(\text{Id} \otimes W_U W_E\)&lt;/span&gt;&lt;/td&gt;
&lt;td&gt;&lt;/td&gt;
&lt;td&gt;&lt;strong&gt;route 1&lt;/strong&gt;&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;span class="arithmatex"&gt;\(W_U W_E\)&lt;/span&gt;&lt;/td&gt;
&lt;td&gt;&lt;code&gt;[n_vocab, n_vocab]&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;route 1's &lt;strong&gt;bigram table&lt;/strong&gt;&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;span class="arithmatex"&gt;\(\text{Id}\)&lt;/span&gt;&lt;/td&gt;
&lt;td&gt;&lt;code&gt;[n_ctx, n_ctx]&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;"read this position only": route 1 never moves positions&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;span class="arithmatex"&gt;\(\sum_{h \in H} \dots\)&lt;/span&gt;&lt;/td&gt;
&lt;td&gt;&lt;/td&gt;
&lt;td&gt;&lt;strong&gt;one term per head&lt;/strong&gt;; they just add&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;span class="arithmatex"&gt;\(W_{QK}^h = W_Q^{hT} W_K^h\)&lt;/span&gt;&lt;/td&gt;
&lt;td&gt;&lt;code&gt;[d_model, d_model]&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;the head's query-key map&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;span class="arithmatex"&gt;\(W_E^T W_{QK}^h W_E\)&lt;/span&gt;&lt;/td&gt;
&lt;td&gt;&lt;code&gt;[n_vocab, n_vocab]&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;route 2 &lt;strong&gt;step 1&lt;/strong&gt;: the score table&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;span class="arithmatex"&gt;\(A^h\)&lt;/span&gt;&lt;/td&gt;
&lt;td&gt;&lt;code&gt;[n_ctx, n_ctx]&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;those scores softmaxed: the percentages, carrying the mask&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;span class="arithmatex"&gt;\(W_{OV}^h = W_O^h W_V^h\)&lt;/span&gt;&lt;/td&gt;
&lt;td&gt;&lt;code&gt;[d_model, d_model]&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;the head's value-output map&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;span class="arithmatex"&gt;\(W_U W_{OV}^h W_E\)&lt;/span&gt;&lt;/td&gt;
&lt;td&gt;&lt;code&gt;[n_vocab, n_vocab]&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;route 2 &lt;strong&gt;step 2&lt;/strong&gt;: the lookup table&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;span class="arithmatex"&gt;\(\otimes\)&lt;/span&gt;&lt;/td&gt;
&lt;td&gt;&lt;/td&gt;
&lt;td&gt;glues the &lt;em&gt;which position&lt;/em&gt; part onto the &lt;em&gt;what it does&lt;/em&gt; part&lt;/td&gt;
&lt;/tr&gt;
&lt;/tbody&gt;
&lt;/table&gt;&lt;/div&gt;
&lt;p&gt;Each term splits the same way: &lt;strong&gt;a left factor saying which positions&lt;/strong&gt;, and &lt;strong&gt;a right factor saying what happens to the logits&lt;/strong&gt;.&lt;/p&gt;
&lt;p&gt;Note what is missing: &lt;span class="arithmatex"&gt;\(W_Q\)&lt;/span&gt; and &lt;span class="arithmatex"&gt;\(W_K\)&lt;/span&gt;. The formula takes &lt;span class="arithmatex"&gt;\(A^h\)&lt;/span&gt; as given, so the weights behind it are already absorbed. That is the freeze, visible in the notation.&lt;/p&gt;
&lt;h3 id="the-tensor-product-notation"&gt;The tensor-product notation&lt;/h3&gt;
&lt;p&gt;&lt;span class="arithmatex"&gt;\(A \otimes B\)&lt;/span&gt; is a factored linear map. With input &lt;span class="arithmatex"&gt;\(t\)&lt;/span&gt;, a &lt;code&gt;[n_ctx, n_vocab]&lt;/code&gt; matrix of one-hot rows, &lt;span class="arithmatex"&gt;\((A \otimes B) \cdot t = A t B^T\)&lt;/span&gt;. The factors act on different axes:&lt;/p&gt;
&lt;ul&gt;
&lt;li&gt;&lt;strong&gt;left&lt;/strong&gt; (&lt;span class="arithmatex"&gt;\(\text{Id}\)&lt;/span&gt;, &lt;span class="arithmatex"&gt;\(A^h\)&lt;/span&gt;) acts across &lt;em&gt;positions&lt;/em&gt;: which source reaches which destination&lt;/li&gt;
&lt;li&gt;&lt;strong&gt;right&lt;/strong&gt; (&lt;span class="arithmatex"&gt;\(W_U W_E\)&lt;/span&gt;, &lt;span class="arithmatex"&gt;\(W_U W_{OV}^h W_E\)&lt;/span&gt;) acts across the &lt;em&gt;vocabulary&lt;/em&gt;: what that does to the logits&lt;/li&gt;
&lt;/ul&gt;
&lt;p&gt;Neither knows about the other. That separation is the point.&lt;/p&gt;
&lt;h3 id="the-attention-pattern-a"&gt;The attention pattern &lt;span class="arithmatex"&gt;\(A\)&lt;/span&gt;&lt;/h3&gt;
&lt;p&gt;Softmax every row of the step 1 matrix, not just row 8, and you get &lt;span class="arithmatex"&gt;\(A\)&lt;/span&gt;. Same grid, same &lt;code&gt;x&lt;/code&gt;, scores as percentages.&lt;/p&gt;
&lt;div class="table-wrap"&gt;&lt;table&gt;
&lt;thead&gt;
&lt;tr&gt;
&lt;th&gt;to  from&lt;/th&gt;
&lt;th&gt;1 &lt;code&gt;the&lt;/code&gt;&lt;/th&gt;
&lt;th&gt;2 &lt;code&gt;cat&lt;/code&gt;&lt;/th&gt;
&lt;th&gt;3 &lt;code&gt;sat&lt;/code&gt;&lt;/th&gt;
&lt;th&gt;4 &lt;code&gt;on&lt;/code&gt;&lt;/th&gt;
&lt;th&gt;5 &lt;code&gt;the&lt;/code&gt;&lt;/th&gt;
&lt;th&gt;6 &lt;code&gt;mat&lt;/code&gt;&lt;/th&gt;
&lt;th&gt;7 &lt;code&gt;.&lt;/code&gt;&lt;/th&gt;
&lt;th&gt;8 &lt;code&gt;the&lt;/code&gt;&lt;/th&gt;
&lt;th&gt;9 &lt;code&gt;cat&lt;/code&gt;&lt;/th&gt;
&lt;/tr&gt;
&lt;/thead&gt;
&lt;tbody&gt;
&lt;tr&gt;
&lt;td&gt;1 &lt;code&gt;the&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;100%&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;2 &lt;code&gt;cat&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;55%&lt;/td&gt;
&lt;td&gt;45%&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;3 &lt;code&gt;sat&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;28%&lt;/td&gt;
&lt;td&gt;51%&lt;/td&gt;
&lt;td&gt;21%&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;4 &lt;code&gt;on&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;26%&lt;/td&gt;
&lt;td&gt;36%&lt;/td&gt;
&lt;td&gt;21%&lt;/td&gt;
&lt;td&gt;17%&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;5 &lt;code&gt;the&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;4%&lt;/td&gt;
&lt;td&gt;&lt;strong&gt;82%&lt;/strong&gt;&lt;/td&gt;
&lt;td&gt;5%&lt;/td&gt;
&lt;td&gt;5%&lt;/td&gt;
&lt;td&gt;4%&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;6 &lt;code&gt;mat&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;16%&lt;/td&gt;
&lt;td&gt;24%&lt;/td&gt;
&lt;td&gt;14%&lt;/td&gt;
&lt;td&gt;17%&lt;/td&gt;
&lt;td&gt;16%&lt;/td&gt;
&lt;td&gt;13%&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;7 &lt;code&gt;.&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;12%&lt;/td&gt;
&lt;td&gt;18%&lt;/td&gt;
&lt;td&gt;14%&lt;/td&gt;
&lt;td&gt;12%&lt;/td&gt;
&lt;td&gt;12%&lt;/td&gt;
&lt;td&gt;21%&lt;/td&gt;
&lt;td&gt;11%&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;strong&gt;8 &lt;code&gt;the&lt;/code&gt;&lt;/strong&gt;&lt;/td&gt;
&lt;td&gt;3%&lt;/td&gt;
&lt;td&gt;&lt;strong&gt;60%&lt;/strong&gt;&lt;/td&gt;
&lt;td&gt;4%&lt;/td&gt;
&lt;td&gt;4%&lt;/td&gt;
&lt;td&gt;3%&lt;/td&gt;
&lt;td&gt;20%&lt;/td&gt;
&lt;td&gt;3%&lt;/td&gt;
&lt;td&gt;3%&lt;/td&gt;
&lt;td&gt;x&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;9 &lt;code&gt;cat&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;11%&lt;/td&gt;
&lt;td&gt;9%&lt;/td&gt;
&lt;td&gt;13%&lt;/td&gt;
&lt;td&gt;10%&lt;/td&gt;
&lt;td&gt;11%&lt;/td&gt;
&lt;td&gt;17%&lt;/td&gt;
&lt;td&gt;9%&lt;/td&gt;
&lt;td&gt;11%&lt;/td&gt;
&lt;td&gt;9%&lt;/td&gt;
&lt;/tr&gt;
&lt;/tbody&gt;
&lt;/table&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;&lt;strong&gt;In &lt;span class="arithmatex"&gt;\(A[i,j]\)&lt;/span&gt;, &lt;span class="arithmatex"&gt;\(i\)&lt;/span&gt; is where you're standing and &lt;span class="arithmatex"&gt;\(j\)&lt;/span&gt; is where you're looking.&lt;/strong&gt; &lt;span class="arithmatex"&gt;\(A[8,2] = 60\%\)&lt;/span&gt; is position 8 on &lt;code&gt;the&lt;/code&gt; looking back at position 2 on &lt;code&gt;cat&lt;/code&gt;. Swap them and &lt;span class="arithmatex"&gt;\(A[2,8]\)&lt;/span&gt; is &lt;code&gt;x&lt;/code&gt;: position 2 cannot see the future.&lt;/li&gt;
&lt;li&gt;&lt;strong&gt;Rows sum to 100%, columns do not.&lt;/strong&gt; Each destination splits a fixed budget; a source can be read by many or none. Row 1 is 100% on itself, having nowhere else to look.&lt;/li&gt;
&lt;li&gt;&lt;strong&gt;Read down a column&lt;/strong&gt; for the other question. Who attends to &lt;code&gt;cat&lt;/code&gt; at position 2? Rows 5 and 8, at 82% and 60%, both of them a &lt;code&gt;the&lt;/code&gt;.&lt;/li&gt;
&lt;li&gt;&lt;strong&gt;Row 5 gets 82% where row 8 gets 60%&lt;/strong&gt;, same token asking the same question. Position 5 cannot see &lt;code&gt;mat&lt;/code&gt; yet, so &lt;code&gt;cat&lt;/code&gt; has no competition; by position 8 it does.&lt;/li&gt;
&lt;li&gt;&lt;strong&gt;The shape of the rows names the head.&lt;/strong&gt; A diagonal band is a previous-token head. A row that puts almost all of its percentage on position 1 means the head is parked there, reading the first token and doing nothing useful. That is how a head switches itself off.&lt;/li&gt;
&lt;/ul&gt;
&lt;p&gt;One &lt;span class="arithmatex"&gt;\(A\)&lt;/span&gt; per head, hence &lt;span class="arithmatex"&gt;\(A^h\)&lt;/span&gt;.&lt;/p&gt;
&lt;h3 id="the-two-terms-formally"&gt;The two terms, formally&lt;/h3&gt;
&lt;div class="arithmatex"&gt;\[\text{Id} \otimes W_U W_E\]&lt;/div&gt;
&lt;p&gt;Route 1. A zero-layer transformer has only this term, and training drives &lt;span class="arithmatex"&gt;\(W_U W_E\)&lt;/span&gt; to approximate bigram log-likelihoods.&lt;/p&gt;
&lt;div class="arithmatex"&gt;\[\sum_{h \in H} A^h \otimes (W_U W_{OV}^h W_E)\]&lt;/div&gt;
&lt;p&gt;Route 2. Heads don't interact in a one-layer model, which is why the terms simply add. Each factors into the two steps:&lt;/p&gt;
&lt;ul&gt;
&lt;li&gt;&lt;strong&gt;&lt;span class="arithmatex"&gt;\(W_E^T W_{QK}^h W_E\)&lt;/span&gt;, the QK circuit.&lt;/strong&gt; Step 1, where to look. Entry &lt;span class="arithmatex"&gt;\((a, b)\)&lt;/span&gt; is how much a destination holding token &lt;span class="arithmatex"&gt;\(a\)&lt;/span&gt; wants a source holding token &lt;span class="arithmatex"&gt;\(b\)&lt;/span&gt;. Step 1's matrix is this table read off at the sequence's tokens, which is why its scores depended only on the two tokens involved. Softmax them over the context and you get &lt;span class="arithmatex"&gt;\(A^h\)&lt;/span&gt;, which is the circuit's output, not the circuit itself.&lt;/li&gt;
&lt;li&gt;&lt;strong&gt;&lt;span class="arithmatex"&gt;\(W_U W_{OV}^h W_E\)&lt;/span&gt;, the OV circuit.&lt;/strong&gt; Step 2, what it does. Independent of position &lt;em&gt;and&lt;/em&gt; of the attention pattern.&lt;/li&gt;
&lt;/ul&gt;
&lt;p&gt;Both are &lt;code&gt;[n_vocab, n_vocab]&lt;/code&gt;, and together they are the head.&lt;/p&gt;
&lt;p&gt;One thing the worked example skipped: the head carries a &lt;strong&gt;value vector&lt;/strong&gt;, not the token.&lt;/p&gt;
&lt;p&gt;Each source &lt;span class="arithmatex"&gt;\(j\)&lt;/span&gt; computes &lt;span class="arithmatex"&gt;\(v_j = W_V \cdot (\text{residual stream at } j)\)&lt;/span&gt;. The head outputs the attention-weighted sum, and &lt;span class="arithmatex"&gt;\(W_O\)&lt;/span&gt; maps it back into the residual stream. &lt;span class="arithmatex"&gt;\(W_V\)&lt;/span&gt; is a projection, so only the subspace that head cares about travels. Every attended position lands at once: a blend, not a choice.&lt;/p&gt;
&lt;p&gt;&lt;span class="arithmatex"&gt;\(W_{OV}^h = W_O^h W_V^h\)&lt;/span&gt; is that round trip. It is &lt;code&gt;[d_model, d_model]&lt;/code&gt; but has rank at most &lt;code&gt;d_head&lt;/code&gt;, factoring through the head's narrow dimension. Sandwich it between &lt;span class="arithmatex"&gt;\(W_E\)&lt;/span&gt; and &lt;span class="arithmatex"&gt;\(W_U\)&lt;/span&gt; for the vocab-to-vocab lookup table.&lt;/p&gt;
&lt;h3 id="why-this-framing-is-useful"&gt;Why this framing is useful&lt;/h3&gt;
&lt;p&gt;All of this holds only with &lt;span class="arithmatex"&gt;\(A^h\)&lt;/span&gt; frozen. Given that, the decomposition is exact: behaviour reads straight off the weights as a sum of paths.&lt;/p&gt;
&lt;p&gt;The payoff: a head is interpretable from its QK and OV matrices alone, both &lt;code&gt;[n_vocab, n_vocab]&lt;/code&gt;, never by probing activations. Read them directly, or eigendecompose the OV matrix, where positive eigenvalues mark a head that copies whatever it attends to. The paper's skip-trigram example is &lt;code&gt;[perfect] ... [are] -&amp;gt; [perfect]&lt;/code&gt;.&lt;/p&gt;
&lt;p&gt;It also sets up the two-layer story. Everything above is a single attention block. Stack a second one on top and its input is the first block's output rather than the raw embeddings, so its &lt;span class="arithmatex"&gt;\(A^h\)&lt;/span&gt; is no longer a function of the tokens alone. The cross terms that creates, where the second block's queries, keys or values read what the first block wrote (Q-, K- and V-composition), are where induction heads come from.&lt;/p&gt;</content><category term="transformers"/><category term="interpretability"/><category term="transformer-circuits"/></entry><entry><title>Sampling from a Transformer</title><link href="https://weirdmachine.wtf/posts/sampling-from-a-transformer/" rel="alternate"/><published>2026-10-03T00:00:00-04:00</published><updated>2026-10-03T00:00:00-04:00</updated><author><name>Dominic Wang</name></author><id>tag:weirdmachine.wtf,2026-10-03:/posts/sampling-from-a-transformer/</id><summary type="html">&lt;p&gt;Porting the generation loop into PvML, the strategies to turn logits into a token, and the one-token bug that made a trained model emit nothing but commas.&lt;/p&gt;</summary><content type="html">&lt;p&gt;The last two posts built a transformer and trained one. In the &lt;a href="https://weirdmachine.wtf/posts/training-a-transformer/"&gt;Training a Transformer&lt;/a&gt; post, I used TransformerLens's sampler on the TinyStories weights I trained. This post ports the sampler, using the techniques from the ARENA curriculum: &lt;code&gt;greedy_search&lt;/code&gt;, &lt;code&gt;apply_temperature&lt;/code&gt;, &lt;code&gt;apply_frequency_penalty&lt;/code&gt;, &lt;code&gt;sample_basic&lt;/code&gt;, &lt;code&gt;sample_top_k&lt;/code&gt; and &lt;code&gt;sample_top_p&lt;/code&gt;, plus &lt;code&gt;beam_search&lt;/code&gt;.&lt;/p&gt;
&lt;p&gt;PvML: &lt;a href="https://github.com/d0mzw/PvML"&gt;https://github.com/d0mzw/PvML&lt;/a&gt;&lt;/p&gt;
&lt;p&gt;Disclaimer: these notes come from working through the ARENA 3.0 curriculum. I claim no credit for the original material, and this is not affiliated with or endorsed by ARENA.&lt;/p&gt;
&lt;h2 id="sampler"&gt;Sampler&lt;/h2&gt;
&lt;div class="mermaid"&gt;flowchart TB
    prompt["prompt&amp;lt;br/&amp;gt;tokens (1, posn)"]

    prompt --&amp;gt; fwd["Transformer&amp;lt;br/&amp;gt;logits (1, posn, d_vocab)"]
    fwd --&amp;gt; last["final position only&amp;lt;br/&amp;gt;logits (d_vocab,)"]

    last --&amp;gt;|"frequency_penalty"| pen["penalised&amp;lt;br/&amp;gt;(d_vocab,)"]
    pen --&amp;gt;|"/ temperature"| scaled["scaled&amp;lt;br/&amp;gt;(d_vocab,)"]
    scaled --&amp;gt;|"top_k or top_p"| filt["filtered&amp;lt;br/&amp;gt;(k,)"]

    filt --&amp;gt; draw["draw one token&amp;lt;br/&amp;gt;(1,)"]
    draw --&amp;gt;|"append, run again"| prompt&lt;/div&gt;
&lt;p&gt;The forward pass is the one from training, frozen in place. It still emits a score over all 50,257 vocabulary entries at every position, and generation needs only the last, so the rest are discarded every step.&lt;/p&gt;
&lt;div class="language-text highlight"&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;src/pvml/sampling/
    strategies.py    six functions on one logits vector
    args.py          SamplingArgs, the knobs and their checks
    sampler.py       the loop, the tokenizer, the model
    beams.py         beam search
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;h3 id="strategies"&gt;&lt;code&gt;strategies&lt;/code&gt;&lt;/h3&gt;
&lt;p&gt;Six functions on a vector of logits, excluding &lt;code&gt;beam_search&lt;/code&gt;. Four pick a token, two reshape the logits first.&lt;/p&gt;
&lt;div class="table-wrap"&gt;&lt;table&gt;
&lt;thead&gt;
&lt;tr&gt;
&lt;th&gt;Function&lt;/th&gt;
&lt;th&gt;Code&lt;/th&gt;
&lt;th&gt;What it does&lt;/th&gt;
&lt;/tr&gt;
&lt;/thead&gt;
&lt;tbody&gt;
&lt;tr&gt;
&lt;td&gt;&lt;code&gt;greedy_search&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;2 lines&lt;/td&gt;
&lt;td&gt;the argmax, no draw at all&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;code&gt;apply_temperature&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;2 lines&lt;/td&gt;
&lt;td&gt;divide before the softmax&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;code&gt;apply_frequency_penalty&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;6 lines&lt;/td&gt;
&lt;td&gt;subtract a cost per prior occurrence&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;code&gt;sample_basic&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;2 lines&lt;/td&gt;
&lt;td&gt;draw from all 50,257&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;code&gt;sample_top_k&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;4 lines&lt;/td&gt;
&lt;td&gt;keep the k best, renormalise, draw&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;code&gt;sample_top_p&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;9 lines&lt;/td&gt;
&lt;td&gt;keep the smallest set summing to p, draw&lt;/td&gt;
&lt;/tr&gt;
&lt;/tbody&gt;
&lt;/table&gt;&lt;/div&gt;
&lt;h3 id="samplingargs"&gt;&lt;code&gt;SamplingArgs&lt;/code&gt;&lt;/h3&gt;
&lt;div class="language-python highlight"&gt;&lt;span class="code-lang"&gt;python&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="nd"&gt;@dataclass&lt;/span&gt;
&lt;span class="k"&gt;class&lt;/span&gt;&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="nc"&gt;SamplingArgs&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt;
    &lt;span class="n"&gt;max_new_tokens&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt; &lt;span class="nb"&gt;int&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="mi"&gt;50&lt;/span&gt;
    &lt;span class="n"&gt;temperature&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt; &lt;span class="nb"&gt;float&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="mf"&gt;1.0&lt;/span&gt;  &lt;span class="c1"&gt;# 0 means greedy&lt;/span&gt;
    &lt;span class="n"&gt;top_k&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt; &lt;span class="nb"&gt;int&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="mi"&gt;0&lt;/span&gt;  &lt;span class="c1"&gt;# 0 disables&lt;/span&gt;
    &lt;span class="n"&gt;top_p&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt; &lt;span class="nb"&gt;float&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="mf"&gt;0.0&lt;/span&gt;  &lt;span class="c1"&gt;# 0 disables&lt;/span&gt;
    &lt;span class="n"&gt;frequency_penalty&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt; &lt;span class="nb"&gt;float&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="mf"&gt;0.0&lt;/span&gt;
    &lt;span class="n"&gt;seed&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt; &lt;span class="nb"&gt;int&lt;/span&gt; &lt;span class="o"&gt;|&lt;/span&gt; &lt;span class="kc"&gt;None&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="kc"&gt;None&lt;/span&gt;

    &lt;span class="k"&gt;def&lt;/span&gt;&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="nf"&gt;__post_init__&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt; &lt;span class="o"&gt;-&amp;gt;&lt;/span&gt; &lt;span class="kc"&gt;None&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt;
        &lt;span class="k"&gt;if&lt;/span&gt; &lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;max_new_tokens&lt;/span&gt; &lt;span class="o"&gt;&amp;lt;=&lt;/span&gt; &lt;span class="mi"&gt;0&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt;
            &lt;span class="k"&gt;raise&lt;/span&gt; &lt;span class="ne"&gt;ValueError&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="sa"&gt;f&lt;/span&gt;&lt;span class="s2"&gt;&amp;quot;max_new_tokens must be positive, got &lt;/span&gt;&lt;span class="si"&gt;{&lt;/span&gt;&lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;max_new_tokens&lt;/span&gt;&lt;span class="si"&gt;}&lt;/span&gt;&lt;span class="s2"&gt;&amp;quot;&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
        &lt;span class="k"&gt;if&lt;/span&gt; &lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;temperature&lt;/span&gt; &lt;span class="o"&gt;&amp;lt;&lt;/span&gt; &lt;span class="mi"&gt;0&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt;
            &lt;span class="k"&gt;raise&lt;/span&gt; &lt;span class="ne"&gt;ValueError&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="sa"&gt;f&lt;/span&gt;&lt;span class="s2"&gt;&amp;quot;temperature must be non-negative, got &lt;/span&gt;&lt;span class="si"&gt;{&lt;/span&gt;&lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;temperature&lt;/span&gt;&lt;span class="si"&gt;}&lt;/span&gt;&lt;span class="s2"&gt;&amp;quot;&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
        &lt;span class="k"&gt;if&lt;/span&gt; &lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;top_k&lt;/span&gt; &lt;span class="o"&gt;&amp;lt;&lt;/span&gt; &lt;span class="mi"&gt;0&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt;
            &lt;span class="k"&gt;raise&lt;/span&gt; &lt;span class="ne"&gt;ValueError&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="sa"&gt;f&lt;/span&gt;&lt;span class="s2"&gt;&amp;quot;top_k must be non-negative, got &lt;/span&gt;&lt;span class="si"&gt;{&lt;/span&gt;&lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;top_k&lt;/span&gt;&lt;span class="si"&gt;}&lt;/span&gt;&lt;span class="s2"&gt;&amp;quot;&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
        &lt;span class="k"&gt;if&lt;/span&gt; &lt;span class="ow"&gt;not&lt;/span&gt; &lt;span class="mi"&gt;0&lt;/span&gt; &lt;span class="o"&gt;&amp;lt;=&lt;/span&gt; &lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;top_p&lt;/span&gt; &lt;span class="o"&gt;&amp;lt;=&lt;/span&gt; &lt;span class="mf"&gt;1.0&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt;
            &lt;span class="k"&gt;raise&lt;/span&gt; &lt;span class="ne"&gt;ValueError&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="sa"&gt;f&lt;/span&gt;&lt;span class="s2"&gt;&amp;quot;top_p must be a probability in [0, 1], got &lt;/span&gt;&lt;span class="si"&gt;{&lt;/span&gt;&lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;top_p&lt;/span&gt;&lt;span class="si"&gt;}&lt;/span&gt;&lt;span class="s2"&gt;&amp;quot;&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
        &lt;span class="k"&gt;if&lt;/span&gt; &lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;top_k&lt;/span&gt; &lt;span class="ow"&gt;and&lt;/span&gt; &lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;top_p&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt;
            &lt;span class="k"&gt;raise&lt;/span&gt; &lt;span class="ne"&gt;ValueError&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;
                &lt;span class="sa"&gt;f&lt;/span&gt;&lt;span class="s2"&gt;&amp;quot;set at most one of top_k and top_p, got top_k=&lt;/span&gt;&lt;span class="si"&gt;{&lt;/span&gt;&lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;top_k&lt;/span&gt;&lt;span class="si"&gt;}&lt;/span&gt;&lt;span class="s2"&gt; top_p=&lt;/span&gt;&lt;span class="si"&gt;{&lt;/span&gt;&lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;top_p&lt;/span&gt;&lt;span class="si"&gt;}&lt;/span&gt;&lt;span class="s2"&gt;&amp;quot;&lt;/span&gt;
            &lt;span class="p"&gt;)&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;&lt;code&gt;top_k&lt;/code&gt; and &lt;code&gt;top_p&lt;/code&gt; are alternatives, not a pair. Setting both raises, rather than silently applying whichever the dispatcher checks first&lt;/li&gt;
&lt;li&gt;The checks run once at construction, not per token, so a bad combination fails before the first forward pass&lt;/li&gt;
&lt;li&gt;&lt;code&gt;raise&lt;/code&gt; rather than &lt;code&gt;assert&lt;/code&gt;, because &lt;code&gt;python -O&lt;/code&gt; strips asserts and the sampler treats these as a guarantee&lt;/li&gt;
&lt;/ul&gt;
&lt;h3 id="next_token"&gt;&lt;code&gt;next_token&lt;/code&gt;&lt;/h3&gt;
&lt;div class="language-python highlight"&gt;&lt;span class="code-lang"&gt;python&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="k"&gt;if&lt;/span&gt; &lt;span class="n"&gt;args&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;frequency_penalty&lt;/span&gt; &lt;span class="o"&gt;!=&lt;/span&gt; &lt;span class="mf"&gt;0.0&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt;
    &lt;span class="n"&gt;logits&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;apply_frequency_penalty&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;input_ids&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;logits&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;args&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;frequency_penalty&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;span class="k"&gt;if&lt;/span&gt; &lt;span class="n"&gt;args&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;temperature&lt;/span&gt; &lt;span class="o"&gt;==&lt;/span&gt; &lt;span class="mi"&gt;0&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt;
    &lt;span class="k"&gt;return&lt;/span&gt; &lt;span class="n"&gt;greedy_search&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;logits&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;span class="k"&gt;if&lt;/span&gt; &lt;span class="n"&gt;args&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;temperature&lt;/span&gt; &lt;span class="o"&gt;!=&lt;/span&gt; &lt;span class="mf"&gt;1.0&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt;
    &lt;span class="n"&gt;logits&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;apply_temperature&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;logits&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;args&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;temperature&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;span class="k"&gt;if&lt;/span&gt; &lt;span class="n"&gt;args&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;top_k&lt;/span&gt; &lt;span class="o"&gt;&amp;gt;&lt;/span&gt; &lt;span class="mi"&gt;0&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt;
    &lt;span class="k"&gt;return&lt;/span&gt; &lt;span class="n"&gt;sample_top_k&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;logits&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;args&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;top_k&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;span class="k"&gt;if&lt;/span&gt; &lt;span class="n"&gt;args&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;top_p&lt;/span&gt; &lt;span class="o"&gt;&amp;gt;&lt;/span&gt; &lt;span class="mf"&gt;0.0&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt;
    &lt;span class="k"&gt;return&lt;/span&gt; &lt;span class="n"&gt;sample_top_p&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;logits&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;args&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;top_p&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;span class="k"&gt;return&lt;/span&gt; &lt;span class="n"&gt;sample_basic&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;logits&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;The penalty is a correction to the raw scores so it goes first, and &lt;code&gt;top_k&lt;/code&gt; and &lt;code&gt;top_p&lt;/code&gt; read the distribution temperature produces so they come after it&lt;/li&gt;
&lt;li&gt;&lt;code&gt;greedy_search&lt;/code&gt; is checked before &lt;code&gt;apply_temperature&lt;/code&gt; because temperature 0 would divide by zero. As temperature falls towards 0 the distribution concentrates on the argmax, so &lt;code&gt;greedy_search&lt;/code&gt; is the limit it never reaches&lt;/li&gt;
&lt;li&gt;Deliberately different from ARENA, which applies temperature first and returns greedy before the penalty. Here the penalty lands first, so it applies under greedy too, and dividing by temperature afterwards scales its strength by &lt;code&gt;1/temperature&lt;/code&gt;&lt;/li&gt;
&lt;/ul&gt;
&lt;h3 id="the-loop"&gt;The loop&lt;/h3&gt;
&lt;div class="language-python highlight"&gt;&lt;span class="code-lang"&gt;python&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="k"&gt;if&lt;/span&gt; &lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;prepend_bos&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt;
    &lt;span class="n"&gt;window&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;t&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;cat&lt;/span&gt;&lt;span class="p"&gt;([&lt;/span&gt;&lt;span class="n"&gt;input_ids&lt;/span&gt;&lt;span class="p"&gt;[:&lt;/span&gt;&lt;span class="mi"&gt;1&lt;/span&gt;&lt;span class="p"&gt;],&lt;/span&gt; &lt;span class="n"&gt;input_ids&lt;/span&gt;&lt;span class="p"&gt;[&lt;/span&gt;&lt;span class="mi"&gt;1&lt;/span&gt;&lt;span class="p"&gt;:][&lt;/span&gt;&lt;span class="o"&gt;-&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;cfg&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;n_ctx&lt;/span&gt; &lt;span class="o"&gt;-&lt;/span&gt; &lt;span class="mi"&gt;1&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt; &lt;span class="p"&gt;:]])&lt;/span&gt;
&lt;span class="k"&gt;else&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt;
    &lt;span class="n"&gt;window&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;input_ids&lt;/span&gt;&lt;span class="p"&gt;[&lt;/span&gt;&lt;span class="o"&gt;-&lt;/span&gt;&lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;cfg&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;n_ctx&lt;/span&gt; &lt;span class="p"&gt;:]&lt;/span&gt;
&lt;span class="n"&gt;logits&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;model&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;window&lt;/span&gt;&lt;span class="p"&gt;[&lt;/span&gt;&lt;span class="kc"&gt;None&lt;/span&gt;&lt;span class="p"&gt;])[&lt;/span&gt;&lt;span class="mi"&gt;0&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="o"&gt;-&lt;/span&gt;&lt;span class="mi"&gt;1&lt;/span&gt;&lt;span class="p"&gt;]&lt;/span&gt;
&lt;span class="n"&gt;next_id&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;next_token&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;input_ids&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;logits&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;args&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;span class="n"&gt;input_ids&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;t&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;cat&lt;/span&gt;&lt;span class="p"&gt;([&lt;/span&gt;&lt;span class="n"&gt;input_ids&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;t&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;tensor&lt;/span&gt;&lt;span class="p"&gt;([&lt;/span&gt;&lt;span class="n"&gt;next_id&lt;/span&gt;&lt;span class="p"&gt;],&lt;/span&gt; &lt;span class="n"&gt;device&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;device&lt;/span&gt;&lt;span class="p"&gt;)],&lt;/span&gt; &lt;span class="n"&gt;dim&lt;/span&gt;&lt;span class="o"&gt;=-&lt;/span&gt;&lt;span class="mi"&gt;1&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;span class="k"&gt;if&lt;/span&gt; &lt;span class="n"&gt;next_id&lt;/span&gt; &lt;span class="o"&gt;==&lt;/span&gt; &lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;tokenizer&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;eos_token_id&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt;
    &lt;span class="k"&gt;break&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;div class="language-text highlight"&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;input_ids        (seq_len,)
window           (min(seq_len, n_ctx),)   BOS plus the last n_ctx - 1
window[None]     (1, posn)                None adds the batch axis
model(...)       (1, posn, d_vocab)
[0, -1]          (d_vocab,)               drop the batch, keep the last position
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;&lt;code&gt;input_ids&lt;/code&gt; is 1 dim, the model wants &lt;code&gt;(batch, posn)&lt;/code&gt;&lt;/li&gt;
&lt;li&gt;The window is built in two pieces so BOS stays at position 0. Slicing the whole sequence to the last &lt;code&gt;n_ctx&lt;/code&gt; would drop it the moment generation outgrows the context length&lt;/li&gt;
&lt;li&gt;&lt;code&gt;next_token&lt;/code&gt; gets the full &lt;code&gt;input_ids&lt;/code&gt;, not the window, so &lt;code&gt;apply_frequency_penalty&lt;/code&gt; counts tokens the model can no longer see&lt;/li&gt;
&lt;/ul&gt;
&lt;h3 id="the-missing-bos-token"&gt;The missing BOS token&lt;/h3&gt;
&lt;div class="language-python highlight"&gt;&lt;span class="code-lang"&gt;python&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="k"&gt;if&lt;/span&gt; &lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;prepend_bos&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt;
    &lt;span class="n"&gt;bos&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;t&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;tensor&lt;/span&gt;&lt;span class="p"&gt;([&lt;/span&gt;&lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;tokenizer&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;bos_token_id&lt;/span&gt;&lt;span class="p"&gt;],&lt;/span&gt; &lt;span class="n"&gt;device&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;device&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
    &lt;span class="n"&gt;input_ids&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;t&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;cat&lt;/span&gt;&lt;span class="p"&gt;([&lt;/span&gt;&lt;span class="n"&gt;bos&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;input_ids&lt;/span&gt;&lt;span class="p"&gt;],&lt;/span&gt; &lt;span class="n"&gt;dim&lt;/span&gt;&lt;span class="o"&gt;=-&lt;/span&gt;&lt;span class="mi"&gt;1&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;p&gt;Without these a trained TinyStories model writes this:&lt;/p&gt;
&lt;div class="language-text highlight"&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;without BOS  in  [  7454,    2402,  257,     640]
                 [&amp;#39;Once&amp;#39;, &amp;#39; upon&amp;#39;, &amp;#39; a&amp;#39;, &amp;#39; time&amp;#39;]
             out &amp;#39;Once upon a time,,,,,,,,,,,,,,,,,,,,,,,,,,,,,,&amp;#39;

with BOS     in  [          50256,   7454,    2402,  257,     640]
                 [&amp;#39;&amp;lt;|endoftext|&amp;gt;&amp;#39;, &amp;#39;Once&amp;#39;, &amp;#39; upon&amp;#39;, &amp;#39; a&amp;#39;, &amp;#39; time&amp;#39;]
             out &amp;#39;Once upon a time, there was a little girl named Lily. She loved to
                  play outside in the sunshine. One day, she went to the park&amp;#39;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;&lt;code&gt;tokenize_and_concatenate(..., add_bos_token=True)&lt;/code&gt; puts &lt;code&gt;50256&lt;/code&gt; at the start of every training row, so the model never saw anything else at position 0&lt;/li&gt;
&lt;li&gt;Prepending it once is not enough. Slicing the whole sequence to the last &lt;code&gt;n_ctx&lt;/code&gt; drops BOS again the moment generation outgrows the window, so the window is built as BOS plus the last &lt;code&gt;n_ctx - 1&lt;/code&gt; tokens. That also matches training, where &lt;code&gt;seq_len = max_length - 1&lt;/code&gt; leaves room for it&lt;/li&gt;
&lt;/ul&gt;
&lt;h2 id="beam-search"&gt;Beam search&lt;/h2&gt;
&lt;div class="mermaid"&gt;flowchart TB
    beams["Beams&amp;lt;br/&amp;gt;logprob_sums (k,)&amp;lt;br/&amp;gt;tokens (k, seq)"]

    beams --&amp;gt; fwd["Transformer&amp;lt;br/&amp;gt;logits (k, seq, d_vocab)"]
    fwd --&amp;gt; lp["log_softmax at the last position&amp;lt;br/&amp;gt;(k, d_vocab)"]

    lp --&amp;gt;|"topk per beam"| cand["candidates&amp;lt;br/&amp;gt;(k, k)"]
    cand --&amp;gt;|"flatten"| all["k x k beams&amp;lt;br/&amp;gt;sums added, tokens appended"]

    all --&amp;gt;|"filter: top k"| beams
    all --&amp;gt;|"ended on eos"| done["finished"]&lt;/div&gt;
&lt;div class="language-python highlight"&gt;&lt;span class="code-lang"&gt;python&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="nd"&gt;@dataclass&lt;/span&gt;
&lt;span class="k"&gt;class&lt;/span&gt;&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="nc"&gt;Beams&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt;
    &lt;span class="n"&gt;model&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt; &lt;span class="n"&gt;t&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;nn&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;Module&lt;/span&gt;
    &lt;span class="n"&gt;tokenizer&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt; &lt;span class="nb"&gt;object&lt;/span&gt;
    &lt;span class="n"&gt;logprob_sums&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt; &lt;span class="n"&gt;Float&lt;/span&gt;&lt;span class="p"&gt;[&lt;/span&gt;&lt;span class="n"&gt;Tensor&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="s2"&gt;&amp;quot;batch&amp;quot;&lt;/span&gt;&lt;span class="p"&gt;]&lt;/span&gt;
    &lt;span class="n"&gt;tokens&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt; &lt;span class="n"&gt;Int&lt;/span&gt;&lt;span class="p"&gt;[&lt;/span&gt;&lt;span class="n"&gt;Tensor&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="s2"&gt;&amp;quot;batch seq&amp;quot;&lt;/span&gt;&lt;span class="p"&gt;]&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;p&gt;Every other strategy collapses the distribution to one id and moves on. Beam search carries k forward, so that &lt;code&gt;batch&lt;/code&gt; axis is the beams.&lt;/p&gt;
&lt;div class="language-python highlight"&gt;&lt;span class="code-lang"&gt;python&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="n"&gt;new_sums&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;einops&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;repeat&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;logprob_sums&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="s2"&gt;&amp;quot;b -&amp;gt; b k&amp;quot;&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;k&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="n"&gt;k&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt; &lt;span class="o"&gt;+&lt;/span&gt; &lt;span class="n"&gt;topk_logprobs&lt;/span&gt;
&lt;span class="n"&gt;new_tokens&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;t&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;concat&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;
    &lt;span class="p"&gt;[&lt;/span&gt;&lt;span class="n"&gt;einops&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;repeat&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;tokens&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="s2"&gt;&amp;quot;b s -&amp;gt; b k s&amp;quot;&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;k&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="n"&gt;k&lt;/span&gt;&lt;span class="p"&gt;),&lt;/span&gt; &lt;span class="n"&gt;topk_ids&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;unsqueeze&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="o"&gt;-&lt;/span&gt;&lt;span class="mi"&gt;1&lt;/span&gt;&lt;span class="p"&gt;)],&lt;/span&gt; &lt;span class="n"&gt;dim&lt;/span&gt;&lt;span class="o"&gt;=-&lt;/span&gt;&lt;span class="mi"&gt;1&lt;/span&gt;
&lt;span class="p"&gt;)&lt;/span&gt;
&lt;span class="k"&gt;return&lt;/span&gt; &lt;span class="n"&gt;Beams&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;model&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;tokenizer&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;new_sums&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;flatten&lt;/span&gt;&lt;span class="p"&gt;(),&lt;/span&gt; &lt;span class="n"&gt;new_tokens&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;flatten&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="mi"&gt;0&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="mi"&gt;1&lt;/span&gt;&lt;span class="p"&gt;))&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;p&gt;The shapes through two rounds at &lt;code&gt;k=3&lt;/code&gt;:&lt;/p&gt;
&lt;div class="language-text highlight"&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;start        logprob_sums (1,)   tokens (1, 4)
generate(3)  logprob_sums (3,)   tokens (3, 5)
filter(3)    keep (3, 5)         terminated (0, 5)   none ended yet
generate(3)  logprob_sums (9,)   tokens (9, 6)
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;&lt;code&gt;flatten&lt;/code&gt; makes the search global: &lt;code&gt;filter&lt;/code&gt; ranks a 1 dim tensor, so all k x k candidates compete together&lt;/li&gt;
&lt;/ul&gt;
&lt;div class="table-wrap"&gt;&lt;table&gt;
&lt;thead&gt;
&lt;tr&gt;
&lt;th&gt;num_beams&lt;/th&gt;
&lt;th&gt;best score&lt;/th&gt;
&lt;th&gt;completion&lt;/th&gt;
&lt;/tr&gt;
&lt;/thead&gt;
&lt;tbody&gt;
&lt;tr&gt;
&lt;td&gt;1&lt;/td&gt;
&lt;td&gt;-18.076&lt;/td&gt;
&lt;td&gt;&lt;code&gt;When I was a kid, I was a little bit of a nerd.&lt;/code&gt;&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;3&lt;/td&gt;
&lt;td&gt;-14.830&lt;/td&gt;
&lt;td&gt;&lt;code&gt;When I was a kid, I used to go to the movies with my&lt;/code&gt;&lt;/td&gt;
&lt;/tr&gt;
&lt;/tbody&gt;
&lt;/table&gt;&lt;/div&gt;
&lt;h3 id="module-output"&gt;Module output&lt;/h3&gt;
&lt;div class="language-bash highlight"&gt;&lt;span class="code-lang"&gt;shell&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="o"&gt;(&lt;/span&gt;.venv&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt; &lt;/span&gt;dom@dom-z13:~/Desktop/PvML$&lt;span class="w"&gt; &lt;/span&gt;python&lt;span class="w"&gt; &lt;/span&gt;-m&lt;span class="w"&gt; &lt;/span&gt;pvml.sampling.strategies
logits&lt;span class="w"&gt;  &lt;/span&gt;&lt;span class="o"&gt;[&lt;/span&gt;&lt;span class="m"&gt;3&lt;/span&gt;.0,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;2&lt;/span&gt;.0,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;.0,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.0,&lt;span class="w"&gt; &lt;/span&gt;-1.0&lt;span class="o"&gt;]&lt;/span&gt;
probs&lt;span class="w"&gt;   &lt;/span&gt;&lt;span class="o"&gt;[&lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;0.636&amp;#39;&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;0.234&amp;#39;&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;0.086&amp;#39;&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;0.032&amp;#39;&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;0.012&amp;#39;&lt;/span&gt;&lt;span class="o"&gt;]&lt;/span&gt;
cumulative&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;[&lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;0.636&amp;#39;&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;0.871&amp;#39;&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;0.957&amp;#39;&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;0.988&amp;#39;&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;1.000&amp;#39;&lt;/span&gt;&lt;span class="o"&gt;]&lt;/span&gt;

greedy_search&lt;span class="w"&gt;              &lt;/span&gt;-&amp;gt;&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;&lt;span class="w"&gt;   &lt;/span&gt;&lt;span class="c1"&gt;# the argmax&lt;/span&gt;
apply_temperature&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.5&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;     &lt;/span&gt;-&amp;gt;&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;[&lt;/span&gt;&lt;span class="m"&gt;6&lt;/span&gt;.0,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;4&lt;/span&gt;.0,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;2&lt;/span&gt;.0,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.0,&lt;span class="w"&gt; &lt;/span&gt;-2.0&lt;span class="o"&gt;]&lt;/span&gt;
apply_temperature&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;2&lt;/span&gt;.0&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;     &lt;/span&gt;-&amp;gt;&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;[&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;.5,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;.0,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.5,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.0,&lt;span class="w"&gt; &lt;/span&gt;-0.5&lt;span class="o"&gt;]&lt;/span&gt;

sample_basic&lt;span class="w"&gt; &lt;/span&gt;over&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;10&lt;/span&gt;,000&lt;span class="w"&gt; &lt;/span&gt;draws
&lt;span class="w"&gt;  &lt;/span&gt;&lt;span class="o"&gt;[&lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;0.640&amp;#39;&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;0.230&amp;#39;&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;0.086&amp;#39;&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;0.032&amp;#39;&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;0.011&amp;#39;&lt;/span&gt;&lt;span class="o"&gt;]&lt;/span&gt;

sample_top_k&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;2&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt; &lt;/span&gt;over&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;10&lt;/span&gt;,000&lt;span class="w"&gt; &lt;/span&gt;draws
&lt;span class="w"&gt;  &lt;/span&gt;&lt;span class="o"&gt;[&lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;0.736&amp;#39;&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;0.264&amp;#39;&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;0.000&amp;#39;&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;0.000&amp;#39;&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;0.000&amp;#39;&lt;/span&gt;&lt;span class="o"&gt;]&lt;/span&gt;

sample_top_p&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.7&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt; &lt;/span&gt;over&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;10&lt;/span&gt;,000&lt;span class="w"&gt; &lt;/span&gt;draws
&lt;span class="w"&gt;  &lt;/span&gt;&lt;span class="o"&gt;[&lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;0.734&amp;#39;&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;0.267&amp;#39;&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;0.000&amp;#39;&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;0.000&amp;#39;&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;0.000&amp;#39;&lt;/span&gt;&lt;span class="o"&gt;]&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;&lt;code&gt;sample_basic&lt;/code&gt; reproduces the true distribution to three decimals&lt;/li&gt;
&lt;li&gt;&lt;code&gt;top_k(2)&lt;/code&gt; and &lt;code&gt;top_p(0.7)&lt;/code&gt; agree because the cumulative probability reaches 0.7 at the second token, so both keep the same pair. Renormalised, 0.636 of 0.870 is 0.731 against the measured 0.736&lt;/li&gt;
&lt;/ul&gt;
&lt;div class="language-bash highlight"&gt;&lt;span class="code-lang"&gt;shell&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="o"&gt;(&lt;/span&gt;.venv&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt; &lt;/span&gt;dom@dom-z13:~/Desktop/PvML$&lt;span class="w"&gt; &lt;/span&gt;python&lt;span class="w"&gt; &lt;/span&gt;-m&lt;span class="w"&gt; &lt;/span&gt;pvml.sampling.sampler
prompt&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;Jingle bells, jingle bells, jingle all the way&amp;#39;&lt;/span&gt;
ours&lt;span class="w"&gt;      &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;Jingle bells, jingle bells, jingle all the way down to the top of the mountain.&amp;#39;&lt;/span&gt;
reference&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;Jingle bells, jingle bells, jingle all the way down to the top of the mountain.&amp;#39;&lt;/span&gt;

matches&lt;span class="w"&gt; &lt;/span&gt;transformer_lens&lt;span class="w"&gt; &lt;/span&gt;:&lt;span class="w"&gt; &lt;/span&gt;True
matches&lt;span class="w"&gt; &lt;/span&gt;ARENA&lt;span class="err"&gt;&amp;#39;&lt;/span&gt;s&lt;span class="w"&gt; &lt;/span&gt;expected&lt;span class="w"&gt; &lt;/span&gt;:&lt;span class="w"&gt; &lt;/span&gt;True
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;Greedy is the only setting of the stochastic strategies with no draw in it, so it is the only one checkable exactly&lt;/li&gt;
&lt;/ul&gt;
&lt;div class="language-bash highlight"&gt;&lt;span class="code-lang"&gt;shell&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="o"&gt;(&lt;/span&gt;.venv&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt; &lt;/span&gt;dom@dom-z13:~/Desktop/PvML$&lt;span class="w"&gt; &lt;/span&gt;python&lt;span class="w"&gt; &lt;/span&gt;-m&lt;span class="w"&gt; &lt;/span&gt;pvml.sampling.beams
prompt&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;When I was&amp;#39;&lt;/span&gt;

&lt;span class="nv"&gt;num_beams&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="m"&gt;3&lt;/span&gt;
&lt;span class="w"&gt;   &lt;/span&gt;-14.830&lt;span class="w"&gt;  &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;When I was a kid, I used to go to the movies with my&amp;#39;&lt;/span&gt;
&lt;span class="w"&gt;   &lt;/span&gt;-16.321&lt;span class="w"&gt;  &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;When I was a kid, I used to go to the movies. I&amp;#39;&lt;/span&gt;
&lt;span class="w"&gt;   &lt;/span&gt;-16.899&lt;span class="w"&gt;  &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;When I was a kid, I used to go to the movies and watch&amp;#39;&lt;/span&gt;

&lt;span class="nv"&gt;no_repeat_ngram_size&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="m"&gt;3&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;25&lt;/span&gt;&lt;span class="w"&gt; &lt;/span&gt;tokens
&lt;span class="w"&gt;   &lt;/span&gt;-33.603&lt;span class="w"&gt;  &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;When I was a kid, I used to go to the movies with my dad. He was a big movie star, but he was also&amp;#39;&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;Beam search is deterministic, so it gets the exact check the rest of the sampler cannot. All three sequences match &lt;code&gt;GPT2LMHeadModel.generate(num_beams=3, num_return_sequences=3)&lt;/code&gt; token for token&lt;/li&gt;
&lt;/ul&gt;
&lt;h3 id="what-runs-on-sampler"&gt;What runs on &lt;code&gt;Sampler&lt;/code&gt;&lt;/h3&gt;
&lt;p&gt;&lt;code&gt;experiments/sample.py&lt;/code&gt;, the REPL:&lt;/p&gt;
&lt;div class="language-python highlight"&gt;&lt;span class="code-lang"&gt;python&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="c1"&gt;# before&lt;/span&gt;
&lt;span class="nb"&gt;print&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="sa"&gt;f&lt;/span&gt;&lt;span class="s2"&gt;&amp;quot;&lt;/span&gt;&lt;span class="se"&gt;\n&lt;/span&gt;&lt;span class="si"&gt;{&lt;/span&gt;&lt;span class="n"&gt;hooked&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;generate&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;line&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt;&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="n"&gt;verbose&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="kc"&gt;False&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt;&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;**&lt;/span&gt;&lt;span class="n"&gt;sampler&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;kwargs&lt;/span&gt;&lt;span class="p"&gt;())&lt;/span&gt;&lt;span class="si"&gt;}&lt;/span&gt;&lt;span class="se"&gt;\n&lt;/span&gt;&lt;span class="s2"&gt;&amp;quot;&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;

&lt;span class="c1"&gt;# after&lt;/span&gt;
&lt;span class="nb"&gt;print&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="sa"&gt;f&lt;/span&gt;&lt;span class="s2"&gt;&amp;quot;&lt;/span&gt;&lt;span class="se"&gt;\n&lt;/span&gt;&lt;span class="si"&gt;{&lt;/span&gt;&lt;span class="n"&gt;sampler&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;sample&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;line&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt;&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="n"&gt;args&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;&lt;span class="si"&gt;}&lt;/span&gt;&lt;span class="se"&gt;\n&lt;/span&gt;&lt;span class="s2"&gt;&amp;quot;&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;p&gt;&lt;code&gt;experiments/_run.py&lt;/code&gt;, the sample printed after every training epoch:&lt;/p&gt;
&lt;div class="language-python highlight"&gt;&lt;span class="code-lang"&gt;python&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="c1"&gt;# before&lt;/span&gt;
&lt;span class="k"&gt;return&lt;/span&gt; &lt;span class="n"&gt;hooked&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;generate&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;prompt&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;max_new_tokens&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="mi"&gt;50&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;temperature&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="mf"&gt;0.7&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;top_p&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="mf"&gt;0.95&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;verbose&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="kc"&gt;False&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;

&lt;span class="c1"&gt;# after&lt;/span&gt;
&lt;span class="k"&gt;return&lt;/span&gt; &lt;span class="n"&gt;Sampler&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;model&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;tokenizer&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;sample&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;prompt&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;args&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;p&gt;&lt;code&gt;src/pvml/reference/gpt2.py&lt;/code&gt; is now the only file in the repo that imports &lt;code&gt;HookedTransformer&lt;/code&gt;, and it only does so to load GPT-2 for the parity checks.&lt;/p&gt;
&lt;h3 id="whats-next"&gt;What's next&lt;/h3&gt;
&lt;ul&gt;
&lt;li&gt;Skipping the KV cache section in the curriculum and moving on to mechanistic interpretability, which is what I came for&lt;/li&gt;
&lt;/ul&gt;</content><category term="transformers"/><category term="arena"/><category term="samplers"/></entry><entry><title>Training a Transformer</title><link href="https://weirdmachine.wtf/posts/training-a-transformer/" rel="alternate"/><published>2026-09-29T00:00:00-04:00</published><updated>2026-09-29T00:00:00-04:00</updated><author><name>Dominic Wang</name></author><id>tag:weirdmachine.wtf,2026-09-29:/posts/training-a-transformer/</id><summary type="html">&lt;p&gt;Continuation of the Understanding Transformers post, where the weights stop being downloaded and start being learned. A loss, a data pipeline, a trainer, and experiments.&lt;/p&gt;</summary><content type="html">&lt;p&gt;A continuation of &lt;a href="https://weirdmachine.wtf/posts/understanding-transformers/"&gt;Understanding Transformers&lt;/a&gt;, where I ported ARENA's GPT-2 implementation into &lt;code&gt;PvML&lt;/code&gt; one module at a time. To check each module I loaded GPT-2's own weights into it and compared the output against the reference, which is how I know the implementation is right.&lt;/p&gt;
&lt;p&gt;But every weight in it was downloaded. Nothing had been learned.&lt;/p&gt;
&lt;p&gt;This post is the other half: a loss, a dataset, and a training loop, then a model of my own trained on TinyStories.&lt;/p&gt;
&lt;p&gt;PvML: &lt;a href="https://github.com/d0mzw/PvML"&gt;https://github.com/d0mzw/PvML&lt;/a&gt;
Weights: &lt;a href="https://huggingface.co/d0mzw/pvml-tinystories"&gt;https://huggingface.co/d0mzw/pvml-tinystories&lt;/a&gt;&lt;/p&gt;
&lt;p&gt;Disclaimer: these notes come from working through the ARENA 3.0 curriculum. I claim no credit for the original material, and this is not affiliated with or endorsed by ARENA.&lt;/p&gt;
&lt;h2 id="the-training-loop"&gt;The training loop&lt;/h2&gt;
&lt;div class="mermaid"&gt;flowchart TB
    data["TinyStories&amp;lt;br/&amp;gt;tokenized, chunked to n_ctx"]
    data --&amp;gt; loader["DataLoader&amp;lt;br/&amp;gt;batch (batch, posn)"]

    loader --&amp;gt; model["Transformer&amp;lt;br/&amp;gt;logits (batch, posn, d_vocab)"]

    model --&amp;gt; loss["get_log_probs&amp;lt;br/&amp;gt;-mean() -&amp;gt; loss"]
    loader --&amp;gt; loss

    loss --&amp;gt; back["loss.backward()&amp;lt;br/&amp;gt;gradients on every parameter"]
    back --&amp;gt; step["optimizer.step()&amp;lt;br/&amp;gt;AdamW"]
    step --&amp;gt;|"next batch"| loader&lt;/div&gt;
&lt;p&gt;The model is the same one from the last post. Everything here is the machinery around it: where the batches come from, how a prediction is graded, and what moves the weights.&lt;/p&gt;
&lt;h2 id="loss"&gt;Loss&lt;/h2&gt;
&lt;div class="mermaid"&gt;flowchart TB
    logits["logits&amp;lt;br/&amp;gt;(batch, posn, d_vocab)"]
    tokens["tokens&amp;lt;br/&amp;gt;(batch, posn)"]

    logits --&amp;gt;|"log_softmax(-1)"| lp["log_probs&amp;lt;br/&amp;gt;(batch, posn, d_vocab)"]
    lp --&amp;gt;|"[:, :-1]&amp;lt;br/&amp;gt;drop the last position"| preds["predictions&amp;lt;br/&amp;gt;(batch, posn-1, d_vocab)"]
    tokens --&amp;gt;|"[:, 1:]&amp;lt;br/&amp;gt;drop the first token"| targets["targets&amp;lt;br/&amp;gt;(batch, posn-1)"]

    preds --&amp;gt; gather["gather along d_vocab"]
    targets --&amp;gt; gather
    gather --&amp;gt; out["log_probs&amp;lt;br/&amp;gt;(batch, posn-1)"]&lt;/div&gt;
&lt;p&gt;How much probability the model gave to the token that actually came next, at every position.&lt;/p&gt;
&lt;h3 id="get_log_probs"&gt;&lt;code&gt;get_log_probs&lt;/code&gt;&lt;/h3&gt;
&lt;div class="language-python highlight"&gt;&lt;span class="code-lang"&gt;python&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="n"&gt;log_probs&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;logits&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;log_softmax&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;dim&lt;/span&gt;&lt;span class="o"&gt;=-&lt;/span&gt;&lt;span class="mi"&gt;1&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;

&lt;span class="k"&gt;return&lt;/span&gt; &lt;span class="n"&gt;log_probs&lt;/span&gt;&lt;span class="p"&gt;[:,&lt;/span&gt; &lt;span class="p"&gt;:&lt;/span&gt;&lt;span class="o"&gt;-&lt;/span&gt;&lt;span class="mi"&gt;1&lt;/span&gt;&lt;span class="p"&gt;]&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;gather&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;dim&lt;/span&gt;&lt;span class="o"&gt;=-&lt;/span&gt;&lt;span class="mi"&gt;1&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;index&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="n"&gt;tokens&lt;/span&gt;&lt;span class="p"&gt;[:,&lt;/span&gt; &lt;span class="mi"&gt;1&lt;/span&gt;&lt;span class="p"&gt;:]&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;unsqueeze&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="o"&gt;-&lt;/span&gt;&lt;span class="mi"&gt;1&lt;/span&gt;&lt;span class="p"&gt;))&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;squeeze&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="o"&gt;-&lt;/span&gt;&lt;span class="mi"&gt;1&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;
&lt;p&gt;&lt;code&gt;logits&lt;/code&gt; is &lt;code&gt;(batch, posn, d_vocab)&lt;/code&gt;, a raw score for every token in the vocabulary at every position, which &lt;code&gt;log_softmax&lt;/code&gt; turns into log-probabilities.&lt;/p&gt;
&lt;ul&gt;
&lt;li&gt;Our sentence &lt;code&gt;"The cat sat on the mat"&lt;/code&gt; is 7 tokens because of the prepended BOS, so &lt;code&gt;(1, 7, 50257)&lt;/code&gt;&lt;/li&gt;
&lt;/ul&gt;
&lt;/li&gt;
&lt;li&gt;
&lt;p&gt;&lt;code&gt;log_probs[:, :-1]&lt;/code&gt; drops the last position, which predicts a token that is not in the sequence, so there is no answer to grade it against. &lt;code&gt;(1, 7, 50257)&lt;/code&gt; -&amp;gt; &lt;code&gt;(1, 6, 50257)&lt;/code&gt;&lt;/p&gt;
&lt;/li&gt;
&lt;/ul&gt;
&lt;div class="language-text highlight"&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;position:   0      1      2      3      4      5      6
token:    &amp;lt;BOS&amp;gt;   The    cat    sat     on    the    mat
predicts:  The    cat    sat     on    the    mat     ?
                                                      └── nothing in the sequence to compare against
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;
&lt;p&gt;&lt;code&gt;tokens&lt;/code&gt; is &lt;code&gt;(batch, posn)&lt;/code&gt;, the ids the text was tokenized into. Every token except &lt;code&gt;&amp;lt;BOS&amp;gt;&lt;/code&gt; is also an answer: the one the position before it should have predicted.&lt;/p&gt;
&lt;/li&gt;
&lt;li&gt;
&lt;p&gt;&lt;code&gt;tokens[:, 1:]&lt;/code&gt; is the sequence shifted left by one, which is exactly that list of answers. &lt;code&gt;&amp;lt;BOS&amp;gt;&lt;/code&gt; drops out because nothing predicts it. &lt;code&gt;(1, 7)&lt;/code&gt; -&amp;gt; &lt;code&gt;(1, 6)&lt;/code&gt;&lt;/p&gt;
&lt;/li&gt;
&lt;li&gt;
&lt;p&gt;&lt;code&gt;.gather(dim=-1, index=...)&lt;/code&gt; uses each answer token as an index into that position's 50257 scores, and returns the one it lands on. &lt;code&gt;(1, 6, 50257)&lt;/code&gt; -&amp;gt; &lt;code&gt;(1, 6, 1)&lt;/code&gt;&lt;/p&gt;
&lt;ul&gt;
&lt;li&gt;For example, at position 0 (&lt;code&gt;&amp;lt;BOS&amp;gt;&lt;/code&gt;), shifting left by one gives the answer token &lt;code&gt;The&lt;/code&gt;, and &lt;code&gt;gather&lt;/code&gt; returns the log-probability &lt;code&gt;log_probs&lt;/code&gt; assigned to &lt;code&gt;The&lt;/code&gt; at that position.&lt;/li&gt;
&lt;/ul&gt;
&lt;/li&gt;
&lt;li&gt;
&lt;p&gt;&lt;code&gt;.unsqueeze(-1)&lt;/code&gt; adds an axis of length 1 on the end, &lt;code&gt;.squeeze(-1)&lt;/code&gt; takes it away again: &lt;code&gt;(1, 6)&lt;/code&gt; -&amp;gt; &lt;code&gt;(1, 6, 1)&lt;/code&gt; going in, &lt;code&gt;(1, 6, 1)&lt;/code&gt; -&amp;gt; &lt;code&gt;(1, 6)&lt;/code&gt; coming out. They are only there because &lt;code&gt;gather&lt;/code&gt; needs the index to have the same number of axes as the thing it indexes into. What is left is one log-probability per graded prediction.&lt;/p&gt;
&lt;/li&gt;
&lt;/ul&gt;
&lt;h3 id="module-output"&gt;Module output&lt;/h3&gt;
&lt;div class="language-bash highlight"&gt;&lt;span class="code-lang"&gt;shell&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="o"&gt;(&lt;/span&gt;.venv&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt; &lt;/span&gt;dom@dom-z13:~/Desktop/PvML$&lt;span class="w"&gt; &lt;/span&gt;python&lt;span class="w"&gt; &lt;/span&gt;-m&lt;span class="w"&gt; &lt;/span&gt;pvml.training.losses
logits&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;50257&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;
log_probs&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;6&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;   &lt;/span&gt;&lt;span class="c1"&gt;# one per prediction, not per token&lt;/span&gt;

cross&lt;span class="w"&gt; &lt;/span&gt;entropy,&lt;span class="w"&gt; &lt;/span&gt;trained&lt;span class="w"&gt; &lt;/span&gt;model&lt;span class="w"&gt; &lt;/span&gt;:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;4&lt;/span&gt;.578&lt;span class="w"&gt; &lt;/span&gt;nats
cross&lt;span class="w"&gt; &lt;/span&gt;entropy,&lt;span class="w"&gt; &lt;/span&gt;uniform&lt;span class="w"&gt; &lt;/span&gt;guess&lt;span class="w"&gt; &lt;/span&gt;:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;10&lt;/span&gt;.825&lt;span class="w"&gt; &lt;/span&gt;nats
mean&lt;span class="w"&gt; &lt;/span&gt;probability&lt;span class="w"&gt; &lt;/span&gt;of&lt;span class="w"&gt; &lt;/span&gt;the&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="nb"&gt;true&lt;/span&gt;&lt;span class="w"&gt; &lt;/span&gt;next&lt;span class="w"&gt; &lt;/span&gt;token:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.114

per&lt;span class="w"&gt; &lt;/span&gt;prediction
&lt;span class="w"&gt;     &lt;/span&gt;&amp;lt;BOS&amp;gt;&lt;span class="w"&gt; &lt;/span&gt;-&amp;gt;&lt;span class="w"&gt; &lt;/span&gt;The&lt;span class="w"&gt;      &lt;/span&gt;logprob&lt;span class="w"&gt;  &lt;/span&gt;-3.278&lt;span class="w"&gt;   &lt;/span&gt;p&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.0377
&lt;span class="w"&gt;       &lt;/span&gt;The&lt;span class="w"&gt; &lt;/span&gt;-&amp;gt;&lt;span class="w"&gt;  &lt;/span&gt;cat&lt;span class="w"&gt;     &lt;/span&gt;logprob&lt;span class="w"&gt;  &lt;/span&gt;-9.222&lt;span class="w"&gt;   &lt;/span&gt;p&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.0001
&lt;span class="w"&gt;       &lt;/span&gt;cat&lt;span class="w"&gt; &lt;/span&gt;-&amp;gt;&lt;span class="w"&gt;  &lt;/span&gt;sat&lt;span class="w"&gt;     &lt;/span&gt;logprob&lt;span class="w"&gt;  &lt;/span&gt;-6.844&lt;span class="w"&gt;   &lt;/span&gt;p&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.0011
&lt;span class="w"&gt;       &lt;/span&gt;sat&lt;span class="w"&gt; &lt;/span&gt;-&amp;gt;&lt;span class="w"&gt;  &lt;/span&gt;on&lt;span class="w"&gt;      &lt;/span&gt;logprob&lt;span class="w"&gt;  &lt;/span&gt;-1.493&lt;span class="w"&gt;   &lt;/span&gt;p&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.2248
&lt;span class="w"&gt;        &lt;/span&gt;on&lt;span class="w"&gt; &lt;/span&gt;-&amp;gt;&lt;span class="w"&gt;  &lt;/span&gt;the&lt;span class="w"&gt;     &lt;/span&gt;logprob&lt;span class="w"&gt;  &lt;/span&gt;-0.870&lt;span class="w"&gt;   &lt;/span&gt;p&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.4191
&lt;span class="w"&gt;       &lt;/span&gt;the&lt;span class="w"&gt; &lt;/span&gt;-&amp;gt;&lt;span class="w"&gt;  &lt;/span&gt;mat&lt;span class="w"&gt;     &lt;/span&gt;logprob&lt;span class="w"&gt;  &lt;/span&gt;-5.761&lt;span class="w"&gt;   &lt;/span&gt;p&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.0031
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;7 tokens give 6 predictions. The last position has no next token to be graded on&lt;/li&gt;
&lt;li&gt;&lt;code&gt;log(d_vocab) = 10.825&lt;/code&gt; nats is what uniform guessing scores. GPT-2 gets 4.578 on this sentence, and the gap is what it knows&lt;/li&gt;
&lt;li&gt;Some predictions are nearly determined by context (&lt;code&gt;on -&amp;gt; the&lt;/code&gt;, p 0.42), others not at all (&lt;code&gt;The -&amp;gt; cat&lt;/code&gt;, p 0.0001). Average loss is always a mix of the two, which is why it never approaches zero&lt;/li&gt;
&lt;/ul&gt;
&lt;h2 id="data"&gt;Data&lt;/h2&gt;
&lt;div class="mermaid"&gt;flowchart TB
    raw["TinyStories&amp;lt;br/&amp;gt;2,119,719 stories"]

    raw --&amp;gt;|"tokenize_and_concatenate"| stream["one token stream&amp;lt;br/&amp;gt;glued end to end"]
    stream --&amp;gt;|"chop every n_ctx"| chunks["3,731,146 chunks&amp;lt;br/&amp;gt;(n_ctx,) each"]

    chunks --&amp;gt;|"train_test_split"| train["train&amp;lt;br/&amp;gt;3,730,146"]
    chunks --&amp;gt; test["test&amp;lt;br/&amp;gt;1,000"]

    train --&amp;gt; loader["DataLoader&amp;lt;br/&amp;gt;(batch, posn)"]
    test --&amp;gt; loader2["DataLoader&amp;lt;br/&amp;gt;(batch, posn)"]&lt;/div&gt;
&lt;p&gt;A training example is 128 contiguous tokens, not a story.&lt;/p&gt;
&lt;h3 id="tinystories_loaders"&gt;&lt;code&gt;tinystories_loaders&lt;/code&gt;&lt;/h3&gt;
&lt;div class="language-python highlight"&gt;&lt;span class="code-lang"&gt;python&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="n"&gt;tokenized&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;tokenize_and_concatenate&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;
    &lt;span class="n"&gt;dataset&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;tokenizer&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;streaming&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="kc"&gt;False&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt;
    &lt;span class="n"&gt;max_length&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="n"&gt;cfg&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;n_ctx&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;column_name&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="s2"&gt;&amp;quot;text&amp;quot;&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt;
    &lt;span class="n"&gt;add_bos_token&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="kc"&gt;True&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;num_proc&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="mi"&gt;8&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt;
&lt;span class="p"&gt;)&lt;/span&gt;

&lt;span class="n"&gt;split&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;tokenized&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;train_test_split&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;test_size&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="n"&gt;test_size&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;seed&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="n"&gt;seed&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;&lt;code&gt;train_test_split&lt;/code&gt; shuffles first, so the held-out 1000 are not all from the tail of the corpus. &lt;code&gt;seed=&lt;/code&gt; keeps the split the same across runs&lt;/li&gt;
&lt;li&gt;The loaders are returned rather than left as a module global, so the trainer takes its data as an argument&lt;/li&gt;
&lt;/ul&gt;
&lt;h3 id="module-output_1"&gt;Module output&lt;/h3&gt;
&lt;div class="language-bash highlight"&gt;&lt;span class="code-lang"&gt;shell&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="o"&gt;(&lt;/span&gt;.venv&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt; &lt;/span&gt;dom@dom-z13:~/Desktop/PvML$&lt;span class="w"&gt; &lt;/span&gt;python&lt;span class="w"&gt; &lt;/span&gt;-m&lt;span class="w"&gt; &lt;/span&gt;pvml.data.tinystories
train&lt;span class="w"&gt; &lt;/span&gt;chunks&lt;span class="w"&gt; &lt;/span&gt;:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;3&lt;/span&gt;,730,146
&lt;span class="nb"&gt;test&lt;/span&gt;&lt;span class="w"&gt; &lt;/span&gt;chunks&lt;span class="w"&gt;  &lt;/span&gt;:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,000
chunk&lt;span class="w"&gt; &lt;/span&gt;length&lt;span class="w"&gt; &lt;/span&gt;:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;128&lt;/span&gt;&lt;span class="w"&gt; &lt;/span&gt;tokens
batch&lt;span class="w"&gt;        &lt;/span&gt;:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;4&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;128&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;  &lt;/span&gt;&lt;span class="nv"&gt;keys&lt;/span&gt;&lt;span class="o"&gt;=[&lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;tokens&amp;#39;&lt;/span&gt;&lt;span class="o"&gt;]&lt;/span&gt;

first&lt;span class="w"&gt; &lt;/span&gt;chunk,&lt;span class="w"&gt; &lt;/span&gt;decoded&lt;span class="w"&gt; &lt;/span&gt;back&lt;span class="w"&gt; &lt;/span&gt;to&lt;span class="w"&gt; &lt;/span&gt;text:

&amp;lt;&lt;span class="p"&gt;|&lt;/span&gt;endoftext&lt;span class="p"&gt;|&lt;/span&gt;&amp;gt;&lt;span class="w"&gt; &lt;/span&gt;park.&lt;span class="w"&gt; &lt;/span&gt;One&lt;span class="w"&gt; &lt;/span&gt;day,&lt;span class="w"&gt; &lt;/span&gt;she&lt;span class="w"&gt; &lt;/span&gt;saw&lt;span class="w"&gt; &lt;/span&gt;a&lt;span class="w"&gt; &lt;/span&gt;big&lt;span class="w"&gt; &lt;/span&gt;white&lt;span class="w"&gt; &lt;/span&gt;airplane&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="k"&gt;in&lt;/span&gt;&lt;span class="w"&gt; &lt;/span&gt;the&lt;span class="w"&gt; &lt;/span&gt;sky.&lt;span class="w"&gt; &lt;/span&gt;She&lt;span class="w"&gt; &lt;/span&gt;pointed&lt;span class="w"&gt; &lt;/span&gt;at&lt;span class="w"&gt; &lt;/span&gt;it&lt;span class="w"&gt; &lt;/span&gt;and&lt;span class="w"&gt; &lt;/span&gt;said,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="s2"&gt;&amp;quot;Look, Mommy! A pilot is flying that plane!&amp;quot;&lt;/span&gt;

Her&lt;span class="w"&gt; &lt;/span&gt;mommy&lt;span class="w"&gt; &lt;/span&gt;said,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="s2"&gt;&amp;quot;Yes, Lily. Pilots fly airplanes.&amp;quot;&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;The chunk begins mid-sentence, on the tail of a story that is not in it. Chunks are cut from the stream, not from stories&lt;/li&gt;
&lt;li&gt;3.7M chunks at batch size 32 is 116,000 steps for one full epoch, so training caps steps per epoch rather than running one to the end&lt;/li&gt;
&lt;/ul&gt;
&lt;h2 id="trainer"&gt;Trainer&lt;/h2&gt;
&lt;p&gt;The loop from the top of this post, with the pieces filled in.&lt;/p&gt;
&lt;h3 id="trainingargs"&gt;&lt;code&gt;TrainingArgs&lt;/code&gt;&lt;/h3&gt;
&lt;div class="language-python highlight"&gt;&lt;span class="code-lang"&gt;python&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="n"&gt;batch_size&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt; &lt;span class="nb"&gt;int&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="mi"&gt;32&lt;/span&gt;
&lt;span class="n"&gt;epochs&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt; &lt;span class="nb"&gt;int&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="mi"&gt;10&lt;/span&gt;
&lt;span class="n"&gt;max_steps_per_epoch&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt; &lt;span class="nb"&gt;int&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="mi"&gt;500&lt;/span&gt;
&lt;span class="n"&gt;lr&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt; &lt;span class="nb"&gt;float&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="mf"&gt;1e-3&lt;/span&gt;
&lt;span class="n"&gt;weight_decay&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt; &lt;span class="nb"&gt;float&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="mf"&gt;1e-2&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;Separate from &lt;code&gt;Config&lt;/code&gt;, which holds the architecture. These change between runs, those change between models&lt;/li&gt;
&lt;li&gt;An epoch is a full pass over the training set, but 3.7M chunks is 116,000 steps, so &lt;code&gt;max_steps_per_epoch&lt;/code&gt; cuts it short&lt;/li&gt;
&lt;/ul&gt;
&lt;h3 id="training_step"&gt;&lt;code&gt;training_step&lt;/code&gt;&lt;/h3&gt;
&lt;div class="language-python highlight"&gt;&lt;span class="code-lang"&gt;python&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="n"&gt;tokens&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;batch&lt;/span&gt;&lt;span class="p"&gt;[&lt;/span&gt;&lt;span class="s2"&gt;&amp;quot;tokens&amp;quot;&lt;/span&gt;&lt;span class="p"&gt;]&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;to&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;device&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;span class="n"&gt;loss&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="o"&gt;-&lt;/span&gt;&lt;span class="n"&gt;get_log_probs&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;model&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;tokens&lt;/span&gt;&lt;span class="p"&gt;),&lt;/span&gt; &lt;span class="n"&gt;tokens&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;mean&lt;/span&gt;&lt;span class="p"&gt;()&lt;/span&gt;
&lt;span class="n"&gt;loss&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;backward&lt;/span&gt;&lt;span class="p"&gt;()&lt;/span&gt;
&lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;optimizer&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;step&lt;/span&gt;&lt;span class="p"&gt;()&lt;/span&gt;
&lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;optimizer&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;zero_grad&lt;/span&gt;&lt;span class="p"&gt;()&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;&lt;code&gt;-mean()&lt;/code&gt; turns the &lt;code&gt;(batch, posn-1)&lt;/code&gt; log-probabilities into one number to minimise&lt;/li&gt;
&lt;li&gt;&lt;code&gt;backward()&lt;/code&gt; fills in a gradient on every parameter, &lt;code&gt;step()&lt;/code&gt; applies it, &lt;code&gt;zero_grad()&lt;/code&gt; clears it for the next batch&lt;/li&gt;
&lt;/ul&gt;
&lt;h3 id="evaluate"&gt;&lt;code&gt;evaluate&lt;/code&gt;&lt;/h3&gt;
&lt;div class="language-python highlight"&gt;&lt;span class="code-lang"&gt;python&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="n"&gt;predictions&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;model&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;tokens&lt;/span&gt;&lt;span class="p"&gt;)[:,&lt;/span&gt; &lt;span class="p"&gt;:&lt;/span&gt;&lt;span class="o"&gt;-&lt;/span&gt;&lt;span class="mi"&gt;1&lt;/span&gt;&lt;span class="p"&gt;]&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;argmax&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;dim&lt;/span&gt;&lt;span class="o"&gt;=-&lt;/span&gt;&lt;span class="mi"&gt;1&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;span class="n"&gt;correct&lt;/span&gt; &lt;span class="o"&gt;+=&lt;/span&gt; &lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;predictions&lt;/span&gt; &lt;span class="o"&gt;==&lt;/span&gt; &lt;span class="n"&gt;tokens&lt;/span&gt;&lt;span class="p"&gt;[:,&lt;/span&gt; &lt;span class="mi"&gt;1&lt;/span&gt;&lt;span class="p"&gt;:])&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;sum&lt;/span&gt;&lt;span class="p"&gt;()&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;item&lt;/span&gt;&lt;span class="p"&gt;()&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;The same &lt;code&gt;[:, :-1]&lt;/code&gt; and &lt;code&gt;[1:]&lt;/code&gt; offset as the loss, on the held-out chunks the model never trains on&lt;/li&gt;
&lt;li&gt;&lt;code&gt;argmax&lt;/code&gt; instead of &lt;code&gt;gather&lt;/code&gt;: how often the model's top pick was right, rather than what it scored the right answer&lt;/li&gt;
&lt;/ul&gt;
&lt;h3 id="module-output_2"&gt;Module output&lt;/h3&gt;
&lt;div class="language-text highlight"&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;{&amp;quot;step&amp;quot;: 1, &amp;quot;loss&amp;quot;: 10.83082389831543, &amp;quot;epoch&amp;quot;: 0}
{&amp;quot;step&amp;quot;: 2, &amp;quot;loss&amp;quot;: 10.77728271484375, &amp;quot;epoch&amp;quot;: 0}
...
{&amp;quot;step&amp;quot;: 60, &amp;quot;loss&amp;quot;: 7.720832824707031, &amp;quot;epoch&amp;quot;: 1}
{&amp;quot;step&amp;quot;: 60, &amp;quot;accuracy&amp;quot;: 0.0453986220472441, &amp;quot;epoch&amp;quot;: 1}
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;h2 id="sampler"&gt;Sampler&lt;/h2&gt;
&lt;div class="mermaid"&gt;flowchart TB
    dir["run_dir&amp;lt;br/&amp;gt;config.json + model.pt"]

    dir --&amp;gt;|"Config(**saved)"| cfg["Config&amp;lt;br/&amp;gt;d_model, n_heads, n_ctx, ..."]
    cfg --&amp;gt; model["Transformer&amp;lt;br/&amp;gt;our modules"]
    dir --&amp;gt;|"t.load(model.pt)"| model

    model --&amp;gt;|"state_dict, strict=False"| hooked["HookedTransformer&amp;lt;br/&amp;gt;our weights, their loop"]

    repl["REPL&amp;lt;br/&amp;gt;a line starting with : sets a knob"] --&amp;gt; knobs["Sampler&amp;lt;br/&amp;gt;temperature, top_k, top_p, len"]
    repl --&amp;gt;|"anything else is a prompt"| gen
    knobs --&amp;gt; gen["hooked.generate()"]
    hooked --&amp;gt; gen
    gen --&amp;gt; text["text"]&lt;/div&gt;
&lt;blockquote&gt;
&lt;p&gt;&lt;strong&gt;Update, 3 October 2026&lt;/strong&gt;&lt;br&gt;
The generation loop has since been ported. See &lt;a href="https://weirdmachine.wtf/posts/sampling-from-a-transformer/"&gt;Sampling from a Transformer&lt;/a&gt;. The rest of this section describes how it worked at the time of writing.&lt;/p&gt;
&lt;/blockquote&gt;
&lt;p&gt;I have not ported the generation loop. &lt;code&gt;experiments/sample.py&lt;/code&gt; loads the trained weights into a TransformerLens &lt;code&gt;HookedTransformer&lt;/code&gt; and calls its &lt;code&gt;generate&lt;/code&gt;.&lt;/p&gt;
&lt;ul&gt;
&lt;li&gt;A run directory is self-contained: &lt;code&gt;config.json&lt;/code&gt; rebuilds the architecture and &lt;code&gt;model.pt&lt;/code&gt; fills it. Nothing else is needed to bring a finished run back&lt;/li&gt;
&lt;li&gt;&lt;code&gt;strict=False&lt;/code&gt; because TransformerLens carries buffers we do not, the causal mask among them. The parameter names match, which is the part that matters&lt;/li&gt;
&lt;li&gt;The &lt;code&gt;Sampler&lt;/code&gt; class holds nothing but the knobs and how they read back. It exists so the REPL can print the current settings and convert them to the arguments &lt;code&gt;generate&lt;/code&gt; expects&lt;/li&gt;
&lt;/ul&gt;
&lt;h3 id="options"&gt;Options&lt;/h3&gt;
&lt;div class="table-wrap"&gt;&lt;table&gt;
&lt;thead&gt;
&lt;tr&gt;
&lt;th&gt;Command&lt;/th&gt;
&lt;th&gt;Effect&lt;/th&gt;
&lt;/tr&gt;
&lt;/thead&gt;
&lt;tbody&gt;
&lt;tr&gt;
&lt;td&gt;&lt;code&gt;:temp 0.7&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;Divides the logits before the softmax. Below 1 sharpens the distribution, above 1 flattens it&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;code&gt;:greedy&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;Always take the most likely token. The same as temperature 0&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;code&gt;:topk 40&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;Keep the 40 best tokens, discard the rest, renormalise&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;code&gt;:topp 0.95&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;Keep the smallest set of tokens whose probabilities sum to 0.95&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;code&gt;:len 60&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;How many tokens to generate&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;/td&gt;
&lt;td&gt;&lt;/td&gt;
&lt;/tr&gt;
&lt;/tbody&gt;
&lt;/table&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;&lt;code&gt;top_k&lt;/code&gt; and &lt;code&gt;top_p&lt;/code&gt; are alternatives, not a pair. Setting one clears the other, because TransformerLens checks &lt;code&gt;top_k&lt;/code&gt; first and it would silently win&lt;/li&gt;
&lt;li&gt;The defaults are temperature 0.7 with &lt;code&gt;top_p&lt;/code&gt; 0.95, which is what every sample in this post used&lt;/li&gt;
&lt;/ul&gt;
&lt;h3 id="module-output_3"&gt;Module output&lt;/h3&gt;
&lt;div class="language-bash highlight"&gt;&lt;span class="code-lang"&gt;shell&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="o"&gt;(&lt;/span&gt;.venv&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt; &lt;/span&gt;dom@dom-z13:~/Desktop/PvML$&lt;span class="w"&gt; &lt;/span&gt;python&lt;span class="w"&gt; &lt;/span&gt;experiments/sample.py&lt;span class="w"&gt; &lt;/span&gt;runs/tinystories-d128-l6-h4-ctx512-40k

runs/tinystories-d128-l6-h4-ctx512-40k
&lt;span class="w"&gt;  &lt;/span&gt;&lt;span class="m"&gt;14&lt;/span&gt;,171,473&lt;span class="w"&gt; &lt;/span&gt;parameters,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;6&lt;/span&gt;&lt;span class="w"&gt; &lt;/span&gt;layers,&lt;span class="w"&gt; &lt;/span&gt;d_model&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;128&lt;/span&gt;
&lt;span class="w"&gt;  &lt;/span&gt;temp&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.7,&lt;span class="w"&gt; &lt;/span&gt;top_p&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.95,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;50&lt;/span&gt;&lt;span class="w"&gt; &lt;/span&gt;tokens
&lt;span class="w"&gt;  &lt;/span&gt;a&lt;span class="w"&gt; &lt;/span&gt;prompt&lt;span class="w"&gt; &lt;/span&gt;generates,&lt;span class="w"&gt; &lt;/span&gt;:help&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="k"&gt;for&lt;/span&gt;&lt;span class="w"&gt; &lt;/span&gt;commands,&lt;span class="w"&gt; &lt;/span&gt;:q&lt;span class="w"&gt; &lt;/span&gt;to&lt;span class="w"&gt; &lt;/span&gt;quit

pvml&amp;gt;&lt;span class="w"&gt; &lt;/span&gt;Once&lt;span class="w"&gt; &lt;/span&gt;upon&lt;span class="w"&gt; &lt;/span&gt;a&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="nb"&gt;time&lt;/span&gt;
Once&lt;span class="w"&gt; &lt;/span&gt;upon&lt;span class="w"&gt; &lt;/span&gt;a&lt;span class="w"&gt; &lt;/span&gt;time,&lt;span class="w"&gt; &lt;/span&gt;there&lt;span class="w"&gt; &lt;/span&gt;was&lt;span class="w"&gt; &lt;/span&gt;a&lt;span class="w"&gt; &lt;/span&gt;little&lt;span class="w"&gt; &lt;/span&gt;girl&lt;span class="w"&gt; &lt;/span&gt;named&lt;span class="w"&gt; &lt;/span&gt;Lily.&lt;span class="w"&gt; &lt;/span&gt;She&lt;span class="w"&gt; &lt;/span&gt;loved&lt;span class="w"&gt; &lt;/span&gt;to&lt;span class="w"&gt; &lt;/span&gt;play&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="k"&gt;in&lt;/span&gt;&lt;span class="w"&gt; &lt;/span&gt;her
backyard,&lt;span class="w"&gt; &lt;/span&gt;especially&lt;span class="w"&gt; &lt;/span&gt;when&lt;span class="w"&gt; &lt;/span&gt;she&lt;span class="w"&gt; &lt;/span&gt;saw&lt;span class="w"&gt; &lt;/span&gt;a&lt;span class="w"&gt; &lt;/span&gt;big,&lt;span class="w"&gt; &lt;/span&gt;scary&lt;span class="w"&gt; &lt;/span&gt;monster!&lt;span class="w"&gt; &lt;/span&gt;She&lt;span class="w"&gt; &lt;/span&gt;was&lt;span class="w"&gt; &lt;/span&gt;so&lt;span class="w"&gt; &lt;/span&gt;scared&lt;span class="w"&gt; &lt;/span&gt;that
she&lt;span class="w"&gt; &lt;/span&gt;started&lt;span class="w"&gt; &lt;/span&gt;to&lt;span class="w"&gt; &lt;/span&gt;cry.

Her&lt;span class="w"&gt; &lt;/span&gt;mom&lt;span class="w"&gt; &lt;/span&gt;heard&lt;span class="w"&gt; &lt;/span&gt;her&lt;span class="w"&gt; &lt;/span&gt;cry&lt;span class="w"&gt; &lt;/span&gt;and&lt;span class="w"&gt; &lt;/span&gt;said,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="s2"&gt;&amp;quot;Don&lt;/span&gt;

&lt;span class="s2"&gt;pvml&amp;gt; :greedy&lt;/span&gt;
&lt;span class="s2"&gt;  greedy, 50 tokens&lt;/span&gt;
&lt;span class="s2"&gt;pvml&amp;gt; Once upon a time&lt;/span&gt;
&lt;span class="s2"&gt;Once upon a time, there was a little girl named Lily. She loved to play outside&lt;/span&gt;
&lt;span class="s2"&gt;in the sunshine. One day, she went to the park to play with her friends. She saw&lt;/span&gt;
&lt;span class="s2"&gt;a big tree and decided to climb it.&lt;/span&gt;

&lt;span class="s2"&gt;As she climbed the&lt;/span&gt;

&lt;span class="s2"&gt;pvml&amp;gt; :temp 1.4&lt;/span&gt;
&lt;span class="s2"&gt;  temp 1.4, top_p 0.95, 50 tokens&lt;/span&gt;
&lt;span class="s2"&gt;pvml&amp;gt; Once upon a time&lt;/span&gt;
&lt;span class="s2"&gt;Once upon a time a painter who was the most courageous, dirty every painting&lt;/span&gt;
&lt;span class="s2"&gt;player had been wanting to take off the beautiful thin balloons home. The weary&lt;/span&gt;
&lt;span class="s2"&gt;square gained swiftly and carried, paying with daily thin brushes for his&lt;/span&gt;
&lt;span class="s2"&gt;reliable folder again to admire them whenever most happily grinding balloons&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;h2 id="experiments"&gt;Experiments&lt;/h2&gt;
&lt;h3 id="environment"&gt;Environment&lt;/h3&gt;
&lt;div class="table-wrap"&gt;&lt;table&gt;
&lt;thead&gt;
&lt;tr&gt;
&lt;th&gt;Component&lt;/th&gt;
&lt;th&gt;Detail&lt;/th&gt;
&lt;/tr&gt;
&lt;/thead&gt;
&lt;tbody&gt;
&lt;tr&gt;
&lt;td&gt;Device&lt;/td&gt;
&lt;td&gt;ASUS ROG Flow Z13, Strix Halo (Ryzen AI MAX+ 395)&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;OS&lt;/td&gt;
&lt;td&gt;Ubuntu 26.04.1 LTS, kernel 7.0.0-31-generic&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;GPU&lt;/td&gt;
&lt;td&gt;Radeon 8060S integrated graphics, 128GB unified memory&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;PyTorch&lt;/td&gt;
&lt;td&gt;2.12.0a0+rocm7.13, ROCm nightly for gfx1151&lt;/td&gt;
&lt;/tr&gt;
&lt;/tbody&gt;
&lt;/table&gt;&lt;/div&gt;
&lt;h3 id="tinystories"&gt;&lt;code&gt;TinyStories&lt;/code&gt;&lt;/h3&gt;
&lt;ul&gt;
&lt;li&gt;&lt;a href="https://huggingface.co/datasets/roneneldan/TinyStories"&gt;https://huggingface.co/datasets/roneneldan/TinyStories&lt;/a&gt;&lt;/li&gt;
&lt;/ul&gt;
&lt;div class="language-text highlight"&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;Dataset({ features: [&amp;#39;text&amp;#39;], num_rows: 2119719 })
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;2.1M stories written with the vocabulary of a three or four year old, generated by GPT-3.5 and GPT-4 for the paper &lt;em&gt;TinyStories: How Small Can Language Models Be and Still Speak Coherent English?&lt;/em&gt;&lt;/li&gt;
&lt;li&gt;Simple enough for a model of a few million parameters to learn on one machine&lt;/li&gt;
&lt;/ul&gt;
&lt;p&gt;Story length in GPT-2 tokens, on a 20,000 story sample:&lt;/p&gt;
&lt;div class="table-wrap"&gt;&lt;table&gt;
&lt;thead&gt;
&lt;tr&gt;
&lt;th&gt;Mean&lt;/th&gt;
&lt;th&gt;Median&lt;/th&gt;
&lt;th&gt;p90&lt;/th&gt;
&lt;th&gt;p99&lt;/th&gt;
&lt;th&gt;Longest&lt;/th&gt;
&lt;/tr&gt;
&lt;/thead&gt;
&lt;tbody&gt;
&lt;tr&gt;
&lt;td&gt;222&lt;/td&gt;
&lt;td&gt;191&lt;/td&gt;
&lt;td&gt;358&lt;/td&gt;
&lt;td&gt;641&lt;/td&gt;
&lt;td&gt;1106&lt;/td&gt;
&lt;/tr&gt;
&lt;/tbody&gt;
&lt;/table&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;The corpus is 478M tokens, 3,731,146 chunks at &lt;code&gt;n_ctx=128&lt;/code&gt;&lt;/li&gt;
&lt;li&gt;20,000 steps at batch 32 and &lt;code&gt;n_ctx&lt;/code&gt; 128 is 81.9M tokens, 0.17 epochs. No run here is limited by running out of data&lt;/li&gt;
&lt;/ul&gt;
&lt;h3 id="plan"&gt;Plan&lt;/h3&gt;
&lt;div class="table-wrap"&gt;&lt;table&gt;
&lt;thead&gt;
&lt;tr&gt;
&lt;th&gt;Run&lt;/th&gt;
&lt;th&gt;Question&lt;/th&gt;
&lt;/tr&gt;
&lt;/thead&gt;
&lt;tbody&gt;
&lt;tr&gt;
&lt;td&gt;&lt;code&gt;tinystories-d32-l4-h16-ctx128-5k&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;how far does the smallest model get&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;code&gt;tinystories-d128-l6-h4-ctx128-5k&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;does more capacity help at the same steps&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;code&gt;tinystories-d256-l8-h8-ctx128-5k&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;does capacity keep paying, or level off&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;code&gt;tinystories-d32-l4-h16-ctx128-20k&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;is the small model limited by size or by training&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;code&gt;tinystories-d128-l6-h4-ctx128-20k&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;does the medium model have a ceiling of its own&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;code&gt;tinystories-d128-l6-h4-ctx512-40k&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;added later: was the wandering a context limit&lt;/td&gt;
&lt;/tr&gt;
&lt;/tbody&gt;
&lt;/table&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;Steps are held at 5,000 for the first three and 20,000 for the next two, so capacity and training length move one at a time&lt;/li&gt;
&lt;li&gt;&lt;code&gt;n_ctx&lt;/code&gt; is 128 across those five. Context changes what a model can see, not just how big it is, so it gets its own run&lt;/li&gt;
&lt;li&gt;The sixth was not planned. It came out of reading the samples from the first five&lt;/li&gt;
&lt;/ul&gt;
&lt;h3 id="results"&gt;Results&lt;/h3&gt;
&lt;p&gt;The first five ran 22:50 to 01:35, the sixth 02:59 to 13:40 the next day.&lt;/p&gt;
&lt;div class="table-wrap"&gt;&lt;table&gt;
&lt;thead&gt;
&lt;tr&gt;
&lt;th&gt;Run&lt;/th&gt;
&lt;th&gt;n_ctx&lt;/th&gt;
&lt;th&gt;Steps&lt;/th&gt;
&lt;th&gt;Loss&lt;/th&gt;
&lt;th&gt;Accuracy&lt;/th&gt;
&lt;th&gt;Minutes&lt;/th&gt;
&lt;th&gt;Last 500 steps&lt;/th&gt;
&lt;/tr&gt;
&lt;/thead&gt;
&lt;tbody&gt;
&lt;tr&gt;
&lt;td&gt;&lt;code&gt;d128-l6-h4-ctx512-40k&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;512&lt;/td&gt;
&lt;td&gt;40,000&lt;/td&gt;
&lt;td&gt;1.732&lt;/td&gt;
&lt;td&gt;0.573&lt;/td&gt;
&lt;td&gt;641.3&lt;/td&gt;
&lt;td&gt;-0.002&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;code&gt;d128-l6-h4-ctx128-20k&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;128&lt;/td&gt;
&lt;td&gt;20,000&lt;/td&gt;
&lt;td&gt;2.080&lt;/td&gt;
&lt;td&gt;0.517&lt;/td&gt;
&lt;td&gt;67.3&lt;/td&gt;
&lt;td&gt;-0.005&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;code&gt;d128-l6-h4-ctx128-5k&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;128&lt;/td&gt;
&lt;td&gt;5,000&lt;/td&gt;
&lt;td&gt;2.369&lt;/td&gt;
&lt;td&gt;0.476&lt;/td&gt;
&lt;td&gt;17.0&lt;/td&gt;
&lt;td&gt;-0.036&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;code&gt;d256-l8-h8-ctx128-5k&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;128&lt;/td&gt;
&lt;td&gt;5,000&lt;/td&gt;
&lt;td&gt;2.497&lt;/td&gt;
&lt;td&gt;0.461&lt;/td&gt;
&lt;td&gt;36.6&lt;/td&gt;
&lt;td&gt;-0.064&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;code&gt;d32-l4-h16-ctx128-20k&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;128&lt;/td&gt;
&lt;td&gt;20,000&lt;/td&gt;
&lt;td&gt;2.782&lt;/td&gt;
&lt;td&gt;0.421&lt;/td&gt;
&lt;td&gt;34.6&lt;/td&gt;
&lt;td&gt;+0.006&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;&lt;code&gt;d32-l4-h16-ctx128-5k&lt;/code&gt;&lt;/td&gt;
&lt;td&gt;128&lt;/td&gt;
&lt;td&gt;5,000&lt;/td&gt;
&lt;td&gt;2.996&lt;/td&gt;
&lt;td&gt;0.395&lt;/td&gt;
&lt;td&gt;8.4&lt;/td&gt;
&lt;td&gt;-0.024&lt;/td&gt;
&lt;/tr&gt;
&lt;/tbody&gt;
&lt;/table&gt;&lt;/div&gt;
&lt;p&gt;Only the five at &lt;code&gt;n_ctx&lt;/code&gt; 128 are comparable with each other. Predicting a token from 511 tokens of context is an easier problem than from 127, so part of the top row's lead is the task, not the model. Re-chunking also changes the split, so the two windows were not evaluated on the same held-out text.&lt;/p&gt;
&lt;ul&gt;
&lt;li&gt;Capacity beats steps. &lt;code&gt;d128&lt;/code&gt; reaches 2.369 in 5,000 steps and 17 minutes. &lt;code&gt;d32&lt;/code&gt; needs 20,000 steps and 35 minutes to reach only 2.782, so four times the steps and twice the clock still lands short&lt;/li&gt;
&lt;li&gt;&lt;code&gt;d32&lt;/code&gt; uses 16 heads against &lt;code&gt;d128&lt;/code&gt;'s 4, so its &lt;code&gt;d_head&lt;/code&gt; is 2 where theirs is 32. The capacity comparison therefore mixes width with unusually narrow heads, and a rerun at &lt;code&gt;d32&lt;/code&gt; with 4 heads would separate them&lt;/li&gt;
&lt;li&gt;Of the five, only &lt;code&gt;d32&lt;/code&gt; at 20,000 steps has converged. Its last 500 steps went up 0.006, noise around a floor near 2.78. &lt;code&gt;d128&lt;/code&gt; at 20,000 is still moving and &lt;code&gt;d256&lt;/code&gt; was descending fastest of anything when it stopped&lt;/li&gt;
&lt;li&gt;&lt;code&gt;d256&lt;/code&gt; losing to &lt;code&gt;d128&lt;/code&gt; at 5,000 steps looks like undertraining rather than a ceiling, though nothing here rules the ceiling out. First epoch 5.404 against &lt;code&gt;d128&lt;/code&gt;'s 4.343, then the steepest end slope of the five&lt;/li&gt;
&lt;li&gt;Every run uses a constant &lt;code&gt;lr=1e-3&lt;/code&gt; with no warmup and no decay, inherited from the tiny model. Warmup plus cosine decay is the first thing to try on &lt;code&gt;d256&lt;/code&gt; before concluding anything about its capacity&lt;/li&gt;
&lt;li&gt;&lt;code&gt;d256&lt;/code&gt; at 20,000 steps is still unrun. The end slopes suggest it would overtake, which is a hypothesis and not a result&lt;/li&gt;
&lt;/ul&gt;
&lt;h3 id="what-the-best-of-these-writes"&gt;What the best of these writes&lt;/h3&gt;
&lt;p&gt;The best and the worst of the five on the same two prompts. Temperature 0.7, &lt;code&gt;top_p&lt;/code&gt; 0.95, seed 0, 60 new tokens. Each block holds only what the model wrote after the prompt above it.&lt;/p&gt;
&lt;p&gt;&lt;code&gt;d128-l6-h4-ctx128-20k&lt;/code&gt;, final loss 2.080.&lt;/p&gt;
&lt;p&gt;Prompt &lt;code&gt;Once upon a time&lt;/code&gt;:&lt;/p&gt;
&lt;div class="language-text highlight"&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;there was a little girl called Sally. She loved her garden and always liked to
play in the garden. One day, Sally was playing in the garden when she heard a
loud noise.

Suddenly, the noise disturbed her. It was so loud that the girl couldn&amp;#39;t see
her house. She
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;p&gt;Prompt &lt;code&gt;Lily went to the park and saw&lt;/code&gt;:&lt;/p&gt;
&lt;div class="language-text highlight"&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;a big tree. She wanted to climb it, but her mom said no. She said, &amp;quot;You can&amp;#39;t
climb the tree, Lily. It is not safe. It is not safe. You can fall from the
tree.&amp;quot;

Lily was sad, but she did not want to climb
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;p&gt;&lt;code&gt;d32-l4-h16-ctx128-20k&lt;/code&gt;, final loss 2.782, same prompts and settings.&lt;/p&gt;
&lt;p&gt;Prompt &lt;code&gt;Once upon a time&lt;/code&gt;:&lt;/p&gt;
&lt;div class="language-text highlight"&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;there was a little girl named Sally. One day, Sally went to the park to play.
She saw a big tree. She wanted to play outside. So, she asked her mommy for
help. Her mommy said yes.

Suddenly, the squirrel started to run and higher.
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;p&gt;Prompt &lt;code&gt;Lily went to the park and saw&lt;/code&gt;:&lt;/p&gt;
&lt;div class="language-text highlight"&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;a big, orange bear. She was so happy and looked around the tree. She said, &amp;quot;I
like the bear. I found a pear. I like it.&amp;quot; She said, &amp;quot;Yes, I want to play with
it. Can I help you?&amp;quot;

She ran to the ball and
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;Both are fluent and both hold one name throughout, so neither is the difference&lt;/li&gt;
&lt;li&gt;&lt;code&gt;d32&lt;/code&gt; introduces things nothing set up: a squirrel, a pear, a ball. &lt;code&gt;d128&lt;/code&gt; keeps referring to what it already named&lt;/li&gt;
&lt;li&gt;&lt;code&gt;d128&lt;/code&gt; holds a short causal chain: mom refuses, Lily is sad. &lt;code&gt;d32&lt;/code&gt; has Lily answer &lt;code&gt;"Yes,"&lt;/code&gt; to a question nobody asked&lt;/li&gt;
&lt;li&gt;&lt;code&gt;d128&lt;/code&gt; still repeats. &lt;code&gt;It is not safe&lt;/code&gt; twice within four words&lt;/li&gt;
&lt;li&gt;The samples stop mid-sentence because the cap is 60 tokens, not because the model ran out&lt;/li&gt;
&lt;/ul&gt;
&lt;h3 id="plots"&gt;Plots&lt;/h3&gt;
&lt;div class="language-bash highlight"&gt;&lt;span class="code-lang"&gt;shell&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;python&lt;span class="w"&gt; &lt;/span&gt;experiments/plot.py
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;p&gt;Reads every run with a &lt;code&gt;summary.json&lt;/code&gt;. Runs are grouped by window, so the four charts below hold only the five at &lt;code&gt;n_ctx&lt;/code&gt; 128.&lt;/p&gt;
&lt;p&gt;Who plateaued and who was still descending when it stopped.&lt;/p&gt;
&lt;p&gt;&lt;img alt="Loss against step for the five n_ctx 128 runs" src="https://weirdmachine.wtf/images/training-a-transformer/loss-by-step.png"&gt;&lt;/p&gt;
&lt;p&gt;The same runs priced in minutes rather than steps.&lt;/p&gt;
&lt;p&gt;&lt;img alt="Loss against wall-clock minutes for the five n_ctx 128 runs" src="https://weirdmachine.wtf/images/training-a-transformer/loss-by-time.png"&gt;&lt;/p&gt;
&lt;p&gt;Final loss against non-embedding parameters.&lt;/p&gt;
&lt;p&gt;&lt;img alt="Final loss against non-embedding parameters" src="https://weirdmachine.wtf/images/training-a-transformer/loss-by-capacity.png"&gt;&lt;/p&gt;
&lt;p&gt;Next-token accuracy per epoch.&lt;/p&gt;
&lt;p&gt;&lt;img alt="Next-token accuracy per epoch for the five n_ctx 128 runs" src="https://weirdmachine.wtf/images/training-a-transformer/accuracy.png"&gt;&lt;/p&gt;
&lt;ul&gt;
&lt;li&gt;Per-step loss is too noisy to read, so it is averaged into 100-step windows&lt;/li&gt;
&lt;li&gt;Only non-embedding parameters are plotted. The vocabulary tables dominate these models and do not change with depth or width, so including them would stack every run in the same place&lt;/li&gt;
&lt;li&gt;Loss against wall clock is the chart that changes a decision. Per step a bigger model always looks better, per minute it sometimes does not&lt;/li&gt;
&lt;/ul&gt;
&lt;h2 id="a-longer-window"&gt;A longer window&lt;/h2&gt;
&lt;p&gt;The five runs above all used &lt;code&gt;n_ctx=128&lt;/code&gt;, and every one of them wrote stories that wandered. The obvious suspect was capacity. It was not.&lt;/p&gt;
&lt;h3 id="why-128-was-the-wrong-number"&gt;Why 128 was the wrong number&lt;/h3&gt;
&lt;p&gt;How much of a story each window holds, against a median of 191 tokens:&lt;/p&gt;
&lt;div class="table-wrap"&gt;&lt;table&gt;
&lt;thead&gt;
&lt;tr&gt;
&lt;th&gt;n_ctx&lt;/th&gt;
&lt;th&gt;Stories that fit whole&lt;/th&gt;
&lt;/tr&gt;
&lt;/thead&gt;
&lt;tbody&gt;
&lt;tr&gt;
&lt;td&gt;128&lt;/td&gt;
&lt;td&gt;5.2%&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;512&lt;/td&gt;
&lt;td&gt;96.8%&lt;/td&gt;
&lt;/tr&gt;
&lt;tr&gt;
&lt;td&gt;1024&lt;/td&gt;
&lt;td&gt;99.9%&lt;/td&gt;
&lt;/tr&gt;
&lt;/tbody&gt;
&lt;/table&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;&lt;code&gt;tokenize_and_concatenate&lt;/code&gt; glues the corpus into one stream and chops it into &lt;code&gt;n_ctx&lt;/code&gt; chunks, so a chunk starts and ends wherever it lands. At 128 tokens that is almost always mid-story&lt;/li&gt;
&lt;li&gt;Endings were never the missing piece. Chunks land across story boundaries at a rate of roughly 0.58 per chunk, so the earlier models read endings constantly. What they never saw was one whole arc, start to finish, inside a single window&lt;/li&gt;
&lt;li&gt;512 is chosen on coverage, not on cost. It holds 96.8% of stories whole and 1024 adds three points. Measured at equal tokens per step, doubling to 1024 costs 1.28x, which would have been affordable&lt;/li&gt;
&lt;/ul&gt;
&lt;h3 id="the-run"&gt;The run&lt;/h3&gt;
&lt;div class="language-text highlight"&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;tinystories-d128-l6-h4-ctx512-40k
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;Same model as the best ctx128 run: &lt;code&gt;d_model&lt;/code&gt; 128, 6 layers, 4 heads. Same &lt;code&gt;batch_size&lt;/code&gt; 32, &lt;code&gt;lr&lt;/code&gt; 1e-3, &lt;code&gt;seed&lt;/code&gt; 0&lt;/li&gt;
&lt;li&gt;40,000 steps at this window is 655M tokens, 1.37 epochs of the corpus and 8x the token budget of the best ctx128 run&lt;/li&gt;
&lt;li&gt;&lt;code&gt;max_steps_per_epoch&lt;/code&gt; moved from 500 to 1000, which halves the number of evals for the same total. It changes how often training pauses, not what the model sees&lt;/li&gt;
&lt;/ul&gt;
&lt;p&gt;It reached 1.732 and 0.573, against 2.080 and 0.517 for the same model at &lt;code&gt;n_ctx&lt;/code&gt; 128. For the reason given above, that gap is not all model, so the samples are the evidence that counts.&lt;/p&gt;
&lt;p&gt;&lt;img alt="Loss and accuracy of the n_ctx 512 run against the n_ctx 128 run of the same model" src="https://weirdmachine.wtf/images/training-a-transformer/ctx512.png"&gt;&lt;/p&gt;
&lt;h3 id="what-changed"&gt;What changed&lt;/h3&gt;
&lt;p&gt;Temperature 0.7, &lt;code&gt;top_p&lt;/code&gt; 0.95, seed 0, and a 250 token cap neither of these two reaches.&lt;/p&gt;
&lt;p&gt;Prompt &lt;code&gt;Once upon a time&lt;/code&gt;:&lt;/p&gt;
&lt;div class="language-text highlight"&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;there was a little girl called Sally. She loved to explore and try new
things.

One day Sally went to a big, modern building. She was so excited! Inside the
building, she found a box. Inside, there were lots of toys, but Sally was so
excited. She could not wait to play with them.

Sally was so happy. She could not believe it! She had so many new toys to
play with. Sally was so excited that she couldn&amp;#39;t wait to play with them
again!

The next day, Sally took her new toys to the museum and had lots of fun. She
found lots of new things to play with. She was so happy that she had a great
time!
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;p&gt;Prompt &lt;code&gt;Lily went to the park and saw&lt;/code&gt;:&lt;/p&gt;
&lt;div class="language-text highlight"&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;a big tree. She wanted to climb it and see what was on the other side. She
climbed and climbed and reached the top. She was very happy and excited. She
saw the whole park and saw the flowers and the birds. She saw a butterfly and
the birds and the squirrels. She felt very happy and proud. She felt very
special and loved. She thanked the squirrel and said thank you. The squirrel
smiled and said thank you. Lily was glad she could go to the park and see more
things. She felt very happy and proud.
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;Both stop on their own, the first after 150 tokens. Raising the cap from 250 to 400 changes neither. That is these two samples, not the general case: the five-seed count below is seven of ten&lt;/li&gt;
&lt;li&gt;It has a shape: a setup, a next day, a close. At &lt;code&gt;n_ctx&lt;/code&gt; 128 the model only ever saw slices of a story, never a whole one&lt;/li&gt;
&lt;li&gt;Repetition got worse. &lt;code&gt;so excited&lt;/code&gt; three times, &lt;code&gt;wait to play with them&lt;/code&gt; twice&lt;/li&gt;
&lt;li&gt;Reference drifts. The &lt;code&gt;modern building&lt;/code&gt; becomes &lt;code&gt;the museum&lt;/code&gt;, and Lily thanks a squirrel that was only an item in a list of things she saw&lt;/li&gt;
&lt;/ul&gt;
&lt;p&gt;Across five seeds per prompt at the same settings, seven of the ten stories reached end-of-text on their own and three ran to the 250 token cap. The three that hit the cap are the three that fell into a repetition loop, &lt;code&gt;I love you&lt;/code&gt; and &lt;code&gt;She wished she had&lt;/code&gt; repeating until the cap stopped them. So the window bought story structure, and repetition is what still breaks it. Coherence within a story is uneven rather than absent: &lt;code&gt;Timmy&lt;/code&gt; fixing a picture with glue holds together, while a bear carries candy out of a tree that was never mentioned.&lt;/p&gt;
&lt;h3 id="cost"&gt;Cost&lt;/h3&gt;
&lt;ul&gt;
&lt;li&gt;10.7 hours on the Z13, against 67 minutes for the ctx128 run of the same model&lt;/li&gt;
&lt;li&gt;0.962 sec/step against 0.202, both measured end to end on the real runs. That is 4.8x for 4x the tokens per step&lt;/li&gt;
&lt;li&gt;Profiling the forward pass alone at batch 32, so no backward and no optimiser, puts most of the time in the unembedding rather than attention. At &lt;code&gt;n_ctx&lt;/code&gt; 128 it is 28.0ms of 37.0ms, 75.7%. At 512 it is 108.8ms of 199.4ms, 54.6%. Projecting &lt;code&gt;d_model&lt;/code&gt; 128 up to 50,257 logits at every position is a larger matmul than anything in the blocks&lt;/li&gt;
&lt;li&gt;The blocks are what make the cost superlinear. Going 128 to 512 is 4x the tokens, and they take 10.1x the time while the unembedding takes 3.9x. So attention's quadratic term explains the excess over 4x, while the unembedding explains the bulk of the absolute cost&lt;/li&gt;
&lt;li&gt;The &lt;code&gt;n_ctx=512&lt;/code&gt; chunking is not the same dataset as &lt;code&gt;n_ctx=128&lt;/code&gt;, so the corpus is tokenized again. That is a one-time pass and 4GB of cache&lt;/li&gt;
&lt;/ul&gt;</content><category term="transformers"/><category term="tinystories"/><category term="arena"/></entry><entry><title>Understanding Transformers</title><link href="https://weirdmachine.wtf/posts/understanding-transformers/" rel="alternate"/><published>2026-09-27T00:00:00-04:00</published><updated>2026-09-27T00:00:00-04:00</updated><author><name>Dominic Wang</name></author><id>tag:weirdmachine.wtf,2026-09-27:/posts/understanding-transformers/</id><summary type="html">&lt;p&gt;Notes from rebuilding GPT-2 from scratch while working through ARENA's curriculum, module by module, with the shapes and diagrams I needed to follow it.&lt;/p&gt;</summary><content type="html">&lt;p&gt;I've been self-studying ARENA, an AI-safety curriculum whose first chapter has you rebuild GPT-2 from scratch. Transformers were the gap I kept hitting. They're the core of every LLM, and the interpretability side I actually care about is hard to follow without understanding what sits underneath it.&lt;/p&gt;
&lt;p&gt;So I ported the whole chapter into my own modular implementation, loading GPT-2's weights into each module as I went and comparing the output against the reference, until every module matched. These are my notes from doing that, module by module, with the shapes and diagrams I needed to follow it.&lt;/p&gt;
&lt;p&gt;PvML: &lt;a href="https://github.com/d0mzw/PvML"&gt;https://github.com/d0mzw/PvML&lt;/a&gt;&lt;/p&gt;
&lt;p&gt;Disclaimer: these notes come from working through the ARENA 3.0 curriculum. I claim no credit for the original material, and this is not affiliated with or endorsed by ARENA.&lt;/p&gt;
&lt;h2 id="config"&gt;Config&lt;/h2&gt;
&lt;div class="mermaid"&gt;flowchart TB
    cfg["Config&amp;lt;br/&amp;gt;GPT-2 small"]

    cfg --&amp;gt;|"d_vocab, d_model"| embed["Embed&amp;lt;br/&amp;gt;W_E (50257, 768)"]
    cfg --&amp;gt;|"n_ctx, d_model"| pos["PosEmbed&amp;lt;br/&amp;gt;W_pos (1024, 768)"]
    cfg --&amp;gt;|"d_model, layer_norm_eps"| ln["LayerNorm&amp;lt;br/&amp;gt;w, b (768,)"]
    cfg --&amp;gt;|"n_heads, d_head, d_model"| attn["Attention&amp;lt;br/&amp;gt;W_Q (12, 768, 64)"]
    cfg --&amp;gt;|"d_model, d_mlp"| mlp["MLP&amp;lt;br/&amp;gt;W_in (768, 3072)"]
    cfg --&amp;gt;|"d_model, d_vocab"| unembed["Unembed&amp;lt;br/&amp;gt;W_U (768, 50257)"]
    cfg --&amp;gt;|"n_layers"| block["TransformerBlock&amp;lt;br/&amp;gt;x12"]&lt;/div&gt;
&lt;div class="language-python highlight"&gt;&lt;span class="code-lang"&gt;python&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="nd"&gt;@dataclass&lt;/span&gt;
&lt;span class="k"&gt;class&lt;/span&gt;&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="nc"&gt;Config&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt;
    &lt;span class="n"&gt;d_model&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt; &lt;span class="nb"&gt;int&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="mi"&gt;768&lt;/span&gt;
    &lt;span class="n"&gt;debug&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt; &lt;span class="nb"&gt;bool&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="kc"&gt;False&lt;/span&gt;  &lt;span class="c1"&gt;# per-module __main__ turns it on&lt;/span&gt;
    &lt;span class="n"&gt;layer_norm_eps&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt; &lt;span class="nb"&gt;float&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="mf"&gt;1e-5&lt;/span&gt;
    &lt;span class="n"&gt;d_vocab&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt; &lt;span class="nb"&gt;int&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="mi"&gt;50257&lt;/span&gt;
    &lt;span class="n"&gt;init_range&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt; &lt;span class="nb"&gt;float&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="mf"&gt;0.02&lt;/span&gt;
    &lt;span class="n"&gt;n_ctx&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt; &lt;span class="nb"&gt;int&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="mi"&gt;1024&lt;/span&gt;
    &lt;span class="n"&gt;d_head&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt; &lt;span class="nb"&gt;int&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="mi"&gt;64&lt;/span&gt;
    &lt;span class="n"&gt;d_mlp&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt; &lt;span class="nb"&gt;int&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="mi"&gt;3072&lt;/span&gt;
    &lt;span class="n"&gt;n_heads&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt; &lt;span class="nb"&gt;int&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="mi"&gt;12&lt;/span&gt;
    &lt;span class="n"&gt;n_layers&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt; &lt;span class="nb"&gt;int&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="mi"&gt;12&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;h2 id="embedding-module"&gt;Embedding module&lt;/h2&gt;
&lt;div class="mermaid"&gt;flowchart TB
    tokens["tokens&amp;lt;br/&amp;gt;(batch, posn)&amp;lt;br/&amp;gt;&amp;lt;i&amp;gt;token ids&amp;lt;/i&amp;gt;"]

    tokens --&amp;gt;|"W_E[tokens]"| embed["Embed out&amp;lt;br/&amp;gt;(batch, posn, d_model)"]
    tokens --&amp;gt;|"shape only"| pos["PosEmbed out&amp;lt;br/&amp;gt;(batch, posn, d_model)"]

    embed --&amp;gt; sum(("+"))
    pos --&amp;gt; sum
    sum --&amp;gt; resid["residual stream&amp;lt;br/&amp;gt;(batch, posn, d_model)"]&lt;/div&gt;
&lt;p&gt;Integers become vectors here, and the two results are added to form the residual stream that every later layer reads from and writes back into.&lt;/p&gt;
&lt;p&gt;&lt;code&gt;The cat sat on the mat&lt;/code&gt; is tokenized first, in the shape &lt;code&gt;(batch, posn) = (1, 7)&lt;/code&gt;.&lt;/p&gt;
&lt;div class="language-text highlight"&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;ids             50256    464    3797    3332    319     262    2603
decoded &amp;lt;|endoftext|&amp;gt;  &amp;#39;The&amp;#39;  &amp;#39; cat&amp;#39;  &amp;#39; sat&amp;#39;  &amp;#39; on&amp;#39;  &amp;#39; the&amp;#39;  &amp;#39; mat&amp;#39;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;&lt;code&gt;50256&lt;/code&gt; is &lt;code&gt;&amp;lt;|endoftext|&amp;gt;&lt;/code&gt;, prepended as the beginning-of-sequence marker&lt;/li&gt;
&lt;/ul&gt;
&lt;h3 id="embed"&gt;&lt;code&gt;Embed&lt;/code&gt;&lt;/h3&gt;
&lt;p&gt;Each token id is replaced by its row of &lt;code&gt;W_E&lt;/code&gt; through lookup.&lt;/p&gt;
&lt;div class="language-python highlight"&gt;&lt;span class="code-lang"&gt;python&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;W_E&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;nn&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;Parameter&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;t&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;empty&lt;/span&gt;&lt;span class="p"&gt;((&lt;/span&gt;&lt;span class="n"&gt;cfg&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;d_vocab&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;cfg&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;d_model&lt;/span&gt;&lt;span class="p"&gt;)))&lt;/span&gt;

&lt;span class="n"&gt;out&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;W_E&lt;/span&gt;&lt;span class="p"&gt;[&lt;/span&gt;&lt;span class="n"&gt;tokens&lt;/span&gt;&lt;span class="p"&gt;]&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;div class="language-text highlight"&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;(d_vocab, d_model)[(batch, posn)]  -&amp;gt;  (batch, posn, d_model)
 50257    768       1      7            1      7     768
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;h3 id="posembed"&gt;&lt;code&gt;PosEmbed&lt;/code&gt;&lt;/h3&gt;
&lt;p&gt;Only the shape of &lt;code&gt;tokens&lt;/code&gt; is read, because position embeddings encode where a token sits rather than which token it is.&lt;/p&gt;
&lt;div class="language-python highlight"&gt;&lt;span class="code-lang"&gt;python&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;W_pos&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;nn&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;Parameter&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;t&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;empty&lt;/span&gt;&lt;span class="p"&gt;((&lt;/span&gt;&lt;span class="n"&gt;cfg&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;n_ctx&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;cfg&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;d_model&lt;/span&gt;&lt;span class="p"&gt;)))&lt;/span&gt;

&lt;span class="n"&gt;batch&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;seq_len&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;tokens&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;shape&lt;/span&gt;
&lt;span class="n"&gt;sliced&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;W_pos&lt;/span&gt;&lt;span class="p"&gt;[:&lt;/span&gt;&lt;span class="n"&gt;seq_len&lt;/span&gt;&lt;span class="p"&gt;]&lt;/span&gt;
&lt;span class="n"&gt;out&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;einops&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;repeat&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;sliced&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="s2"&gt;&amp;quot;seq d_model -&amp;gt; batch seq d_model&amp;quot;&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;batch&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="n"&gt;batch&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;&lt;code&gt;W_pos&lt;/code&gt; is an &lt;code&gt;(n_ctx, d_model)&lt;/code&gt; matrix, and its &lt;code&gt;n_ctx&lt;/code&gt; rows are sliced down to &lt;code&gt;seq_len&lt;/code&gt; = &lt;code&gt;posn&lt;/code&gt; = 7&lt;/li&gt;
&lt;li&gt;&lt;code&gt;repeat&lt;/code&gt; copies that sliced table for every sequence in the batch&lt;/li&gt;
&lt;/ul&gt;
&lt;h3 id="module-output"&gt;Module output&lt;/h3&gt;
&lt;div class="language-bash highlight"&gt;&lt;span class="code-lang"&gt;shell&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="o"&gt;(&lt;/span&gt;.venv&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt; &lt;/span&gt;dom@dom-z13:~/Desktop/PvML$&lt;span class="w"&gt; &lt;/span&gt;python&lt;span class="w"&gt; &lt;/span&gt;-m&lt;span class="w"&gt; &lt;/span&gt;pvml.modules.embedding
&lt;span class="nv"&gt;text&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;The cat sat on the mat&amp;#39;&lt;/span&gt;
ref.to_str_tokens&lt;span class="o"&gt;(&lt;/span&gt;text&lt;span class="o"&gt;)=[&lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;&amp;lt;|endoftext|&amp;gt;&amp;#39;&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;The&amp;#39;&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39; cat&amp;#39;&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39; sat&amp;#39;&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39; on&amp;#39;&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39; the&amp;#39;&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39; mat&amp;#39;&lt;/span&gt;&lt;span class="o"&gt;]&lt;/span&gt;
tokens.shape&lt;span class="o"&gt;=&lt;/span&gt;torch.Size&lt;span class="o"&gt;([&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;&lt;span class="o"&gt;])&lt;/span&gt;

&lt;span class="o"&gt;[&lt;/span&gt;embed&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                       &lt;/span&gt;&lt;span class="k"&gt;in&lt;/span&gt;:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;          &lt;/span&gt;&lt;span class="c1"&gt;# token ids, not activations&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;embed&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                      &lt;/span&gt;W_E:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;50257&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="c1"&gt;# one row per vocab entry&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;embed&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                      &lt;/span&gt;out:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;pos_embed&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                   &lt;/span&gt;&lt;span class="k"&gt;in&lt;/span&gt;:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;          &lt;/span&gt;&lt;span class="c1"&gt;# values ignored, only the shape is read&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;pos_embed&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                &lt;/span&gt;W_pos:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1024&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;     &lt;/span&gt;&lt;span class="c1"&gt;# one row per position, up to n_ctx rows&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;pos_embed&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;               &lt;/span&gt;sliced:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;        &lt;/span&gt;&lt;span class="c1"&gt;# W_pos[:seq_len]&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;pos_embed&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                  &lt;/span&gt;out:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;     &lt;/span&gt;&lt;span class="c1"&gt;# same table repeated for each sequence&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;h2 id="layernorm"&gt;LayerNorm&lt;/h2&gt;
&lt;div class="mermaid"&gt;flowchart TB
    resid["residual&amp;lt;br/&amp;gt;(batch, posn, d_model)"]

    resid --&amp;gt; mean["residual_mean&amp;lt;br/&amp;gt;(batch, posn, 1)"]
    resid --&amp;gt; std["residual_std&amp;lt;br/&amp;gt;(batch, posn, 1)"]

    resid --&amp;gt; norm["(residual - mean) / std&amp;lt;br/&amp;gt;(batch, posn, d_model)"]
    mean --&amp;gt; norm
    std --&amp;gt; norm

    norm --&amp;gt;|"* w + b"| out["out&amp;lt;br/&amp;gt;(batch, posn, d_model)"]&lt;/div&gt;
&lt;p&gt;A residual stream of shape &lt;code&gt;(batch, posn, d_model)&lt;/code&gt; comes in, and each position's &lt;code&gt;d_model&lt;/code&gt; vector is standardised on its own by &lt;code&gt;(residual - mean) / sqrt(var + eps)&lt;/code&gt;, with the mean and variance taken across &lt;code&gt;d_model&lt;/code&gt; rather than across positions or the batch. The learned &lt;code&gt;w&lt;/code&gt; and &lt;code&gt;b&lt;/code&gt; then scale and shift it, and the normalised residual stream leaves at the shape it arrived in.&lt;/p&gt;
&lt;h3 id="statistics"&gt;Statistics&lt;/h3&gt;
&lt;p&gt;The two numbers each position gets standardised by.&lt;/p&gt;
&lt;div class="language-python highlight"&gt;&lt;span class="code-lang"&gt;python&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="n"&gt;residual_mean&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;residual&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;mean&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;dim&lt;/span&gt;&lt;span class="o"&gt;=-&lt;/span&gt;&lt;span class="mi"&gt;1&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;keepdim&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="kc"&gt;True&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;span class="n"&gt;residual_std&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;residual&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;var&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;dim&lt;/span&gt;&lt;span class="o"&gt;=-&lt;/span&gt;&lt;span class="mi"&gt;1&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;keepdim&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="kc"&gt;True&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;unbiased&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="kc"&gt;False&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt; &lt;span class="o"&gt;+&lt;/span&gt; &lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;cfg&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;layer_norm_eps&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;sqrt&lt;/span&gt;&lt;span class="p"&gt;()&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;Reduced along &lt;code&gt;d_model&lt;/code&gt;, deriving both the mean and the std in the shape &lt;code&gt;(batch, posn, 1)&lt;/code&gt; = &lt;code&gt;(1, 7, 1)&lt;/code&gt;. Each of those 7 values is computed from one position's own 768 numbers&lt;/li&gt;
&lt;li&gt;&lt;code&gt;keepdim=True&lt;/code&gt; leaves the reduced axis as size 1 so it broadcasts back: &lt;code&gt;(1, 7, 768) - (1, 7, 1)&lt;/code&gt; works, &lt;code&gt;(1, 7, 768) - (1, 7)&lt;/code&gt; does not&lt;/li&gt;
&lt;li&gt;&lt;code&gt;unbiased=False&lt;/code&gt; divides by &lt;code&gt;N&lt;/code&gt;, not &lt;code&gt;N-1&lt;/code&gt;, which is 768 not 767&lt;/li&gt;
&lt;li&gt;&lt;code&gt;layer_norm_eps&lt;/code&gt; is a tiny constant, 1e-5, added to the variance before the sqrt so the denominator can never be zero. A position whose 768 numbers are all identical has a variance of 0, and dividing by that would give NaN&lt;/li&gt;
&lt;/ul&gt;
&lt;h3 id="scale-and-shift"&gt;Scale and shift&lt;/h3&gt;
&lt;p&gt;The learned half of the layer, applied once the standardising is done.&lt;/p&gt;
&lt;div class="language-python highlight"&gt;&lt;span class="code-lang"&gt;python&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;w&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;nn&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;Parameter&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;t&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;ones&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;cfg&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;d_model&lt;/span&gt;&lt;span class="p"&gt;))&lt;/span&gt;
&lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;b&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;nn&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;Parameter&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;t&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;zeros&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;cfg&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;d_model&lt;/span&gt;&lt;span class="p"&gt;))&lt;/span&gt;

&lt;span class="n"&gt;residual&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;residual&lt;/span&gt; &lt;span class="o"&gt;-&lt;/span&gt; &lt;span class="n"&gt;residual_mean&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt; &lt;span class="o"&gt;/&lt;/span&gt; &lt;span class="n"&gt;residual_std&lt;/span&gt;
&lt;span class="n"&gt;out&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;residual&lt;/span&gt; &lt;span class="o"&gt;*&lt;/span&gt; &lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;w&lt;/span&gt; &lt;span class="o"&gt;+&lt;/span&gt; &lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;b&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;&lt;code&gt;w&lt;/code&gt; and &lt;code&gt;b&lt;/code&gt; are both &lt;code&gt;(d_model,)&lt;/code&gt; = &lt;code&gt;(768,)&lt;/code&gt;, and they start as all ones and all zeros, so before any training &lt;code&gt;out&lt;/code&gt; is exactly the standardised residual and the layer does nothing of its own&lt;/li&gt;
&lt;li&gt;One &lt;code&gt;w&lt;/code&gt; and one &lt;code&gt;b&lt;/code&gt; per LayerNorm, reused at all 7 positions, while the mean and std were worked out per position. The normalisation is local, the learned correction on top of it is not&lt;/li&gt;
&lt;li&gt;Broadcasting pads the missing leading axes: in &lt;code&gt;(1, 7, 768) * (768,)&lt;/code&gt; the &lt;code&gt;w&lt;/code&gt; is read as &lt;code&gt;(1, 1, 768)&lt;/code&gt; and stretched to &lt;code&gt;(1, 7, 768)&lt;/code&gt;&lt;/li&gt;
&lt;li&gt;Standardising throws the input's scale away, so &lt;code&gt;[1, 2, 3, 4]&lt;/code&gt; and &lt;code&gt;[10, 20, 30, 40]&lt;/code&gt; come out identical and after centering only the direction survives. &lt;code&gt;w&lt;/code&gt; and &lt;code&gt;b&lt;/code&gt; put a scale back, one the model has learned rather than one the input happened to arrive with&lt;/li&gt;
&lt;/ul&gt;
&lt;h3 id="module-output_1"&gt;Module output&lt;/h3&gt;
&lt;div class="language-bash highlight"&gt;&lt;span class="code-lang"&gt;shell&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="o"&gt;(&lt;/span&gt;.venv&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt; &lt;/span&gt;dom@dom-z13:~/Desktop/PvML$&lt;span class="w"&gt; &lt;/span&gt;python&lt;span class="w"&gt; &lt;/span&gt;-m&lt;span class="w"&gt; &lt;/span&gt;pvml.modules.normalization
&lt;span class="o"&gt;[&lt;/span&gt;ln&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                    &lt;/span&gt;residual:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;     &lt;/span&gt;&lt;span class="c1"&gt;# as it arrives&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;ln&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;               &lt;/span&gt;residual_mean:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;ln&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                &lt;/span&gt;residual_std:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;ln&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                    &lt;/span&gt;residual:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;     &lt;/span&gt;&lt;span class="c1"&gt;# reassigned: (residual - mean) / std&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;ln&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                        &lt;/span&gt;w,&lt;span class="w"&gt; &lt;/span&gt;b:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;,&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;          &lt;/span&gt;&lt;span class="c1"&gt;# padded to (1, 1, d_model), stretched to out&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;ln&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                         &lt;/span&gt;out:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;     &lt;/span&gt;&lt;span class="c1"&gt;# * w + b&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;h2 id="attention-module"&gt;Attention module&lt;/h2&gt;
&lt;div class="mermaid"&gt;flowchart TB
    resid["normalized_resid_pre&amp;lt;br/&amp;gt;(batch, posn, d_model)"]

    resid --&amp;gt;|"@ W_Q + b_Q"| q["q&amp;lt;br/&amp;gt;(batch, posn, nheads, d_head)"]
    resid --&amp;gt;|"@ W_K + b_K"| k["k&amp;lt;br/&amp;gt;(batch, posn, nheads, d_head)"]
    resid --&amp;gt;|"@ W_V + b_V"| v["v&amp;lt;br/&amp;gt;(batch, posn, nheads, d_head)"]

    q --&amp;gt; scores["attn_scores&amp;lt;br/&amp;gt;(batch, nheads, posn_Q, posn_K)&amp;lt;br/&amp;gt;&amp;lt;i&amp;gt;contracts d_head&amp;lt;/i&amp;gt;"]
    k --&amp;gt; scores

    scores --&amp;gt; scale["/ sqrt(d_head)"]
    scale --&amp;gt; mask["apply_causal_mask&amp;lt;br/&amp;gt;upper triangle to -inf"]
    mask --&amp;gt; soft["softmax(-1)"]
    soft --&amp;gt; pattern["attn_pattern&amp;lt;br/&amp;gt;(batch, nheads, posn_Q, posn_K)&amp;lt;br/&amp;gt;&amp;lt;i&amp;gt;rows sum to 1&amp;lt;/i&amp;gt;"]

    pattern --&amp;gt; z["z&amp;lt;br/&amp;gt;(batch, posn, nheads, d_head)&amp;lt;br/&amp;gt;&amp;lt;i&amp;gt;contracts posn_K&amp;lt;/i&amp;gt;"]
    v --&amp;gt; z

    z --&amp;gt;|"@ W_O + b_O"| out["attn_out&amp;lt;br/&amp;gt;(batch, posn, d_model)&amp;lt;br/&amp;gt;&amp;lt;i&amp;gt;contracts nheads and d_head&amp;lt;/i&amp;gt;"]&lt;/div&gt;
&lt;p&gt;Each position asks what it needs (&lt;code&gt;q&lt;/code&gt;), every earlier position advertises what it has (&lt;code&gt;k&lt;/code&gt;), and the match decides how much of their content (&lt;code&gt;v&lt;/code&gt;) gets copied back.&lt;/p&gt;
&lt;h3 id="q-k-v"&gt;&lt;code&gt;q, k, v&lt;/code&gt;&lt;/h3&gt;
&lt;p&gt;All three read the same input through three different learned matrices. Same einsum, same output shape, different weights.&lt;/p&gt;
&lt;div class="language-python highlight"&gt;&lt;span class="code-lang"&gt;python&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="n"&gt;q&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="p"&gt;(&lt;/span&gt;
    &lt;span class="n"&gt;einops&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;einsum&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;
        &lt;span class="n"&gt;normalized_resid_pre&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt;
        &lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;W_Q&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt;
        &lt;span class="s2"&gt;&amp;quot;batch posn d_model, nheads d_model d_head -&amp;gt; batch posn nheads d_head&amp;quot;&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt;
    &lt;span class="p"&gt;)&lt;/span&gt;
    &lt;span class="o"&gt;+&lt;/span&gt; &lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;b_Q&lt;/span&gt;
&lt;span class="p"&gt;)&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;p&gt;&lt;code&gt;k&lt;/code&gt; and &lt;code&gt;v&lt;/code&gt; are the same three lines again, reading the same &lt;code&gt;normalized_resid_pre&lt;/code&gt; through &lt;code&gt;W_K&lt;/code&gt;, &lt;code&gt;b_K&lt;/code&gt; and &lt;code&gt;W_V&lt;/code&gt;, &lt;code&gt;b_V&lt;/code&gt;.&lt;/p&gt;
&lt;div class="language-text highlight"&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;batch posn d_model,  nheads d_model d_head  -&amp;gt;  batch posn nheads d_head
1     7    768       12     768     64          1     7     12    64
           ▲                ▲
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;&lt;code&gt;d_model&lt;/code&gt; appears in both inputs and not in the output, so it is summed away. Each position's 768 numbers become 64 per head&lt;/li&gt;
&lt;li&gt;&lt;code&gt;nheads&lt;/code&gt; and &lt;code&gt;d_head&lt;/code&gt; come from the weight, &lt;code&gt;batch&lt;/code&gt; and &lt;code&gt;posn&lt;/code&gt; pass straight through&lt;/li&gt;
&lt;li&gt;&lt;code&gt;W_Q&lt;/code&gt; is &lt;code&gt;(nheads, d_model, d_head)&lt;/code&gt; = &lt;code&gt;(12, 768, 64)&lt;/code&gt;, one projection per head stacked on a leading axis, so all 12 heads run as a single batched matmul rather than 12 separate modules&lt;/li&gt;
&lt;li&gt;&lt;code&gt;b_Q&lt;/code&gt; is &lt;code&gt;(nheads, d_head)&lt;/code&gt; = &lt;code&gt;(12, 64)&lt;/code&gt;, read as &lt;code&gt;(1, 1, 12, 64)&lt;/code&gt; and stretched: the same bias at every position, a different one per head&lt;/li&gt;
&lt;/ul&gt;
&lt;h3 id="attn_scores"&gt;&lt;code&gt;attn_scores&lt;/code&gt;&lt;/h3&gt;
&lt;p&gt;For each head, dot every query vector (&lt;code&gt;q&lt;/code&gt;) with every key vector (&lt;code&gt;k&lt;/code&gt;) over &lt;code&gt;d_head&lt;/code&gt;, giving a &lt;code&gt;(posn_Q, posn_K)&lt;/code&gt; grid of match scores per head.&lt;/p&gt;
&lt;div class="language-python highlight"&gt;&lt;span class="code-lang"&gt;python&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="n"&gt;attn_scores&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;einops&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;einsum&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;
    &lt;span class="n"&gt;q&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt;
    &lt;span class="n"&gt;k&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt;
    &lt;span class="s2"&gt;&amp;quot;batch posn_Q nheads d_head, batch posn_K nheads d_head -&amp;gt; batch nheads posn_Q posn_K&amp;quot;&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt;
&lt;span class="p"&gt;)&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;&lt;code&gt;d_head&lt;/code&gt; is contracted, so the 64 numbers of a query and a key collapse to one score&lt;/li&gt;
&lt;li&gt;&lt;code&gt;batch&lt;/code&gt; and &lt;code&gt;nheads&lt;/code&gt; appear in both inputs &lt;em&gt;and&lt;/em&gt; the output, so they are batched over rather than summed: an independent 7 x 7 grid per sequence per head&lt;/li&gt;
&lt;li&gt;The two position axes get different names so they survive as separate dimensions. &lt;code&gt;posn_Q&lt;/code&gt; is who is asking, &lt;code&gt;posn_K&lt;/code&gt; is who is being read&lt;/li&gt;
&lt;li&gt;These scores are raw. For entries of roughly unit variance the dot product of two 64-dimensional vectors has a standard deviation that grows like &lt;code&gt;sqrt(d_head)&lt;/code&gt;, so dividing by &lt;code&gt;sqrt(64)&lt;/code&gt; = 8 is exactly what cancels it. That pulls them back to a range where softmax does not saturate into a near one-hot row and kill the gradients&lt;/li&gt;
&lt;/ul&gt;
&lt;h3 id="apply_causal_maskattn_scores"&gt;&lt;code&gt;apply_causal_mask(attn_scores)&lt;/code&gt;&lt;/h3&gt;
&lt;p&gt;Set everything above the diagonal of the scaled scores to &lt;code&gt;-inf&lt;/code&gt;, so no position can attend to a later one.&lt;/p&gt;
&lt;div class="language-python highlight"&gt;&lt;span class="code-lang"&gt;python&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="n"&gt;all_ones&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;t&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;ones&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;attn_scores&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;size&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="o"&gt;-&lt;/span&gt;&lt;span class="mi"&gt;2&lt;/span&gt;&lt;span class="p"&gt;),&lt;/span&gt; &lt;span class="n"&gt;attn_scores&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;size&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="o"&gt;-&lt;/span&gt;&lt;span class="mi"&gt;1&lt;/span&gt;&lt;span class="p"&gt;),&lt;/span&gt; &lt;span class="n"&gt;device&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="n"&gt;attn_scores&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;device&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;span class="n"&gt;mask&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;t&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;triu&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;all_ones&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;diagonal&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="mi"&gt;1&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;bool&lt;/span&gt;&lt;span class="p"&gt;()&lt;/span&gt;
&lt;span class="n"&gt;attn_scores&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;masked_fill_&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;mask&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;IGNORE&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;div class="language-text highlight"&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;        k0 k1 k2 k3 k4 k5 k6
    q0 [ 0  1  1  1  1  1  1 ]    1 = masked, the future
    q1 [ 0  0  1  1  1  1  1 ]
    q2 [ 0  0  0  1  1  1  1 ]
    q3 [ 0  0  0  0  1  1  1 ]
    q4 [ 0  0  0  0  0  1  1 ]
    q5 [ 0  0  0  0  0  0  1 ]
    q6 [ 0  0  0  0  0  0  0 ]
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;&lt;code&gt;diagonal=1&lt;/code&gt; leaves the diagonal unmasked, so a position can attend to itself&lt;/li&gt;
&lt;li&gt;The mask is 2-D &lt;code&gt;(7, 7)&lt;/code&gt; and broadcasts over &lt;code&gt;batch&lt;/code&gt; and &lt;code&gt;nheads&lt;/code&gt;. Causality depends only on position, so every sequence and every head gets the same one&lt;/li&gt;
&lt;li&gt;&lt;code&gt;IGNORE&lt;/code&gt; is a registered buffer, so &lt;code&gt;-inf&lt;/code&gt; follows &lt;code&gt;.to(device)&lt;/code&gt;. &lt;code&gt;all_ones&lt;/code&gt; needs an explicit &lt;code&gt;device=&lt;/code&gt; because a tensor made inside a method does not&lt;/li&gt;
&lt;li&gt;&lt;code&gt;masked_fill_&lt;/code&gt; is in place: it mutates &lt;code&gt;attn_scores&lt;/code&gt;&lt;/li&gt;
&lt;/ul&gt;
&lt;h3 id="attn_pattern"&gt;&lt;code&gt;attn_pattern&lt;/code&gt;&lt;/h3&gt;
&lt;p&gt;Softmax each row into a distribution over the keys it can see.&lt;/p&gt;
&lt;div class="language-python highlight"&gt;&lt;span class="code-lang"&gt;python&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="n"&gt;attn_scores_masked&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;apply_causal_mask&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;attn_scores&lt;/span&gt; &lt;span class="o"&gt;/&lt;/span&gt; &lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;cfg&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;d_head&lt;/span&gt;&lt;span class="o"&gt;**&lt;/span&gt;&lt;span class="mf"&gt;0.5&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;span class="n"&gt;attn_pattern&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;attn_scores_masked&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;softmax&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="o"&gt;-&lt;/span&gt;&lt;span class="mi"&gt;1&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;The mask has to land before the softmax, not after, so &lt;code&gt;exp(-inf)&lt;/code&gt; = 0 never enters the denominator and the surviving weights still sum to 1&lt;/li&gt;
&lt;li&gt;Shape is unchanged throughout, &lt;code&gt;(1, 12, 7, 7)&lt;/code&gt;. Row &lt;code&gt;q3&lt;/code&gt; of each head now holds four non-zero weights summing to 1, and three zeros&lt;/li&gt;
&lt;/ul&gt;
&lt;h3 id="z"&gt;&lt;code&gt;z&lt;/code&gt;&lt;/h3&gt;
&lt;p&gt;Weighted average of &lt;code&gt;v&lt;/code&gt; for each query position, contracting &lt;code&gt;posn_K&lt;/code&gt;.&lt;/p&gt;
&lt;div class="language-python highlight"&gt;&lt;span class="code-lang"&gt;python&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="n"&gt;z&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;einops&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;einsum&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;
    &lt;span class="n"&gt;v&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt;
    &lt;span class="n"&gt;attn_pattern&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt;
    &lt;span class="s2"&gt;&amp;quot;batch posn_K nheads d_head, batch nheads posn_Q posn_K -&amp;gt; batch posn_Q nheads d_head&amp;quot;&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt;
&lt;span class="p"&gt;)&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;&lt;code&gt;posn_K&lt;/code&gt; is the contracted axis, so the sum runs over the positions being &lt;em&gt;read&lt;/em&gt;. &lt;code&gt;posn_Q&lt;/code&gt; survives, so the result is indexed by the position doing the asking&lt;/li&gt;
&lt;li&gt;Concretely, at query position 3: its row of &lt;code&gt;attn_pattern&lt;/code&gt; holds four non-zero weights, and &lt;code&gt;z&lt;/code&gt; at that position is those four value vectors each multiplied by its weight and added together, separately for each of the 12 heads. The three masked positions carry weight 0, so they contribute nothing&lt;/li&gt;
&lt;li&gt;&lt;code&gt;z&lt;/code&gt; comes out the same shape as &lt;code&gt;v&lt;/code&gt;, &lt;code&gt;(1, 7, 12, 64)&lt;/code&gt;. Nothing changed size. What changed is that each position's vector is now a blend of the positions it could see, instead of only its own&lt;/li&gt;
&lt;li&gt;This is the only step in the whole model where information crosses between positions. Embedding, LayerNorm and the MLP all act on each position independently, so every bit of context a token ever gets arrives through this one sum&lt;/li&gt;
&lt;/ul&gt;
&lt;h3 id="attn_out"&gt;&lt;code&gt;attn_out&lt;/code&gt;&lt;/h3&gt;
&lt;p&gt;Project each head's &lt;code&gt;d_head&lt;/code&gt; back up to &lt;code&gt;d_model&lt;/code&gt; and sum over the heads.&lt;/p&gt;
&lt;div class="language-python highlight"&gt;&lt;span class="code-lang"&gt;python&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="n"&gt;attn_out&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="p"&gt;(&lt;/span&gt;
    &lt;span class="n"&gt;einops&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;einsum&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;
        &lt;span class="n"&gt;z&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt;
        &lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;W_O&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt;
        &lt;span class="s2"&gt;&amp;quot;batch posn_Q nheads d_head, nheads d_head d_model -&amp;gt; batch posn_Q d_model&amp;quot;&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt;
    &lt;span class="p"&gt;)&lt;/span&gt;
    &lt;span class="o"&gt;+&lt;/span&gt; &lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;b_O&lt;/span&gt;
&lt;span class="p"&gt;)&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;&lt;code&gt;W_O&lt;/code&gt; is &lt;code&gt;(nheads, d_head, d_model)&lt;/code&gt; = &lt;code&gt;(12, 64, 768)&lt;/code&gt;, the reverse of &lt;code&gt;W_Q&lt;/code&gt;, &lt;code&gt;W_K&lt;/code&gt; and &lt;code&gt;W_V&lt;/code&gt;. Those project the residual stream down into each head's 64 dimensions, this one projects the result back up to 768&lt;/li&gt;
&lt;li&gt;Two names are contracted at once, &lt;code&gt;nheads&lt;/code&gt; and &lt;code&gt;d_head&lt;/code&gt;, so the 12 heads are &lt;strong&gt;summed, not concatenated&lt;/strong&gt;. Each head's 64 numbers become a full 768-wide vector, and the 12 vectors are added&lt;/li&gt;
&lt;li&gt;&lt;code&gt;b_O&lt;/code&gt; is &lt;code&gt;(d_model,)&lt;/code&gt; = &lt;code&gt;(768,)&lt;/code&gt;, with no head axis, because it is added once after the heads are already summed. Compare &lt;code&gt;b_Q&lt;/code&gt;, which is &lt;code&gt;(12, 64)&lt;/code&gt; and applies per head&lt;/li&gt;
&lt;li&gt;Output is &lt;code&gt;(1, 7, 768)&lt;/code&gt;, back to the residual stream's width, ready to be added into it&lt;/li&gt;
&lt;/ul&gt;
&lt;h3 id="module-output_2"&gt;Module output&lt;/h3&gt;
&lt;div class="language-bash highlight"&gt;&lt;span class="code-lang"&gt;shell&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="o"&gt;(&lt;/span&gt;.venv&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt; &lt;/span&gt;dom@dom-z13:~/Desktop/PvML$&lt;span class="w"&gt; &lt;/span&gt;python&lt;span class="w"&gt; &lt;/span&gt;-m&lt;span class="w"&gt; &lt;/span&gt;pvml.modules.attention
&lt;span class="w"&gt; &lt;/span&gt;W_Q:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;12&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;64&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;
&lt;span class="w"&gt; &lt;/span&gt;W_K:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;12&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;64&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;
&lt;span class="w"&gt; &lt;/span&gt;W_V:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;12&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;64&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;
&lt;span class="w"&gt; &lt;/span&gt;W_O:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;12&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;64&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;
&lt;span class="w"&gt; &lt;/span&gt;b_Q:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;12&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;64&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;
&lt;span class="w"&gt; &lt;/span&gt;b_K:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;12&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;64&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;
&lt;span class="w"&gt; &lt;/span&gt;b_V:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;12&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;64&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;
&lt;span class="w"&gt; &lt;/span&gt;b_O:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;,&lt;span class="o"&gt;)&lt;/span&gt;
total:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;2&lt;/span&gt;,362,368&lt;span class="w"&gt; &lt;/span&gt;parameters

Loaded&lt;span class="w"&gt; &lt;/span&gt;pretrained&lt;span class="w"&gt; &lt;/span&gt;model&lt;span class="w"&gt; &lt;/span&gt;gpt2-small&lt;span class="w"&gt; &lt;/span&gt;into&lt;span class="w"&gt; &lt;/span&gt;HookedTransformer
&lt;span class="o"&gt;[&lt;/span&gt;attn&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;      &lt;/span&gt;normalized_resid_pre:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;     &lt;/span&gt;&lt;span class="c1"&gt;# (batch, posn, d_model)&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;attn&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                       &lt;/span&gt;W_Q:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;12&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;64&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;   &lt;/span&gt;&lt;span class="c1"&gt;# (n_heads, d_model, d_head)&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;attn&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                       &lt;/span&gt;b_Q:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;12&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;64&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;        &lt;/span&gt;&lt;span class="c1"&gt;# (n_heads, d_head), broadcast over posn&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;attn&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                         &lt;/span&gt;q:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;12&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;64&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;  &lt;/span&gt;&lt;span class="c1"&gt;# (batch, posn, n_heads, d_head)&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;attn&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                       &lt;/span&gt;W_K:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;12&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;64&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;attn&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                       &lt;/span&gt;b_K:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;12&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;64&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;attn&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                         &lt;/span&gt;k:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;12&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;64&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;attn&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                       &lt;/span&gt;W_V:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;12&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;64&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;attn&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                       &lt;/span&gt;b_V:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;12&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;64&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;attn&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                         &lt;/span&gt;v:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;12&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;64&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;attn&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;               &lt;/span&gt;attn_scores:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;12&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;   &lt;/span&gt;&lt;span class="c1"&gt;# (batch, n_heads, query_pos, key_pos)&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;attn&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                      &lt;/span&gt;mask:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;          &lt;/span&gt;&lt;span class="c1"&gt;# True above the diagonal = cannot attend&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;attn&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;              &lt;/span&gt;attn_pattern:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;12&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;   &lt;/span&gt;&lt;span class="c1"&gt;# rows sum to 1 over visible keys&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;attn&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                         &lt;/span&gt;z:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;12&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;64&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;  &lt;/span&gt;&lt;span class="c1"&gt;# weighted average of v, per query position&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;attn&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                       &lt;/span&gt;W_O:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;12&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;64&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;   &lt;/span&gt;&lt;span class="c1"&gt;# (n_heads, d_head, d_model), projects back up&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;attn&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                  &lt;/span&gt;attn_out:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;     &lt;/span&gt;&lt;span class="c1"&gt;# back to the residual stream&amp;#39;s width&lt;/span&gt;

head&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;&lt;span class="w"&gt; &lt;/span&gt;attention&lt;span class="w"&gt; &lt;/span&gt;pattern
&lt;span class="w"&gt;   &lt;/span&gt;query&lt;span class="w"&gt; &lt;/span&gt;/&lt;span class="w"&gt; &lt;/span&gt;key&lt;span class="w"&gt;   &lt;/span&gt;&amp;lt;BOS&amp;gt;&lt;span class="w"&gt;     &lt;/span&gt;The&lt;span class="w"&gt;     &lt;/span&gt;cat&lt;span class="w"&gt;     &lt;/span&gt;sat&lt;span class="w"&gt;      &lt;/span&gt;on&lt;span class="w"&gt;     &lt;/span&gt;the&lt;span class="w"&gt;     &lt;/span&gt;mat
&lt;span class="w"&gt;         &lt;/span&gt;&amp;lt;BOS&amp;gt;&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;.00&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.00&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.00&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.00&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.00&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.00&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.00
&lt;span class="w"&gt;           &lt;/span&gt;The&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.93&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.07&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.00&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.00&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.00&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.00&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.00
&lt;span class="w"&gt;           &lt;/span&gt;cat&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.71&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.10&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.18&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.00&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.00&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.00&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.00
&lt;span class="w"&gt;           &lt;/span&gt;sat&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.64&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.14&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.04&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.18&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.00&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.00&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.00
&lt;span class="w"&gt;            &lt;/span&gt;on&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.48&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.15&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.12&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.23&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.03&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.00&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.00
&lt;span class="w"&gt;           &lt;/span&gt;the&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.60&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.12&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.09&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.16&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.02&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.02&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.00
&lt;span class="w"&gt;           &lt;/span&gt;mat&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.37&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.09&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.06&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.03&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.08&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.10&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.26
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;h2 id="mlp"&gt;MLP&lt;/h2&gt;
&lt;div class="mermaid"&gt;flowchart TB
    resid["normalized_resid_mid&amp;lt;br/&amp;gt;(batch, posn, d_model)"]

    resid --&amp;gt;|"@ W_in + b_in"| pre["pre&amp;lt;br/&amp;gt;(batch, posn, d_mlp)&amp;lt;br/&amp;gt;&amp;lt;i&amp;gt;contracts d_model&amp;lt;/i&amp;gt;"]
    pre --&amp;gt;|"gelu_new"| post["post&amp;lt;br/&amp;gt;(batch, posn, d_mlp)"]
    post --&amp;gt;|"@ W_out + b_out"| out["mlp_out&amp;lt;br/&amp;gt;(batch, posn, d_model)&amp;lt;br/&amp;gt;&amp;lt;i&amp;gt;contracts d_mlp&amp;lt;/i&amp;gt;"]&lt;/div&gt;
&lt;p&gt;Expand to four times the residual width, apply the nonlinearity, project back. Every position goes through the same matrices independently, so nothing here mixes positions.&lt;/p&gt;
&lt;h3 id="pre"&gt;&lt;code&gt;pre&lt;/code&gt;&lt;/h3&gt;
&lt;p&gt;The first of the two projections, into the wide space.&lt;/p&gt;
&lt;div class="language-python highlight"&gt;&lt;span class="code-lang"&gt;python&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;W_in&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;nn&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;Parameter&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;t&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;empty&lt;/span&gt;&lt;span class="p"&gt;((&lt;/span&gt;&lt;span class="n"&gt;cfg&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;d_model&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;cfg&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;d_mlp&lt;/span&gt;&lt;span class="p"&gt;)))&lt;/span&gt;

&lt;span class="n"&gt;pre&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="p"&gt;(&lt;/span&gt;
    &lt;span class="n"&gt;einops&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;einsum&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;
        &lt;span class="n"&gt;normalized_resid_mid&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt;
        &lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;W_in&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt;
        &lt;span class="s2"&gt;&amp;quot;batch position d_model, d_model d_mlp -&amp;gt; batch position d_mlp&amp;quot;&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt;
    &lt;span class="p"&gt;)&lt;/span&gt;
    &lt;span class="o"&gt;+&lt;/span&gt; &lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;b_in&lt;/span&gt;
&lt;span class="p"&gt;)&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;&lt;code&gt;W_in&lt;/code&gt; is 2-D. The MLP has no head axis, unlike attention's weights&lt;/li&gt;
&lt;li&gt;&lt;code&gt;d_model&lt;/code&gt; contracts, &lt;code&gt;d_mlp&lt;/code&gt; appears: 768 in, 3072 out&lt;/li&gt;
&lt;li&gt;Input is the residual stream after attention has been added, through the block's &lt;code&gt;ln2&lt;/code&gt;&lt;/li&gt;
&lt;/ul&gt;
&lt;h3 id="post"&gt;&lt;code&gt;post&lt;/code&gt;&lt;/h3&gt;
&lt;p&gt;The only nonlinearity in the MLP.&lt;/p&gt;
&lt;div class="language-python highlight"&gt;&lt;span class="code-lang"&gt;python&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="n"&gt;post&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;gelu_new&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;pre&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;Without it &lt;code&gt;W_in&lt;/code&gt; and &lt;code&gt;W_out&lt;/code&gt; would collapse into a single matrix and the whole detour through 3072 dimensions would buy nothing&lt;/li&gt;
&lt;li&gt;Applied elementwise in the wide space, so the shape does not change&lt;/li&gt;
&lt;/ul&gt;
&lt;h3 id="mlp_out"&gt;&lt;code&gt;mlp_out&lt;/code&gt;&lt;/h3&gt;
&lt;p&gt;The second projection, undoing the expansion.&lt;/p&gt;
&lt;div class="language-python highlight"&gt;&lt;span class="code-lang"&gt;python&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;W_out&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;nn&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;Parameter&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;t&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;empty&lt;/span&gt;&lt;span class="p"&gt;((&lt;/span&gt;&lt;span class="n"&gt;cfg&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;d_mlp&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;cfg&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;d_model&lt;/span&gt;&lt;span class="p"&gt;)))&lt;/span&gt;

&lt;span class="n"&gt;mlp_out&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="p"&gt;(&lt;/span&gt;
    &lt;span class="n"&gt;einops&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;einsum&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;
        &lt;span class="n"&gt;post&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt;
        &lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;W_out&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt;
        &lt;span class="s2"&gt;&amp;quot;batch position d_mlp, d_mlp d_model -&amp;gt; batch position d_model&amp;quot;&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt;
    &lt;span class="p"&gt;)&lt;/span&gt;
    &lt;span class="o"&gt;+&lt;/span&gt; &lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;b_out&lt;/span&gt;
&lt;span class="p"&gt;)&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;&lt;code&gt;d_mlp&lt;/code&gt; contracts: 3072 in, 768 out, back to the width the block adds into&lt;/li&gt;
&lt;li&gt;&lt;code&gt;b_out&lt;/code&gt; is &lt;code&gt;(d_model,)&lt;/code&gt;, broadcast over batch and position&lt;/li&gt;
&lt;/ul&gt;
&lt;h3 id="module-output_3"&gt;Module output&lt;/h3&gt;
&lt;div class="language-bash highlight"&gt;&lt;span class="code-lang"&gt;shell&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="o"&gt;(&lt;/span&gt;.venv&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt; &lt;/span&gt;dom@dom-z13:~/Desktop/PvML$&lt;span class="w"&gt; &lt;/span&gt;python&lt;span class="w"&gt; &lt;/span&gt;-m&lt;span class="w"&gt; &lt;/span&gt;pvml.modules.mlp
&lt;span class="w"&gt;  &lt;/span&gt;W_in:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;3072&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;
&lt;span class="w"&gt; &lt;/span&gt;W_out:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;3072&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;
&lt;span class="w"&gt;  &lt;/span&gt;b_in:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;3072&lt;/span&gt;,&lt;span class="o"&gt;)&lt;/span&gt;
&lt;span class="w"&gt; &lt;/span&gt;b_out:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;,&lt;span class="o"&gt;)&lt;/span&gt;
&lt;span class="w"&gt; &lt;/span&gt;total:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;4&lt;/span&gt;,722,432&lt;span class="w"&gt; &lt;/span&gt;parameters

&lt;span class="o"&gt;[&lt;/span&gt;mlp&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;       &lt;/span&gt;normalized_resid_mid:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;     &lt;/span&gt;&lt;span class="c1"&gt;# (batch, posn, d_model)&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;mlp&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                       &lt;/span&gt;W_in:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;3072&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;     &lt;/span&gt;&lt;span class="c1"&gt;# (d_model, d_mlp), contracts d_model&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;mlp&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                        &lt;/span&gt;pre:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;3072&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="c1"&gt;# 4x wider than the residual stream&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;mlp&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                       &lt;/span&gt;post:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;3072&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="c1"&gt;# gelu_new, same shape&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;mlp&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                      &lt;/span&gt;W_out:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;3072&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;     &lt;/span&gt;&lt;span class="c1"&gt;# (d_mlp, d_model), contracts d_mlp&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;mlp&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                    &lt;/span&gt;mlp_out:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;     &lt;/span&gt;&lt;span class="c1"&gt;# back to (batch, posn, d_model)&lt;/span&gt;

max&lt;span class="w"&gt; &lt;/span&gt;abs&lt;span class="w"&gt; &lt;/span&gt;diff&lt;span class="w"&gt; &lt;/span&gt;vs&lt;span class="w"&gt; &lt;/span&gt;gpt-2:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.000e+00
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;4.7M parameters, twice what attention costs. Two thirds of a block is the MLP&lt;/li&gt;
&lt;li&gt;&lt;code&gt;d_mlp = 4 * d_model&lt;/code&gt; is convention, not a requirement&lt;/li&gt;
&lt;/ul&gt;
&lt;h2 id="transformerblock"&gt;TransformerBlock&lt;/h2&gt;
&lt;div class="mermaid"&gt;flowchart TB
    pre["resid_pre&amp;lt;br/&amp;gt;(batch, posn, d_model)"]

    pre --&amp;gt; ln1["ln1&amp;lt;br/&amp;gt;LayerNorm"]
    ln1 --&amp;gt; attn["attn&amp;lt;br/&amp;gt;takes normalized_resid_pre"]
    attn --&amp;gt; add1(("+"))
    pre ---&amp;gt; add1

    add1 --&amp;gt; mid["resid_mid&amp;lt;br/&amp;gt;(batch, posn, d_model)"]

    mid --&amp;gt; ln2["ln2&amp;lt;br/&amp;gt;LayerNorm"]
    ln2 --&amp;gt; mlp["mlp&amp;lt;br/&amp;gt;takes normalized_resid_mid"]
    mlp --&amp;gt; add2(("+"))
    mid ---&amp;gt; add2

    add2 --&amp;gt; post["resid_post&amp;lt;br/&amp;gt;(batch, posn, d_model)"]&lt;/div&gt;
&lt;p&gt;Attention and the MLP each read the stream and add a delta back. Neither replaces it, so every component's contribution stays separable.&lt;/p&gt;
&lt;h3 id="resid_mid"&gt;&lt;code&gt;resid_mid&lt;/code&gt;&lt;/h3&gt;
&lt;p&gt;Attention reads the normalised stream and its output is added back.&lt;/p&gt;
&lt;div class="language-python highlight"&gt;&lt;span class="code-lang"&gt;python&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="n"&gt;resid_mid&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;attn&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;ln1&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;resid_pre&lt;/span&gt;&lt;span class="p"&gt;))&lt;/span&gt; &lt;span class="o"&gt;+&lt;/span&gt; &lt;span class="n"&gt;resid_pre&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;&lt;code&gt;ln1&lt;/code&gt; normalises what attention &lt;strong&gt;reads&lt;/strong&gt;. What gets added back is the unnormalised &lt;code&gt;attn_out&lt;/code&gt;, so the stream itself never passes through a LayerNorm&lt;/li&gt;
&lt;li&gt;The stream keeps its width the whole way: 768 in, 768 out, at every point in all 12 blocks&lt;/li&gt;
&lt;/ul&gt;
&lt;h3 id="resid_post"&gt;&lt;code&gt;resid_post&lt;/code&gt;&lt;/h3&gt;
&lt;p&gt;The MLP's turn, on the stream attention just wrote into.&lt;/p&gt;
&lt;div class="language-python highlight"&gt;&lt;span class="code-lang"&gt;python&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="n"&gt;resid_post&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;mlp&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;ln2&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;resid_mid&lt;/span&gt;&lt;span class="p"&gt;))&lt;/span&gt; &lt;span class="o"&gt;+&lt;/span&gt; &lt;span class="n"&gt;resid_mid&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;&lt;code&gt;ln2&lt;/code&gt; normalises what the MLP reads, exactly as &lt;code&gt;ln1&lt;/code&gt; did for attention&lt;/li&gt;
&lt;li&gt;&lt;code&gt;(1, 7, 3072)&lt;/code&gt; inside the MLP and &lt;code&gt;(1, 7, 12, 64)&lt;/code&gt; inside attention are private scratch space, projected back before anything is added&lt;/li&gt;
&lt;/ul&gt;
&lt;h3 id="module-output_4"&gt;Module output&lt;/h3&gt;
&lt;div class="language-bash highlight"&gt;&lt;span class="code-lang"&gt;shell&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="o"&gt;(&lt;/span&gt;.venv&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt; &lt;/span&gt;dom@dom-z13:~/Desktop/PvML$&lt;span class="w"&gt; &lt;/span&gt;python&lt;span class="w"&gt; &lt;/span&gt;-m&lt;span class="w"&gt; &lt;/span&gt;pvml.modules.block
block&lt;span class="w"&gt; &lt;/span&gt;parameters:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,087,872

&lt;span class="o"&gt;[&lt;/span&gt;block&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                &lt;/span&gt;resid_pre:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;     &lt;/span&gt;&lt;span class="c1"&gt;# (batch, posn, d_model)&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;ln1&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                   &lt;/span&gt;residual:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;     &lt;/span&gt;&lt;span class="c1"&gt;# as it arrives&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;ln1&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;              &lt;/span&gt;residual_mean:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;ln1&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;               &lt;/span&gt;residual_std:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;ln1&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                   &lt;/span&gt;residual:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;     &lt;/span&gt;&lt;span class="c1"&gt;# reassigned: (residual - mean) / std&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;ln1&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                       &lt;/span&gt;w,&lt;span class="w"&gt; &lt;/span&gt;b:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;,&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;          &lt;/span&gt;&lt;span class="c1"&gt;# padded to (1, 1, d_model), stretched to out&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;ln1&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                        &lt;/span&gt;out:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;     &lt;/span&gt;&lt;span class="c1"&gt;# * w + b&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;attn&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;      &lt;/span&gt;normalized_resid_pre:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;     &lt;/span&gt;&lt;span class="c1"&gt;# (batch, posn, d_model)&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;attn&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                         &lt;/span&gt;q:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;12&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;64&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;  &lt;/span&gt;&lt;span class="c1"&gt;# (batch, posn, n_heads, d_head)&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;attn&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;               &lt;/span&gt;attn_scores:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;12&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;   &lt;/span&gt;&lt;span class="c1"&gt;# (batch, n_heads, query_pos, key_pos)&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;attn&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;              &lt;/span&gt;attn_pattern:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;12&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;   &lt;/span&gt;&lt;span class="c1"&gt;# rows sum to 1 over visible keys&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;attn&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                         &lt;/span&gt;z:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;12&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;64&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;  &lt;/span&gt;&lt;span class="c1"&gt;# weighted average of v, per query position&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;attn&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                  &lt;/span&gt;attn_out:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;     &lt;/span&gt;&lt;span class="c1"&gt;# back to the residual stream&amp;#39;s width&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;block&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                &lt;/span&gt;resid_mid:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;     &lt;/span&gt;&lt;span class="c1"&gt;# resid_pre + attn_out&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;ln2&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                   &lt;/span&gt;residual:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;     &lt;/span&gt;&lt;span class="c1"&gt;# as it arrives&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;ln2&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                        &lt;/span&gt;out:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;     &lt;/span&gt;&lt;span class="c1"&gt;# * w + b&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;mlp&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;       &lt;/span&gt;normalized_resid_mid:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;     &lt;/span&gt;&lt;span class="c1"&gt;# (batch, posn, d_model)&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;mlp&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                        &lt;/span&gt;pre:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;3072&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="c1"&gt;# 4x wider than the residual stream&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;mlp&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                       &lt;/span&gt;post:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;3072&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="c1"&gt;# gelu_new, same shape&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;mlp&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                    &lt;/span&gt;mlp_out:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;     &lt;/span&gt;&lt;span class="c1"&gt;# back to (batch, posn, d_model)&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;block&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;               &lt;/span&gt;resid_post:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;     &lt;/span&gt;&lt;span class="c1"&gt;# resid_mid + mlp_out&lt;/span&gt;

max&lt;span class="w"&gt; &lt;/span&gt;abs&lt;span class="w"&gt; &lt;/span&gt;diff&lt;span class="w"&gt; &lt;/span&gt;vs&lt;span class="w"&gt; &lt;/span&gt;gpt-2:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;.144e-05
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;First module where the diff is not exactly zero. Values reach 120, so 1e-5 is float32 accumulation across the four sub-modules of a block, not an error&lt;/li&gt;
&lt;li&gt;&lt;code&gt;tag_tree&lt;/code&gt; names the children by position, so the two LayerNorms print as &lt;code&gt;ln1&lt;/code&gt; and &lt;code&gt;ln2&lt;/code&gt;&lt;/li&gt;
&lt;/ul&gt;
&lt;h2 id="unembed"&gt;Unembed&lt;/h2&gt;
&lt;div class="mermaid"&gt;flowchart TB
    resid["normalized_resid_final&amp;lt;br/&amp;gt;(batch, posn, d_model)"]

    resid --&amp;gt;|"@ W_U + b_U"| logits["logits&amp;lt;br/&amp;gt;(batch, posn, d_vocab)&amp;lt;br/&amp;gt;&amp;lt;i&amp;gt;contracts d_model&amp;lt;/i&amp;gt;"]&lt;/div&gt;
&lt;p&gt;One score per vocabulary entry, per position. This is where &lt;code&gt;d_model&lt;/code&gt; becomes &lt;code&gt;d_vocab&lt;/code&gt; and the model commits to a prediction.&lt;/p&gt;
&lt;h3 id="logits"&gt;&lt;code&gt;logits&lt;/code&gt;&lt;/h3&gt;
&lt;p&gt;A single matmul against a &lt;code&gt;(768, 50257)&lt;/code&gt; table, the mirror of &lt;code&gt;W_E&lt;/code&gt;.&lt;/p&gt;
&lt;div class="language-python highlight"&gt;&lt;span class="code-lang"&gt;python&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;W_U&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;nn&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;Parameter&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;t&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;empty&lt;/span&gt;&lt;span class="p"&gt;((&lt;/span&gt;&lt;span class="n"&gt;cfg&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;d_model&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt; &lt;span class="n"&gt;cfg&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;d_vocab&lt;/span&gt;&lt;span class="p"&gt;)))&lt;/span&gt;
&lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;b_U&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;nn&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;Parameter&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;t&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;zeros&lt;/span&gt;&lt;span class="p"&gt;((&lt;/span&gt;&lt;span class="n"&gt;cfg&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;d_vocab&lt;/span&gt;&lt;span class="p"&gt;)),&lt;/span&gt; &lt;span class="n"&gt;requires_grad&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;&lt;span class="kc"&gt;False&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;

&lt;span class="n"&gt;logits&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="p"&gt;(&lt;/span&gt;
    &lt;span class="n"&gt;einops&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;einsum&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;
        &lt;span class="n"&gt;normalized_resid_final&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt;
        &lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;W_U&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt;
        &lt;span class="s2"&gt;&amp;quot;batch posn d_model, d_model d_vocab -&amp;gt; batch posn d_vocab&amp;quot;&lt;/span&gt;&lt;span class="p"&gt;,&lt;/span&gt;
    &lt;span class="p"&gt;)&lt;/span&gt;
    &lt;span class="o"&gt;+&lt;/span&gt; &lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;b_U&lt;/span&gt;
&lt;span class="p"&gt;)&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;&lt;code&gt;b_U&lt;/code&gt; is frozen. GPT-2 has no unembedding bias, so it stays at zero and only exists to keep the shapes lined up&lt;/li&gt;
&lt;li&gt;Input is the stream after all twelve blocks, through &lt;code&gt;ln_final&lt;/code&gt;&lt;/li&gt;
&lt;li&gt;&lt;code&gt;d_vocab&lt;/code&gt; appears in only two places in the whole model, here and at &lt;code&gt;W_E&lt;/code&gt;. Everything between them is &lt;code&gt;d_model&lt;/code&gt;&lt;/li&gt;
&lt;/ul&gt;
&lt;h3 id="module-output_5"&gt;Module output&lt;/h3&gt;
&lt;div class="language-bash highlight"&gt;&lt;span class="code-lang"&gt;shell&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="o"&gt;(&lt;/span&gt;.venv&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt; &lt;/span&gt;dom@dom-z13:~/Desktop/PvML$&lt;span class="w"&gt; &lt;/span&gt;python&lt;span class="w"&gt; &lt;/span&gt;-m&lt;span class="w"&gt; &lt;/span&gt;pvml.modules.unembedding
&lt;span class="w"&gt; &lt;/span&gt;W_U:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;50257&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;  &lt;/span&gt;&lt;span class="nv"&gt;requires_grad&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;True
&lt;span class="w"&gt; &lt;/span&gt;b_U:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;50257&lt;/span&gt;,&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;  &lt;/span&gt;&lt;span class="nv"&gt;requires_grad&lt;/span&gt;&lt;span class="o"&gt;=&lt;/span&gt;False
total:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;38&lt;/span&gt;,647,633&lt;span class="w"&gt; &lt;/span&gt;parameters

&lt;span class="o"&gt;[&lt;/span&gt;unembed&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;   &lt;/span&gt;normalized_resid_final:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;     &lt;/span&gt;&lt;span class="c1"&gt;# (batch, posn, d_model)&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;unembed&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                      &lt;/span&gt;W_U:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;768&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;50257&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;    &lt;/span&gt;&lt;span class="c1"&gt;# (d_model, d_vocab), the mirror of W_E&lt;/span&gt;
&lt;span class="o"&gt;[&lt;/span&gt;unembed&lt;span class="o"&gt;]&lt;/span&gt;&lt;span class="w"&gt;                   &lt;/span&gt;logits:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;50257&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt;   &lt;/span&gt;&lt;span class="c1"&gt;# one score per vocab entry&lt;/span&gt;

max&lt;span class="w"&gt; &lt;/span&gt;abs&lt;span class="w"&gt; &lt;/span&gt;diff&lt;span class="w"&gt; &lt;/span&gt;vs&lt;span class="w"&gt; &lt;/span&gt;gpt-2:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;0&lt;/span&gt;.000e+00

next-token&lt;span class="w"&gt; &lt;/span&gt;predictions
&lt;span class="w"&gt;     &lt;/span&gt;&amp;lt;BOS&amp;gt;&lt;span class="w"&gt;  &lt;/span&gt;-&amp;gt;&lt;span class="w"&gt;  &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;\n&amp;#39;&lt;/span&gt;
&lt;span class="w"&gt;       &lt;/span&gt;The&lt;span class="w"&gt;  &lt;/span&gt;-&amp;gt;&lt;span class="w"&gt;  &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39; first&amp;#39;&lt;/span&gt;
&lt;span class="w"&gt;       &lt;/span&gt;cat&lt;span class="w"&gt;  &lt;/span&gt;-&amp;gt;&lt;span class="w"&gt;  &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39; was&amp;#39;&lt;/span&gt;
&lt;span class="w"&gt;       &lt;/span&gt;sat&lt;span class="w"&gt;  &lt;/span&gt;-&amp;gt;&lt;span class="w"&gt;  &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39; on&amp;#39;&lt;/span&gt;
&lt;span class="w"&gt;        &lt;/span&gt;on&lt;span class="w"&gt;  &lt;/span&gt;-&amp;gt;&lt;span class="w"&gt;  &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39; the&amp;#39;&lt;/span&gt;
&lt;span class="w"&gt;       &lt;/span&gt;the&lt;span class="w"&gt;  &lt;/span&gt;-&amp;gt;&lt;span class="w"&gt;  &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39; floor&amp;#39;&lt;/span&gt;
&lt;span class="w"&gt;       &lt;/span&gt;mat&lt;span class="w"&gt;  &lt;/span&gt;-&amp;gt;&lt;span class="w"&gt;  &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;,&amp;#39;&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;Every position predicts its own next token from one forward pass. That is what the causal mask is for&lt;/li&gt;
&lt;li&gt;&lt;code&gt;sat -&amp;gt; ' on'&lt;/code&gt; and &lt;code&gt;on -&amp;gt; ' the'&lt;/code&gt; are right. &lt;code&gt;the -&amp;gt; ' floor'&lt;/code&gt; is a guess, because position 5 cannot see &lt;code&gt;' mat'&lt;/code&gt;&lt;/li&gt;
&lt;/ul&gt;
&lt;h2 id="transformer"&gt;Transformer&lt;/h2&gt;
&lt;div class="mermaid"&gt;flowchart TB
    tokens["tokens&amp;lt;br/&amp;gt;(batch, posn)"]

    tokens --&amp;gt; embed["embed&amp;lt;br/&amp;gt;W_E[tokens]"]
    tokens --&amp;gt; pos["pos_embed&amp;lt;br/&amp;gt;W_pos[:seq_len]"]

    embed --&amp;gt; add(("+"))
    pos --&amp;gt; add

    add --&amp;gt; resid["residual stream&amp;lt;br/&amp;gt;(batch, posn, d_model)"]

    resid --&amp;gt; block["blocks[i]&amp;lt;br/&amp;gt;TransformerBlock"]
    block --&amp;gt;|"residual = block(residual)&amp;lt;br/&amp;gt;for i in range(n_layers)"| resid

    resid --&amp;gt;|"after 12 blocks"| lnf["ln_final&amp;lt;br/&amp;gt;LayerNorm"]
    lnf --&amp;gt; unembed["unembed&amp;lt;br/&amp;gt;@ W_U + b_U&amp;lt;br/&amp;gt;&amp;lt;i&amp;gt;contracts d_model&amp;lt;/i&amp;gt;"]
    unembed --&amp;gt; logits["logits&amp;lt;br/&amp;gt;(batch, posn, d_vocab)"]&lt;/div&gt;
&lt;p&gt;Token ids in, logits out. The module owns no parameters of its own, it just wires the others together.&lt;/p&gt;
&lt;h3 id="forward"&gt;&lt;code&gt;forward&lt;/code&gt;&lt;/h3&gt;
&lt;p&gt;Embed, run the blocks in sequence, normalise, unembed.&lt;/p&gt;
&lt;div class="language-python highlight"&gt;&lt;span class="code-lang"&gt;python&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="n"&gt;residual&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;embed&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;tokens&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt; &lt;span class="o"&gt;+&lt;/span&gt; &lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;pos_embed&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;tokens&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;

&lt;span class="k"&gt;for&lt;/span&gt; &lt;span class="n"&gt;block&lt;/span&gt; &lt;span class="ow"&gt;in&lt;/span&gt; &lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;blocks&lt;/span&gt;&lt;span class="p"&gt;:&lt;/span&gt;
    &lt;span class="n"&gt;residual&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="n"&gt;block&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;residual&lt;/span&gt;&lt;span class="p"&gt;)&lt;/span&gt;

&lt;span class="n"&gt;logits&lt;/span&gt; &lt;span class="o"&gt;=&lt;/span&gt; &lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;unembed&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="bp"&gt;self&lt;/span&gt;&lt;span class="o"&gt;.&lt;/span&gt;&lt;span class="n"&gt;ln_final&lt;/span&gt;&lt;span class="p"&gt;(&lt;/span&gt;&lt;span class="n"&gt;residual&lt;/span&gt;&lt;span class="p"&gt;))&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;&lt;code&gt;residual = block(residual)&lt;/code&gt; only works because the shape never changes. That is why blocks stack&lt;/li&gt;
&lt;li&gt;Nothing normalises the stream itself along the way, so &lt;code&gt;ln_final&lt;/code&gt; is needed before the unembedding can read it&lt;/li&gt;
&lt;li&gt;&lt;code&gt;nn.ModuleList&lt;/code&gt; gives the children paths like &lt;code&gt;blocks.3.attn&lt;/code&gt;, which match &lt;code&gt;transformer_lens&lt;/code&gt;'s module paths. Its hooks are points underneath those, such as &lt;code&gt;blocks.3.attn.hook_z&lt;/code&gt;, so the two trees line up and one vocabulary covers the trace and the reference activations&lt;/li&gt;
&lt;/ul&gt;
&lt;h3 id="module-output_6"&gt;Module output&lt;/h3&gt;
&lt;div class="language-bash highlight"&gt;&lt;span class="code-lang"&gt;shell&lt;/span&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;&lt;span class="o"&gt;(&lt;/span&gt;.venv&lt;span class="o"&gt;)&lt;/span&gt;&lt;span class="w"&gt; &lt;/span&gt;dom@dom-z13:~/Desktop/PvML$&lt;span class="w"&gt; &lt;/span&gt;python&lt;span class="w"&gt; &lt;/span&gt;-m&lt;span class="w"&gt; &lt;/span&gt;pvml.modules.transformer
parameters&lt;span class="w"&gt;      &lt;/span&gt;:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;163&lt;/span&gt;,087,441
missing&lt;span class="w"&gt; &lt;/span&gt;keys&lt;span class="w"&gt;    &lt;/span&gt;:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;[]&lt;/span&gt;
unexpected&lt;span class="w"&gt; &lt;/span&gt;keys&lt;span class="w"&gt; &lt;/span&gt;:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;12&lt;/span&gt;&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;([&lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;mask&amp;#39;&lt;/span&gt;&lt;span class="o"&gt;])&lt;/span&gt;

logits&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="o"&gt;(&lt;/span&gt;&lt;span class="m"&gt;1&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;7&lt;/span&gt;,&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;50257&lt;/span&gt;&lt;span class="o"&gt;)&lt;/span&gt;
max&lt;span class="w"&gt; &lt;/span&gt;abs&lt;span class="w"&gt; &lt;/span&gt;diff&lt;span class="w"&gt; &lt;/span&gt;vs&lt;span class="w"&gt; &lt;/span&gt;gpt-2:&lt;span class="w"&gt; &lt;/span&gt;&lt;span class="m"&gt;9&lt;/span&gt;.155e-05
same&lt;span class="w"&gt; &lt;/span&gt;argmax&lt;span class="w"&gt; &lt;/span&gt;everywhere:&lt;span class="w"&gt; &lt;/span&gt;True

next-token&lt;span class="w"&gt; &lt;/span&gt;predictions
&lt;span class="w"&gt;     &lt;/span&gt;&amp;lt;BOS&amp;gt;&lt;span class="w"&gt;  &lt;/span&gt;-&amp;gt;&lt;span class="w"&gt;  &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;\n&amp;#39;&lt;/span&gt;
&lt;span class="w"&gt;       &lt;/span&gt;The&lt;span class="w"&gt;  &lt;/span&gt;-&amp;gt;&lt;span class="w"&gt;  &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39; first&amp;#39;&lt;/span&gt;
&lt;span class="w"&gt;       &lt;/span&gt;cat&lt;span class="w"&gt;  &lt;/span&gt;-&amp;gt;&lt;span class="w"&gt;  &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39; was&amp;#39;&lt;/span&gt;
&lt;span class="w"&gt;       &lt;/span&gt;sat&lt;span class="w"&gt;  &lt;/span&gt;-&amp;gt;&lt;span class="w"&gt;  &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39; on&amp;#39;&lt;/span&gt;
&lt;span class="w"&gt;        &lt;/span&gt;on&lt;span class="w"&gt;  &lt;/span&gt;-&amp;gt;&lt;span class="w"&gt;  &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39; the&amp;#39;&lt;/span&gt;
&lt;span class="w"&gt;       &lt;/span&gt;the&lt;span class="w"&gt;  &lt;/span&gt;-&amp;gt;&lt;span class="w"&gt;  &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39; floor&amp;#39;&lt;/span&gt;
&lt;span class="w"&gt;       &lt;/span&gt;mat&lt;span class="w"&gt;  &lt;/span&gt;-&amp;gt;&lt;span class="w"&gt;  &lt;/span&gt;&lt;span class="s1"&gt;&amp;#39;,&amp;#39;&lt;/span&gt;
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;p&gt;The check loads GPT-2's whole &lt;code&gt;state_dict&lt;/code&gt; into our model with &lt;code&gt;strict=False&lt;/code&gt;, then runs both on the same tokens.&lt;/p&gt;
&lt;ul&gt;
&lt;li&gt;&lt;code&gt;missing keys: []&lt;/code&gt; is the line that matters. Every parameter our model needs was found in GPT-2's. Because &lt;code&gt;strict=False&lt;/code&gt; tolerates a mismatch instead of raising, a single misspelled parameter would leave that weight at its random initialisation, and the model would still run and still produce plausible logits. An empty list is what rules that out&lt;/li&gt;
&lt;li&gt;&lt;code&gt;unexpected keys: 12&lt;/code&gt; is the other direction, entries GPT-2 has that we do not. All 12 are named &lt;code&gt;mask&lt;/code&gt;, one per block. &lt;code&gt;transformer_lens&lt;/code&gt; precomputes its causal mask as a buffer, ours is built inside &lt;code&gt;apply_causal_mask&lt;/code&gt; on every call, so there is nothing to load them into&lt;/li&gt;
&lt;li&gt;Every earlier check was handed its input from the reference's own activations. This one gets only token ids, so the entire forward pass is ours and no borrowed intermediate can prop up a mistake&lt;/li&gt;
&lt;li&gt;&lt;code&gt;9.155e-05&lt;/code&gt; across 50,257 logits at each of 7 positions, and the argmax is identical at every one. That is float32 accumulation over twelve blocks, not an error&lt;/li&gt;
&lt;/ul&gt;
&lt;p&gt;Where those parameters live:&lt;/p&gt;
&lt;div class="language-text highlight"&gt;&lt;pre&gt;&lt;span&gt;&lt;/span&gt;&lt;code&gt;  embed        38,597,376
  pos_embed       786,432
  blocks       85,054,464
  ln_final          1,536
  unembed      38,647,633
  total       163,087,441
&lt;/code&gt;&lt;/pre&gt;&lt;/div&gt;
&lt;ul&gt;
&lt;li&gt;85M of it is the twelve blocks, and two thirds of each block is its MLP&lt;/li&gt;
&lt;li&gt;GPT-2 small is usually quoted as 124M, which is this total minus &lt;code&gt;unembed&lt;/code&gt;. The original ties the two vocabulary tables, &lt;code&gt;W_U = W_E&lt;/code&gt; transposed, and counts those 38,597,376 parameters once. &lt;code&gt;transformer_lens&lt;/code&gt; keeps them as separate tensors, so they are counted twice here&lt;/li&gt;
&lt;/ul&gt;</content><category term="transformers"/><category term="gpt-2"/><category term="arena"/></entry></feed>