Transformer Systems Lab

From text to one gradient update

One short string, two toy tokenizations, one causal attention row, a cross-entropy loss and a single gradient step, small enough to check every number.

Step 1 of 8Predict, then checkAnswers saved in this browserFinite-difference check

Two scales of the same calculation

Follow one target, then account for a whole batch.

PredictionWhich source token receives the largest gradient signal?
ExampleToy text, toy tokens, causal attention, softmax loss.
What you checkAttention table, logit update, loss before and after, finite difference.
New caseChange the target, turn off the residual path, predict again.

Step 1 of 8

From text to one gradient update

Take a toy sentence, choose a tokenization, let the last context token attend to itself and the tokens before it in one causal attention row, then work through the exact softmax cross-entropy update.

Your answerPrediction neededYour first answer is kept even if you change it later.
Checking saved lesson state

Example

One next-token training example

the model learns patterns

The tokenizer here is deliberately tiny. It is not BPE, a unigram language model or real SentencePiece; it shows how token boundaries change the positions the loss is computed at. For the last training position, the context is every token before the last one. The target y is one of three toy output classes; by default it is patterns, the word that ends the sentence. All vectors and weights are hand-set toy numbers.

0the1model2learns3patterns
4 toy tokens3 next-token positionsh=[0.83, 0.76, 0.31]query from position 2

Your prediction

With the key of model scaled by 2x, which source token receives the largest gradient from the loss on the last token, through its value and the residual path?

First answerNo prediction yetChanging your answer later does not replace this one.

Controls

Change one setting, then predict again.

toy tokenizer
target y
learning rate η
scale of the model key
attention temperature τ
causal mask
residual path

Coarse: fewer pieces, so fewer next-token positions to train on.

Readout investigation · frozen hidden rows

Which tokens teach this readout?

The calculation above changed W and b for one chosen target. A training batch scores every position at once, and two choices decide what the step does: which positions count as targets, and what their summed loss is divided by. Here both stay visible through one exact update.

The example. Two right-padded sequences (B = 2, T = 4) over the classes red, blue and green; PAD is padding, not a class. Each position is one row of the batch, with a hand-picked hidden vector h ∈ ℝ² (d = 2), frozen and separate from the attention controls above. A row’s target is the next token; the mask m says whether that target counts, and M is how many count. The readout W (3 × 2) and b (3) start at zero and take one SGD step with η = 0.3. The intended objective is the mean loss over completion targets.

Targets that count

The intended objective. Steps 1 to 3 trace it.

Pair each position with the token after it

A next-token model is scored at position t on token t + 1 of the same sequence. Shift first, then decide which targets count. The mark belongs to the target, not to the position that reads it.

Sequence 1 2 of the M = 3 targets
  1. t = 0readsgreenpromptscored ongreenpromptprompt target, not counted
  2. t = 1readsgreenpromptscored onredcompletioncounts: target 1 of 3; first completion token
  3. t = 2readsredcompletionscored onredcompletioncounts: target 2 of 3
  4. t = 3readsredcompletionscored onnoneend of sequenceno next token
Sequence 2 1 of the M = 3 targets
  1. t = 0readsgreenpromptscored onbluecompletioncounts: target 3 of 3; first completion token
  2. t = 1readsbluecompletionscored onPADpaddingPAD is not a target
  3. t = 2readsPADpaddingscored onPADpaddingPAD is not a target
  4. t = 3readsPADpaddingscored onnoneend of sequenceno next token

Among the tokens scored on, the completion begins one position earlier than among the tokens read, because every token has moved one step left. Dashed lines mark both boundaries. So position 1 of sequence 1, which reads a prompt token, is scored on the first completion token and counts. The last real token of each sequence is scored from the position before it; its own position has no valid next-token target (PAD follows it in sequence 2; nothing follows it in sequence 1).

Written on the input positions before the shift, with −100 meaning “ignore”: sequence 1 [−100, −100, red, red] and sequence 2 [−100, blue, −100, −100]. Pairing each position with the label after it gives the lower rows above.

Counting decides only which positions add a loss term. A position that does not count still supplies context: in a transformer, prompt tokens feed later positions through attention, which a separate attention mask controls, and gradients from completion losses flow back into them. Here the hidden rows are chosen by hand and frozen, so those upstream paths are outside the calculation.

In the source

Transformers v4.46.3, loss_utils.py L38–39: ForCausalLMLoss drops the last logit position and the first label, pairing position t with label t + 1. TRL v0.24.0, sft_trainer.py L192–220: the collator pads labels with −100 (L199–201) separately from the attention mask, which it pads with 0 (L207–210); with completion-only loss it sets excluded labels to −100 (L215), and likewise for assistant masks (L220).

Separate historical excerpts, not one installed stack or today’s defaults; they do not show identical behaviour for every collator, template, version or packed sequence.

Follow one update: raw loss, count, gradient, step, new probabilities

h ↦ z = Wh + b↦ p = softmax(z)↦ ℓ = logsumexp(z) − zyL = Σ mℓ / Mδ = m(p − ey) / M∇W = Σ δhᵀ, ∇b = Σ δ

aEach counted row: its raw loss and its weight 1/M. Choose one to trace.
  • h
    (1, 0)
    raw loss ℓ
    1.0986
    weight
    1/3
    share of L
    0.3662
    δ
    (−2/9, 1/9, 1/9)
  • h
    (1, 0)
    raw loss ℓ
    1.0986
    weight
    1/3
    share of L
    0.3662
    δ
    (−2/9, 1/9, 1/9)
  • h
    (1, 0)
    raw loss ℓ
    1.0986
    weight
    1/3
    share of L
    0.3662
    δ
    (1/9, −2/9, 1/9)

Row sequence 1, t = 1 is scored on red. With W = 0 and b = 0 every logit is 0, so p = (1/3, 1/3, 1/3). Then p − e(red) = (−2/3, 1/3, 1/3), and times its weight 1/3 this is δ = (−2/9, 1/9, 1/9). Its h = (1, 0), so δhᵀ writes δ into column h₁ of ∇W and adds nothing to the other column.

bSummed loss Σℓ = 3.2958 over M = 3 targets, so L = 3.2958 / 3 = 1.0986. The count is taken after the shift; a row that does not count adds nothing to the sum or the count.

cGradient: the rows’ δ added up

∇b, one bar per class

red−1/3
blue0
green1/3

Each piece is one row’s δ for that class, coloured by the class the row is scored on; the solid tick is the sum ∇b. The step moves b by −η∇b, so a negative sum raises that logit.

∇W = Σ δhᵀ
h₁h₂
red−1/30
blue00
green1/30
∇b = Σ δ
value
red−1/3
blue0
green1/3

The outlined column of ∇W is where the selected row writes.

Chain rule, one entry: ∂ℓ/∂zc = pc − [c = y] and ∂zc/∂Wcj = hj, so ∂L/∂Wcj = Σ δchj over the rows that count. For one row that is the outer product δhᵀ: a 3 × 1 column times a 1 × 2 row, the shape of W.

dOne step with η = 0.3: W′ = W − η∇W, b′ = b − η∇b

W′
h₁h₂
red0.10
blue00
green−0.10
b′
value
red0.1
blue0
green−0.1

Recomputed from W = 0 and b = 0 whenever the choice of targets changes; the step is never applied twice.

eNew probabilities for the same rows

Readout probabilities before and after the step
for h =redbluegreen
(1, 0)0.3333 → 0.40180.3333 → 0.32890.3333 → 0.2693
  • sequence 1, t = 1 → red: raw loss 1.0986 → 0.9119
  • sequence 1, t = 2 → red: raw loss 1.0986 → 0.9119
  • sequence 2, t = 0 → blue: raw loss 1.0986 → 1.1119

Summed loss 3.2958 → 2.9357; mean over M = 3: 1.0986 → 0.9786.

Row sequence 1, t = 0 did not count, yet its p(green) moved from 0.3333 to 0.3006 through the shared bias b. Not counting a row removes its loss term; it does not freeze its prediction.

A held-out case, fixed before the step

A red prompt followed by a blue completion. The frozen representation gives it h = (1, 0), the same hidden row and blue target as the blue training row. Their prompts differ, but the readout receives identical inputs and their losses agree. The red completion targets also share this hidden row.

training mean, M = 3
1.0986 → 0.9786 (lower)
held-out loss, target blue
1.0986 → 1.1119 (higher)

The step raised red most for every context represented by (1, 0), and this context is followed by blue. This is a feature collision built into the example: the fixed representation cannot tell these contexts apart. It shows that a lower training loss does not by itself establish better held-out behaviour. It is not a sample and says nothing about how often this happens.

In the source

Transformers v4.46.3, loss_utils.py L24–47: when a count is supplied, fixed_cross_entropy sums the per-token losses (L25), skipping labels equal to the ignore index (L26), and divides the sum by that count (L28); ForCausalLMLoss passes the shifted tensors and the caller’s count to it (L46). It divides by whatever count it is handed. Llama 3 (v3), §4.1.3: supervised fine-tuning applies cross-entropy to target tokens and masks the loss on prompt tokens.

This page keeps only target selection and the readout derivative: not Llama 3’s network, data, optimizer or results.

Regroup the same rows, holding everything else fixed

Gradient accumulation splits a batch into microbatches and takes one step after all of them. Here the rows, hidden vectors, targets, W = 0, b = 0 and η stay the same; only the grouping, and how the pieces are combined, change.

Grouping

Two microbatches: average the two microbatch means.

Each microbatch reports a summed loss and a count
microbatchΣℓcountown mean
sequence 12.197221.0986
sequence 21.098611.0986

Average the microbatch means: (1.0986 + 1.0986) / 2 = 1.0986.

Token mean, one batch: weight per target
This grouping: weight per target
  • sequence 1, t = 1 → red: weight 1/4 (token mean 1/3)
  • sequence 1, t = 2 → red: weight 1/4 (token mean 1/3)
  • sequence 2, t = 0 → blue: weight 1/2 (token mean 1/3)

∇b under this grouping

red−1/6
blue−1/6
green1/3

Each piece is one row’s δ for that class, coloured by the class the row is scored on; the solid tick is the sum ∇b, the hollow tick the one-batch token-mean sum. The step moves b by −η∇b, so a negative sum raises that logit.

Token mean, one batch

loss before the step
1.0986
∇W, row red
(−1/3, 0)
∇W, row blue
(0, 0)
∇W, row green
(1/3, 0)
∇b
(−1/3, 0, 1/3)
token-mean loss after one step
0.9786

Split, average means

loss before the step
1.0986
∇W, row red
(−1/6, 0) (differs)
∇W, row blue
(−1/6, 0) (differs)
∇W, row green
(1/3, 0)
∇b
(−1/6, −1/6, 1/3) (differs)
token-mean loss after one step
1.0083 (differs)

The loss before the step is 1.0986 under both, so a check of the loss alone passes. The gradient does not: averaging means gives each of the 2 microbatches 1/2 of the objective whatever its count, so it gives 1/4 to each target of sequence 1 and 1/2 to each target of sequence 2. The bias gradient changes: red’s from −1/3 to −1/6, blue’s from 0 to −1/6. After the step the token mean is 1.0083 instead of 0.9786. This is the equal-sequence objective, a legitimate objective of its own; it is wrong only when presented as the token mean.

The equality holds because the activations, parameters and the single final step are the same for every grouping. With dropout, activations that depend on the batch, a step between microbatches, or per-microbatch clipping, it need not.

In the source

Hugging Face, 16 October 2024, “Where does it stem from?”: a gradient-accumulation discrepancy arose when per-batch average losses were averaged instead of dividing the summed token losses by the total token count across accumulation steps.

This reproduces that arithmetic mechanism, not their training runs, any particular model or current Trainer behaviour.

Compare the two choices of target set

The mask says what the model should learn from. Changing it changes the targets, the count and the update together; each record below is computed from W = 0 and b = 0.

Completion tokens only

targets M
3
rows beyond the completion targets
none
∇W, column h₂
(0, 0, 0)
∇b
(−1/3, 0, 1/3)
its own training mean
1.0986 → 0.9786
held-out loss
1.0986 → 1.1119

All valid next-token targets

targets M
4 (differs)
rows beyond the completion targets
sequence 1, t = 0 → green (prompt), h = (0, 1) (differs)
∇W, column h₂
(1/12, 1/12, −1/6) (differs)
∇b
(−1/6, 1/12, 1/12) (differs)
its own training mean
1.0986 → 1.0396 (differs)
held-out loss
1.0986 → 1.1280 (differs)

The two training means average over different targets, 3 against 4, so they do not show which choice trains better, and neither number chooses the target set. That choice comes from what the model should learn, checked on held-out cases chosen in advance. On the one held-out case here, both steps make the loss higher.

Which count goes in the denominator?
countCompletion tokens onlyAll valid next-token targets
targets after the shift, M34
labels before the shift that are not −10036 (not M)
input tokens that are not PAD6 (not M)6 (not M)

With all valid targets, each sequence’s first token carries a label, but no position is scored on it: the shift drops it. Dividing the same sum by 6 instead of 4 multiplies the loss and every gradient entry by 4/6, a smaller step that nothing in the loss function flags. It trusts the count its caller supplies (Transformers v4.46.3, L25, L28, L46), so the caller must count eligible targets after the shift.

Watch the weighting change, or reproduce the calculation

Optional 30-second reference explanation. The same three completion targets are regrouped; changing their weights changes the gradient. This is recorded playback.

Read the animation’s explanation
  1. After shifting, sequence 1 supplies two red completion targets and sequence 2 supplies one blue target.
  2. A mean over the three targets gives each a weight of 1/3. The bias gradient is (−1/3, 0, 1/3), in red, blue, green order.
  3. Averaging the two microbatch means gives the red targets weights of 1/4 each and the blue target a weight of 1/2. The bias gradient becomes (−1/6, −1/6, 1/3).
  4. Using equal target weights and a step size of 0.3 lowers the training mean from 1.0986 to 0.9786. The blue training row and the hand-picked blue counterexample both worsen from 1.0986 to 1.1119. They have the same hidden row, (1, 0), and the same blue target, so this readout cannot distinguish them. This calculation does not test sensitivity to their different prompts.

Run the Python file locally with NumPy. With marimo and uv installed, open the notebook with marimo edit --sandbox investigation.py; uv downloads its declared Python dependencies into an isolated environment on your computer. These reference downloads start with completion targets in one batch. The browser runner below uses the current controls; no model training service is involved.

Completion tokens only; split, average means; one step from zero with η = 0.3. This reproduces the grouping selected in step 3.

Reproduce this case in Python

Runs on your device. Python and NumPy load on demand; your inputs stay in this browser. Cached files may be downloaded again if the browser clears them.

Python runs here when you choose Run.

The mathematics: shapes, the composition and the chain rule
Symbols and shapes
symbolshapehere
HB × T × d2 × 4 × 2 hidden rows, chosen by hand and frozen
y, mB × Tthe target after the shift, and m ∈ {0, 1}: whether it counts
W, bV × d, V3 × 2 and 3, starting at 0
z, pV per rowlogits and probabilities
δB × T × Vzero on every row with m = 0
∇WV × dthe sum over rows of δhᵀ, the shape of W

A composition of functions carries one row to its loss: h ↦ z = Wh + b ↦ p = softmax(z) ↦ ℓ = −log py. Mathematically ℓ = logsumexp(z) − zy. It is computed in shifted coordinates s = z − max z: ℓ = log Σc exp(sc) − sy and p = exp(s) / Σc exp(sc). No large positive number is exponentiated, and a large common offset in z is never added back before zy is subtracted, which would cancel the digits the loss needs. A result that double precision cannot hold, such as an overflowing logit spread, sum or updated parameter, is rejected rather than rounded to infinity.

The objective. M = Σi,t mit counts the targets after the shift. If M > 0, L = (1/M) Σi,t mit ℓit. If M = 0 there is no mean: the helper reports that there are no supervised targets, returns no loss, and leaves W and b unchanged. It does not call the loss zero.

The chain rule, one entry at a time. For one row, ∂ℓ/∂zc = pc − [c = y], the softmax cross-entropy derivative. zc depends on W only through row c: ∂zc/∂Wcj = hj and ∂zc/∂bc = 1. Multiplying along the chain and summing over the rows that count,

∂L/∂Wcj = Σi,t (mit/M)(pitc − [c = yit]) hitj = Σi,t δitc hitj.

For each row that is the outer product δhᵀ, a V × 1 column times a 1 × d row, which has the shape of W; and ∇b = Σ δ.

Microbatches. With microbatches k = 1, …, K holding summed losses Sk over Mk targets, the token mean is (Σk Sk) / (Σk Mk). The mean of microbatch means is (1/K) Σk Sk/Mk, which gives each target the weight 1/(KMk). They agree when every Mk is equal. An empty microbatch adds nothing to either total of the token mean and has no mean of its own.

What is left out. In a trainable network the same δ continues as Wᵀδ into h, and from there into attention and embeddings, including through prompt positions that do not count. Here H is frozen, so that path is not computed.

Related lessons: Cross-Entropy for the p − ey derivative, Backpropagation for a received signal times a local derivative and the weight-row outer product, and Gradient Descent for the step.

The code: the same computation in NumPy

Array code over each sequence’s rows rather than one row at a time. The browser runner supplies your current mask and grouping as the dictionary case. Run without a case in Python 3 and NumPy, it prints the completion-only, one-batch reference and checks with assert that two microbatches divided once give the one-batch gradient. The final JSON line keeps full precision for comparison.

import numpy as np

# Output classes 0 red, 1 blue, 2 green (V = 3). PAD marks padding; it is not a class.
RED, BLUE, GREEN, PAD = 0, 1, 2, "PAD"
V, d, eta = 3, 2, 0.3

# Two right-padded sequences, B = 2 and T = 4. Roles mark the completion; they are not tokens.
tokens = [[GREEN, GREEN, RED, RED],
          [GREEN, BLUE, PAD, PAD]]
roles = [["prompt", "prompt", "completion", "completion"],
         ["prompt", "completion", None, None]]

# Authored, frozen hidden rows, one per input position: shape (B, T, d).
H = np.array([[[0.0, 1.0], [1.0, 0.0], [1.0, 0.0], [0.0, 0.0]],
              [[1.0, 0.0], [0.0, 0.0], [0.0, 0.0], [0.0, 0.0]]])
W0, b0 = np.zeros((V, d)), np.zeros(V)


def align(policy):
    """Shift first: row (i, t) is scored on token t + 1 of the same sequence.
    Returns targets y (a class, or None) and eligibility m (0 or 1), each B x T."""
    y, m = [], []
    for seq, role in zip(tokens, roles):
        nxt, nxt_role = seq[1:] + [None], role[1:] + [None]  # the last row has no successor
        y.append([None if tok in (None, PAD) else tok for tok in nxt])
        m.append([int(tok not in (None, PAD) and (policy == "all" or r == "completion"))
                  for tok, r in zip(nxt, nxt_role)])
    return y, m


def finite(*values):
    """Reject any result float64 cannot hold, instead of carrying inf or nan onward."""
    for value in values:
        if not np.isfinite(value).all():
            raise FloatingPointError("a result is not representable in float64")
    return values


def sums(seqs, y, m, W, b):
    """Summed loss, eligible count and summed gradients over the rows of the given sequences.
    Only rows with m = 1 are read; every other row has delta = 0."""
    S, M, GW, Gb = 0.0, 0, np.zeros(W.shape), np.zeros(b.shape)
    for i in seqs:
        rows = [t for t in range(len(y[i])) if m[i][t] == 1]
        if not rows:
            continue
        Hr = np.array([H[i, t].tolist() for t in rows])   # (R, d) eligible hidden rows
        Z = Hr @ W.T + b                                   # (R, d) @ (d, V) + (V,) -> (R, V) logits
        Zs = Z - Z.max(axis=1, keepdims=True)              # shifted: each row's largest is 0
        finite(Z, Zs)
        E = np.exp(Zs)
        total = E.sum(axis=1, keepdims=True)               # (R, 1), between 1 and V
        P = E / total                                      # (R, V) probabilities
        Y = np.array([[float(y[i][t] == c) for c in range(len(b))] for t in rows])  # one-hot targets
        # Loss in shifted coordinates: log sum exp(z - max) - (z_y - max). The offset never returns.
        S += float((np.log(total) - (Zs * Y).sum(axis=1, keepdims=True)).sum())
        M += len(rows)
        R = P - Y                                          # (R, V): p - e_y
        GW = GW + R.T @ Hr                                 # (V, R) @ (R, d) = sum of (p - e_y) h^T
        Gb = Gb + R.sum(axis=0)
        finite(S, GW, Gb)                                  # summed, not yet divided
    return S, M, GW, Gb


def token_mean(parts, y, m, W, b):
    """The declared objective: add every microbatch's sums and counts, then divide once."""
    S, M, GW, Gb = 0.0, 0, np.zeros(W.shape), np.zeros(b.shape)
    for part in parts:
        s, k, gw, gb = sums(part, y, m, W, b)
        S, M, GW, Gb = S + s, M + k, GW + gw, Gb + gb
    finite(S, GW, Gb)
    if M == 0:
        return None                                        # no supervised target: no mean, no step
    return S / M, GW / M, Gb / M


def mean_of_means(parts, y, m, W, b):
    """A different objective: each microbatch's own mean, averaged with equal weight."""
    pieces = [sums(part, y, m, W, b) for part in parts]
    if any(k == 0 for _, k, _, _ in pieces):
        return None                                        # an empty microbatch has no mean
    finite(sum(s for s, _, _, _ in pieces))                # the summed raw loss, as the ledger reports it
    K = len(pieces)                                        # each target weighs 1 / (K k)
    return finite(sum(s / k / K for s, k, _, _ in pieces),
                  sum(gw / k / K for _, k, gw, _ in pieces),
                  sum(gb / k / K for _, k, _, gb in pieces))


def sgd_step(W, b, gW, gb, eta):
    """One plain step. A parameter float64 cannot hold is rejected, never clamped."""
    return finite(W - eta * gW, b - eta * gb)


def readout(W, b, h):
    """Shifted logits, probabilities and log sum exp(z - max) for one hidden vector."""
    z = W @ h + b                                          # (V,)
    zs = z - z.max()
    finite(z, zs)
    e = np.exp(zs)
    return zs, e / e.sum(), float(np.log(e.sum()))


def loss_at(W, b, h, target):
    zs, _, log_total = readout(W, b, h)
    return log_total - float(zs[target])


def show(a):
    return (np.round(a, 6) + 0.0).tolist()                 # + 0.0 prints -0.0 as 0.0


REFERENCE_CASE = {"example": "masked-readout-one-step-v1", "policy": "completion",
                  "grouping": "one-batch", "eta": 0.3}


def reproduce_case(case):
    """This version fixes tokens, H, zero W/b and one step. Export only its declared controls."""
    if (not isinstance(case, dict) or set(case) != set(REFERENCE_CASE)
            or case["example"] != REFERENCE_CASE["example"]
            or case["policy"] not in ("completion", "all")
            or case["grouping"] not in ("one-batch", "split-token-mean", "split-mean-of-means")
            or type(case["eta"]) not in (int, float) or case["eta"] != eta):
        raise ValueError("This masked-readout case is not supported by this code")
    y, m = align(case["policy"])
    parts = [[0, 1]] if case["grouping"] == "one-batch" else [[0], [1]]
    reduction = mean_of_means if case["grouping"] == "split-mean-of-means" else token_mean
    L, gW, gb = reduction(parts, y, m, W0, b0)
    W1, b1 = sgd_step(W0, b0, gW, gb, case["eta"])
    before, _, _ = token_mean([[0, 1]], y, m, W0, b0)
    after, _, _ = token_mean([[0, 1]], y, m, W1, b1)
    h = np.array([1.0, 0.0])
    _, p1, _ = readout(W1, b1, h)
    return {**case, "y": y, "m": m, "targets": sum(map(sum, m)),
            "objective": float(L), "numerator": float(sums([0, 1], y, m, W0, b0)[0]),
            "training_before": float(before), "training_after": float(after),
            "grad_W": gW.tolist(), "grad_b": gb.tolist(), "W_after": W1.tolist(), "b_after": b1.tolist(),
            "row_loss_after": [float(loss_at(W1, b1, H[i, t], y[i][t]))
                               for part in parts for i in part for t in range(len(y[i])) if m[i][t]],
            "held_out_before": float(loss_at(W0, b0, h, BLUE)),
            "held_out_after": float(loss_at(W1, b1, h, BLUE)), "probabilities_after": p1.tolist()}


# Standalone reference demonstration; a supplied case is reproduced below instead.
if "case" not in globals():
    y, m = align("completion")
    L, gW, gb = token_mean([[0, 1]], y, m, W0, b0)
    W1, b1 = sgd_step(W0, b0, gW, gb, eta)                     # one SGD step
    L1, _, _ = token_mean([[0, 1]], y, m, W1, b1)
    h = np.array([1.0, 0.0])
    _, p1, _ = readout(W1, b1, h)

    print("targets M =", sum(map(sum, m)), "| loss", round(L, 6), "->", round(L1, 6))
    print("grad_W", show(gW), "grad_b", show(gb))
    print("W'", show(W1), "b'", show(b1))
    print("p' at h = (1, 0):", show(p1))

    # The same rows in two microbatches, then one step: add sums, divide once.
    _, gW_split, gb_split = token_mean([[0], [1]], y, m, W0, b0)
    assert np.allclose(gW_split, gW) and np.allclose(gb_split, gb)
    # Averaging the two microbatch means starts at the same loss but moves differently.
    L_mm, gW_mm, gb_mm = mean_of_means([[0], [1]], y, m, W0, b0)
    print("mean of means: loss", round(L_mm, 6), "grad_b", show(gb_mm))

    # A held-out case declared in advance: red prompt -> blue completion, also h = (1, 0).
    print("held-out loss", round(loss_at(W0, b0, h, BLUE), 6), "->", round(loss_at(W1, b1, h, BLUE), 6))

    # The other declared target set, and a batch with nothing to learn from.
    y_all, m_all = align("all")
    L_all, gW_all, gb_all = token_mean([[0, 1]], y_all, m_all, W0, b0)
    print("all targets: M =", sum(map(sum, m_all)), "grad_b", show(gb_all))
    assert token_mean([[0, 1]], [[None] * 4] * 2, [[0] * 4] * 2, W0, b0) is None

# The browser sets case from JSON. Locally, use runpy.run_path(..., init_globals={"case": data}).
# Full-precision JSON is the comparison contract; the preceding reference printout is rounded.
import json
result = reproduce_case(globals().get("case", REFERENCE_CASE))
if "case" in globals():
    print("Case:", result["policy"], "|", result["grouping"], "| one step, eta =", result["eta"])
    print("targets M =", result["targets"], "| training mean", result["training_before"], "->", result["training_after"])
    print("grad_W", result["grad_W"], "grad_b", result["grad_b"])
    print("W'", result["W_after"], "b'", result["b_after"])
    print("p' at h = (1, 0):", result["probabilities_after"])
print("MASKED_READOUT_CASE=" + json.dumps(result, allow_nan=False, separators=(",", ":")))
Practice (optional)

Two short checks on the same example. A hint comes before the answer. What you open is noted beside each item on this page only; nothing is saved, and reloading clears it.

  1. Add a third sequence: green (prompt), red (prompt), blue (completion), PAD. With completion tokens as the targets, what is M for all three sequences together?

  2. Now put each of the three sequences in its own microbatch and average the three microbatch means. What weight does the third sequence’s single target get? The token mean gives every target 1/4.

What this does not show. No transformer is trained here: H is chosen by hand and frozen, so attention, embeddings and upstream gradients are outside the calculation. The source notes map operations in separate historical code and one paper; they do not describe one installed software stack, current defaults, or Llama 3’s recipe or results. One hand-picked held-out case is a counterexample, not a measurement of generalization.

Continuation · source records into the same readout

What weighting did this data decision create?

The readout above used a token mean: the summed loss over completion targets, divided by their number. Now its two sequences arrive as source records, and a third record is supplied as an exact mirror of the first. Keeping, dropping or down-weighting a record decides how much each of its targets counts, and so which objective the step follows.

The example. Three training records: A1 and B1 are the two sequences above, and A2 is stated to be an exact mirror of A1, with the same text, roles and hidden rows under its own record ID. Each record r has a chosen weight ar ≥ 0 that multiplies each of its nr completion targets. The denominator is Z = Σ arnr, and each target of record r weighs ar / Z. W and b start at zero and take one step with η = 0.3, as above. The mirror relation is supplied with the example: nothing here detects mirrors or establishes where a record came from.

Weighting (a for A1, A2, B1)

Every supplied record trains at weight 1, so source A’s two targets are counted twice.

Which records train, and how much each counts

The inventory is the same under every weighting. A weight changes how much each of a record’s targets counts; a weight of 0 leaves the record out of this objective without deleting it.

  • A1 source A

    source record

    Written for this page: the first sequence of the readout above.

    completion targets, n
    t = 1 → redt = 2 → red (2)
    weight a
    1
    term in Z, a · n
    1 × 2 = 2
    each target weighs a / Z
    1/5
    share of the objective
    2/5
  • A2 source A

    exact mirror of A1, as supplied

    Supplied with the example: an exact mirror of A1, with the same text, roles and hidden rows.

    completion targets, n
    t = 1 → redt = 2 → red (2)
    weight a
    1
    term in Z, a · n
    1 × 2 = 2
    each target weighs a / Z
    1/5
    share of the objective
    2/5
  • B1 source B

    source record

    Written for this page: the second sequence of the readout above. No mirror is stated.

    completion targets, n
    t = 0 → blue (1)
    weight a
    1
    term in Z, a · n
    1 × 1 = 1
    each target weighs a / Z
    1/5
    share of the objective
    1/5

Z = 1 × 2 + 1 × 2 + 1 × 1 = 5. The step reads 5 eligible targets.

A record’s share is a · n / Z, not a. Here source A carries 4/5 and source B carries 1/5 of the objective. Giving each source half would be another choice, an equal-source mean like the mean of means above; keeping one copy is not that.

What makes A2 an exact mirror of A1 here?
  • tokensA1: green, green, red, redA2: green, green, red, redequal
  • rolesA1: prompt, prompt, completion, completionA2: prompt, prompt, completion, completionequal
  • eligible targetsA1: t = 1 → red, t = 2 → redA2: t = 1 → red, t = 2 → redequal
  • hidden rowsA1: (0, 1), (1, 0), (1, 0), (0, 0)A2: (0, 1), (1, 0), (1, 0), (0, 0)equal

The objective sees a record only through its counted rows: their hidden vectors and targets. Because these are equal, A2’s targets are interchangeable with A1’s at every W and b. The relation itself is supplied with the example. Equal fields in two records written for this example say nothing about where either came from, and similar text or a shared source label would not make the rows equal.

In the source

Lee et al., ACL 2022, §4.1–4.2 (pp. 8426–8428): ExactSubstr matches shared substrings with a suffix array; the experiments use a minimum length of 50 BPE tokens. NearDup proposes candidate pairs with 5-gram MinHash (a 9,000-entry signature in 450 groups of 20), keeps pairs whose token edit similarity exceeds 0.8, and clusters them as connected components of the match graph.

These are operational matching rules: their output is a match, not a record of a document’s origin. A2’s relation here is stated with the example and reproduces neither algorithm.

Follow the weights through the targets to the gradient

Each eligible target of record r weighs ar / Z. A record’s targets add up to its share, and the targets scored on a class add up to that class’s share of the objective L. In the drawing each path is one target, and its width is that weight.

  1. A1, weight 1: t = 1 → red weighs 1/5; t = 2 → red weighs 1/5. Together 2/5 of L.
  2. A2, weight 1: t = 1 → red weighs 1/5; t = 2 → red weighs 1/5. Together 2/5 of L.
  3. B1, weight 1: t = 0 → blue weighs 1/5. Together 1/5 of L.
  4. By class: red 4/5, blue 1/5; green is never a target.
Path width is proportional to target weight; a dashed hairline marks weight 0. Weighting (1, 1, 1) for (A1, A2, B1).

From class shares to the gradient

At W = 0 and b = 0 every row predicts p = (1/3, 1/3, 1/3), and every counted row has h = (1, 0). Each target’s δ is its weight times p − ey, so here the gradient depends on the records only through the class shares:

∇b = 4/5 · (−2/3, 1/3, 1/3) + 1/5 · (1/3, −2/3, 1/3) = (−7/15, 2/15, 1/3)

∇W holds the same vector in column h₁ and zeros in column h₂. The shortcut belongs to this example: with differing hidden rows, each target writes its own δhᵀ into ∇W.

∇W
h₁h₂
red−7/150
blue2/150
green1/30
∇b
value
red−7/15
blue2/15
green1/3
W′ = W − η∇W
h₁h₂
red0.140
blue−0.040
green−0.10
b′ = b − η∇b
value
red0.14
blue−0.04
green−0.1

One step with η = 0.3, recomputed from W = 0 and b = 0 whenever the weighting changes; the step is never applied twice.

Keep one copy: the objective above

Z
3
red : blue share
2/3 : 1/3
∇b (= ∇W column h₁)
(−1/3, 0, 1/3)
b′ (= W′ column h₁)
(0.1, 0, −0.1)
p′ at h = (1, 0)
(0.4018, 0.3289, 0.2693)

Keep both copies

Z
5 (differs)
red : blue share
4/5 : 1/5 (differs)
∇b (= ∇W column h₁)
(−7/15, 2/15, 1/3) (differs)
b′ (= W′ column h₁)
(0.14, −0.04, −0.1) (differs)
p′ at h = (1, 0)
(0.4317, 0.3012, 0.2671) (differs)

Counting source A twice gives red 4/5 of the objective instead of 2/3. The step raises red more (b′ red 0.14 instead of 0.1) and now lowers blue (−0.04 instead of 0). This is a different objective from the one above, not a larger dose of it.

Read the change on criteria fixed in advance

Keeping both copies trains a different objective from keeping one copy, so their own training losses cannot rank them. Compare the weightings on criteria fixed before the step, computed the same way after every weighting. Every criterion starts at log 3 = 1.0986.

  • Its own objective

    each weighting’s own targets and weights: one objective for one copy and the split pair, another for both copies

    Keep both copies (selected)
    0.9120
    Keep one copy
    0.9786
    Split the pair’s weight
    0.9786

    Not comparable between both copies and one copy: the two numbers average different objectives.

  • Fixed audit

    the original two red and one blue targets, equal weights: the readout’s completion-target objective above

    Keep both copies (selected)
    0.9600
    Keep one copy
    0.9786
    Split the pair’s weight
    0.9786

    Both copies: lower than one copy.

  • Blue probe

    the held-out case above: a red prompt, a blue completion, h = (1, 0)

    Keep both copies (selected)
    1.2000
    Keep one copy
    1.1119
    Split the pair’s weight
    1.1119

    Both copies: higher than one copy.

  • Evaluation copy E_A

    known to share A1’s source; its two red completion targets

    Keep both copies (selected)
    0.8400
    Keep one copy
    0.9119
    Split the pair’s weight
    0.9119

    Both copies: lower, on a criterion with known overlap; not evidence about new sources.

On the fixed audit, keeping both copies ends lower (0.9600 against 0.9786); on the blue probe it ends higher (1.2000 against 1.1119). Both hold at once. The audit’s blue target has the probe’s h = (1, 0) and target, so its loss rises by the same amount, but each of its two red targets falls from 0.9119 to 0.8400, which is more in total. The one-copy step follows the audit’s own gradient; one finite step along it need not be the largest decrease available. Neither number chooses the weighting. If the intended objective counts each source’s targets once, keeping both copies does not implement it and splitting the pair’s weight does, exactly. If the extra weight on source A is intended, that is a different, legitimate objective, to be read on criteria like these, fixed in advance.

E_A is known to share A1’s source, and its loss is lower after both copies (0.8400 against 0.9119). A lower loss on a criterion with known overlap with training does not establish performance on new sources and does not measure memorization. The step raises red at h = (1, 0) for every context represented there, whatever its stated source, so the arithmetic does not single out the overlap as the cause. Excluding A2 does not remove the overlap: A1 trains under every weighting. The blue probe is a different kind of evidence, a feature collision built into the example.

In the source

Lee et al., ACL 2022, Table 5 and §6.3 (p. 8430): for Transformer-XL on LM1B, perplexity is 21.77 on the official validation set, 10.11 on the validation examples NearDup flagged as duplicates, and 23.58 on those classified unique.

One existing model on three evaluation subsets: the averaged population changes, not the training recipe.

FineWeb technical report, “Ablations and evaluation setup”, “More deduplication is always better, right?” and “Taking a step back: individual dump dedup”: ablation models have 1.82B parameters and most train on about 28B tokens. For crawl 2013-48, two models trained on 28B tokens each: one drawn from the roughly 31B tokens kept by iterative MinHash deduplication across crawls, the other from 171B tokens obtained by deduplicating, within that crawl alone, the roughly 460B tokens that process removed. The second scored better on their aggregate benchmark; the two pools share an estimated 4B tokens. The preceding full-dataset comparison used 350B-token samples instead.

Evidence against assuming that removing more always helps, in that measured setting. The report’s explanation is a hypothesis; neither a universal deduplication rule nor a benefit from keeping exact mirrors follows. Report contents as checked on 29 September 2026, not a pinned revision.

The mathematics: the weighted objective, exact mirrors and scaling

The objective. For record r with eligibility mrt after the shift, raw losses ℓrt and chosen weight ar ≥ 0, let Sr = Σt mrtℓrt and nr = Σt mrt. Then Z = Σr arnr and, if Z > 0,

L = Σr arSr / Z, δrt = (armrt / Z)(prt − ey), ∇W = Σ δhᵀ, ∇b = Σ δ.

The chain rule is the readout’s, with each target’s factor 1/M replaced by ar/Z. With every ar = 1, Z = M and this is the token mean above. If Z = 0, no target carries weight: there is no mean and no gradient, and W and b stay as they are. It is not a loss of zero.

Exact mirrors. If record q has the same counted hidden rows and targets as record r, then Sq = Sr, nq = nr and their gradient sums agree at every W and b. Their weights then enter only through ar + aq: any split with sum 1 gives the one-copy objective, gradient and step. A different counted feature, target or mask breaks the equality; a shared source label does not create it.

Scaling. Multiplying every ar by the same c > 0 multiplies Z and the weighted sum by c, so L and its gradient are unchanged. Mirroring every record once is such a scaling; mirroring some records is not.

Arithmetic. Z is formed first, then each record’s sums are scaled by ar/Z ≤ 1/nr, so a very large or very small common weight is never multiplied into a loss sum. Weights must be finite and non-negative, one per record. A Z, summed loss, gradient or step that overflows double precision is rejected rather than rounded to infinity or treated as no supervision; a target weight far below the others can underflow to 0, as a negligible probability does.

The code: the same computation in NumPy

Array code over each record’s counted rows. It refuses a malformed record, weighted or not, before any arithmetic. It runs as written with Python 3 and NumPy, prints the values above, and checks with assert that splitting the mirror pair’s weight matches one copy at nonzero W and b, that a common factor cancels, and that no weighted target means no step.

import numpy as np

# The frozen readout above, now fed by source records. Classes 0 red, 1 blue, 2 green (V = 3).
RED, BLUE, GREEN, PAD = 0, 1, 2, "PAD"
V, d, eta = 3, 2, 0.3
W0, b0 = np.zeros((V, d)), np.zeros(V)

# A record: tokens, roles and hand-picked frozen hidden rows H (T x d, one row per position).
# Provenance is carried for the ledger and checked by the page's helper; no function below
# reads "lineage" or "relation".
A1 = {"id": "A1", "lineage": "source A", "relation": None,
      "tokens": [GREEN, GREEN, RED, RED],
      "roles": ["prompt", "prompt", "completion", "completion"],
      "H": [[0.0, 1.0], [1.0, 0.0], [1.0, 0.0], [0.0, 0.0]]}
A2 = dict(A1, id="A2", relation=("exact mirror of", "A1"))  # supplied, not detected here
B1 = {"id": "B1", "lineage": "source B", "relation": None,
      "tokens": [GREEN, BLUE, PAD, PAD],
      "roles": ["prompt", "completion", None, None],
      "H": [[1.0, 0.0], [0.0, 0.0], [0.0, 0.0], [0.0, 0.0]]}
E_A = dict(A1, id="E_A", relation=("evaluation copy of", "A1"))  # scored, never trained
RECORDS = [A1, A2, B1]


def check(records, W, b):
    """Refuse, before any arithmetic, what the page's helper refuses, for every record whether or
    not it carries weight. W is finite (k x d) with k >= 2 classes and d >= 1; b is finite with k
    entries. Each record has the same T >= 1 positions: class tokens 0..V-1 with a prompt or
    completion role, then only PAD (no role); every token after the first is some row's target,
    so it must be one of W's k classes; and T hidden rows of d finite numbers."""
    k, width = W.shape if len(W.shape) == 2 else (0, 0)
    if k < 2 or width < 1 or b.shape != (k,) or not (np.isfinite(W).all() and np.isfinite(b).all()):
        raise ValueError("W must be finite (k x d) with k >= 2, d >= 1, and b finite with k entries")

    def number(v):
        return isinstance(v, (int, float)) and not isinstance(v, bool) and np.isfinite(float(v)).all()

    T = len(records[0]["tokens"]) if records else 0
    if T < 1:
        raise ValueError("at least one record with at least one position is required")
    for r, rec in enumerate(records):
        tokens, roles, H = rec["tokens"], rec["roles"], rec["H"]
        if len(tokens) != T or len(roles) != T or len(H) != T or any(
                len(row) != width or not all(number(v) for v in row) for row in H):
            raise ValueError(f"record {r} needs {T} tokens, roles and rows of {width} finite numbers")
        padded = False
        for t, (tok, role) in enumerate(zip(tokens, roles)):
            if tok == PAD and role is None:
                padded = True
            elif (padded or isinstance(tok, bool) or not isinstance(tok, int) or not 0 <= tok < V
                  or role not in ("prompt", "completion") or (t > 0 and tok >= k)):
                raise ValueError(f"record {r}, position {t}: expected a class token of W with a "
                                 "prompt or completion role, and PAD only at the end")


def targets(rec):
    """Shift inside the record, then keep completion targets: row t is scored on token t + 1."""
    nxt, nxt_role = rec["tokens"][1:] + [None], rec["roles"][1:] + [None]
    return [(t, tok) for t, (tok, role) in enumerate(zip(nxt, nxt_role))
            if tok not in (None, PAD) and role == "completion"]


def finite(*values):
    """Reject any result float64 cannot hold, instead of carrying inf or nan onward."""
    for value in values:
        if not np.isfinite(value).all():
            raise FloatingPointError("a result is not representable in float64")
    return values


def record_sums(rec, W, b):
    """S_r and the summed gradients over one record's eligible rows; no other row is read."""
    rows = targets(rec)
    Hr = np.array([rec["H"][t] for t, _ in rows])        # (n, d) eligible hidden rows
    z = Hr @ W.T + b                                      # (n, d) @ (d, V) + (V,) -> (n, V) logits
    zs = z - z.max(axis=1, keepdims=True)                 # shifted: each row's largest is 0
    finite(z, zs)
    e = np.exp(zs)
    total = e.sum(axis=1, keepdims=True)                  # (n, 1), between 1 and V
    Y = np.array([[float(y == c) for c in range(len(b))] for _, y in rows])  # one-hot targets
    S = float((np.log(total) - (zs * Y).sum(axis=1, keepdims=True)).sum())
    R = e / total - Y                                     # (n, V): p - e_y
    GW, Gb = R.T @ Hr, R.sum(axis=0)                      # (V, n) @ (n, d): sum of (p - e_y) h^T
    finite(S, GW, Gb)                                     # summed, not yet weighted
    return S, GW, Gb


def weighted_objective(records, weights, W, b):
    """Z = sum_r a_r n_r and L = sum_r a_r S_r / Z: each eligible target of record r weighs a_r / Z.
    Returns None when Z = 0: no weighted target, so no mean and no step."""
    check(records, W, b)                                  # every record, including weight 0
    if len(weights) != len(records):
        raise ValueError("give exactly one weight per record")
    for a in weights:
        if not isinstance(a, float) or not np.isfinite(a).all() or a < 0:
            raise ValueError("each weight must be a finite float >= 0")
    Z = 0.0
    for a, rec in zip(weights, records):
        Z += a * len(targets(rec))                        # this record's term a_r n_r
    finite(Z)
    if Z == 0:
        return None
    L, raw, GW, Gb = 0.0, 0.0, np.zeros(W.shape), np.zeros(b.shape)
    for a, rec in zip(weights, records):
        if a == 0 or not targets(rec):
            continue                                      # excluded, or nothing to learn: never read
        S, gw, gb = record_sums(rec, W, b)
        w = a / Z                                         # weight of each of this record's targets
        L, raw, GW, Gb = L + w * S, raw + S, GW + w * gw, Gb + w * gb
    finite(L, raw, GW, Gb)                                # raw: the summed raw loss the ledger reports
    return Z, L, GW, Gb


def sgd_step(W, b, gW, gb, eta):
    """One plain step. A parameter float64 cannot hold is rejected, never clamped."""
    return finite(W - eta * gW, b - eta * gb)


def loss_at(W, b, h, target):
    """Raw loss of one hidden vector against one target, in shifted coordinates."""
    z = W @ np.array(h) + b
    zs = z - z.max()
    finite(z, zs)
    return float(np.log(np.exp(zs).sum())) - float(zs[target])


def show(a):
    return (np.round(a, 6) + 0.0).tolist()                # + 0.0 prints -0.0 as 0.0


WEIGHTINGS = {"both copies": [1.0, 1.0, 1.0],
              "one copy": [1.0, 0.0, 1.0],
              "split pair": [0.5, 0.5, 1.0]}
for name, a in WEIGHTINGS.items():
    Z, L, gW, gb = weighted_objective(RECORDS, a, W0, b0)
    W1, b1 = sgd_step(W0, b0, gW, gb, eta)                # one step from zero, never stacked
    own = weighted_objective(RECORDS, a, W1, b1)[1]
    # Fixed criteria, the same after every weighting, each with equal target weights:
    audit = weighted_objective([A1, B1], [1.0, 1.0], W1, b1)[1]  # original 2 red : 1 blue
    probe = loss_at(W1, b1, [1.0, 0.0], BLUE)             # red prompt -> blue completion
    overlap = weighted_objective([E_A], [1.0], W1, b1)[1]  # known same source as A1
    print(name, "| Z =", Z, "| loss", round(L, 6), "| grad_b", show(gb), "| b'", show(b1))
    print("   own", round(own, 6), "| audit", round(audit, 6), "| blue probe", round(probe, 6),
          "| E_A", round(overlap, 6))

# Exact mirrors: splitting the pair's weight is the one-copy objective at any W and b, not only at 0.
Wx, bx = np.array([[0.4, -0.3], [-0.2, 0.5], [0.1, 0.25]]), np.array([0.05, -0.1, 0.2])
_, L_split, gW_split, gb_split = weighted_objective(RECORDS, [0.5, 0.5, 1.0], Wx, bx)
_, L_one, gW_one, gb_one = weighted_objective(RECORDS, [1.0, 0.0, 1.0], Wx, bx)
assert abs(L_split - L_one) <= 1e-12 and np.allclose(gW_split, gW_one) and np.allclose(gb_split, gb_one)
# Multiplying every weight by the same positive number leaves the objective unchanged.
assert np.allclose(weighted_objective(RECORDS, [2.0, 2.0, 2.0], Wx, bx)[2],
                   weighted_objective(RECORDS, [1.0, 1.0, 1.0], Wx, bx)[2])
# No weighted target: no mean and no step, rather than a loss of zero.
assert weighted_objective(RECORDS, [0.0, 0.0, 0.0], W0, b0) is None
Practice (optional)

Five short checks on the same records. A hint comes before the answer. Beside each item this page notes your first check and any help opened; a later correct check never changes the first. Nothing is saved or sent, and reloading clears it. A shown answer is assistance, not an unassisted success.

  1. Keep both copies of A and add C1, a supplied exact mirror of B1. Every weight is 1. What is Z?

  2. With the same four records, what share of the objective do blue targets carry?

  3. After one step, the fixed-audit loss is 0.9600 with both copies and 0.9786 with one copy. What does that support?

  4. No mirror of B1 is stated. What does that establish about other copies of B1?

  5. In Lee et al.’s Table 5, the same Transformer-XL model has perplexity 21.77 on the official validation set and 10.11 on the validation examples flagged as duplicates. What changed between the two numbers?

What this does not show. The mirror relation is supplied, not detected; nothing here establishes a record’s origin, measures memorization, or recommends a deduplication policy. The readout is the frozen one above, so no network is trained, and one hand-built example with one step is a counterexample to a general claim, not a measurement. The source notes describe what two published works matched and measured; they do not validate this example.

Outline of this lesson
Examplethe model learns patternstokens → attention → softmax and cross-entropy → gradient
01Example

toy text and tokenization

02Question

which earlier token gets the most gradient

03Prediction

choose a source token

04Controls

tokenizer, key scale, mask

05Result

trace, table, finite-difference check

06What stays true

only tokens on a path to the loss get gradient

07New case

turn the residual path off

08Next

RoPE: position as rotation

How the steps workEach step asks for a prediction before it shows the result. Some steps, such as KV memory, also let you skip the prediction; skipping saves no guess.

Transformer Systems Lab: all steps

From the update to position and memory.

Next, see how position changes attention scores, then how attention stores and reaches earlier tokens. Choices in the batch section above stay on this page only. Each step saves its own answers in this browser, and each link opens a different worked example.

  1. 01Text to one update

    Toy tokens, causal attention, loss, gradient

    current
  2. AtlasTokens and position

    Token ids, embeddings, RoPE phase

    background reading
  3. AtlasAttention routing

    QK^T scores, mask, softmax row

    background reading
  4. 02RoPE phase

    Rotated Q/K pair, phase gap, dot score

    next
  5. 03KV memory

    Mem_KV = B * N_layers * T * H_kv * d_head * 2 * bytes

    later
  6. 04Long-context pressure

    Memory, position, retrieval, limits of evaluation

    later
  7. 05Serving and decoding

    Time to first token, time per output token, cache reads, sampling

    later
  8. 06Speculative decoding

    A draft model proposes tokens; the target model checks them

    later
  9. 07Evaluation and falsification

    Claim, slice, metric, counterexample

    later
  10. 08Capstone systems claim

    Assumptions, evidence, where the claim fails

    later