From text to next token · interactive lecture

LLM 核心机制

沿一条真实的数据流,把 BPE、decoder、训练、KV cache、采样与评测连成一个可以操作、可以运行、可以解释的闭环。

文本token idsdecoderlossprefill / decodenext token

读法:每章先观察一个现象,再把它写成 shape 与公式,最后用浏览器实验和独立 Python 文件复现。

来源:机制基准、Stanford CS336 2026 课程材料和固定 commit 的 nanochat 源码分层呈现;缺失内容不会补写成官方结论。

五章只回答一个问题:模型怎样得到下一个 token?

  1. 01
    表示文本如何成为整数序列
  2. 02
    计算decoder block 如何变换每个位置
  3. 03
    学习怎样用真实后继 token 构造 loss
  4. 04
    推理prompt 与逐 token 生成为何是两种负载
  5. 05
    决策怎样从 logits 选择 token 并停止

CHAPTER 01 · REPRESENTATION

文本如何变成模型输入

Tokenizer 不是前处理细节:它决定词表、序列长度、特殊边界和 loss 的计数单位。BPE 在 byte 与整词之间学习一套数据驱动的压缩规则。

先问:为什么不直接按“词”切分?

整词词表会遇到未登录词和巨大词表;逐 byte 永远可表示,却让序列过长。byte-level BPE 从 256 个 byte token 出发,反复合并训练语料里最常见的相邻 pair。频繁片段变短,罕见文本仍可退回 bytes,因此不会出现真正的未知字符。

可计算对象

一次 merge 到底做什么

\[ (a^*,b^*)=\arg\max_{(a,b)}\operatorname{count}(a,b),\qquad (a^*,b^*)\mapsto v_{\mathrm{new}}. \]

统计当前 token 序列中的相邻 pair,选择频率最高者,分配新 id,再替换所有不重叠出现。下一轮统计基于更新后的序列,因此规则有顺序。

数据流

类型与长度

str→ UTF-8bytes [T₀]→ mergesids [T]→ lookupx [T,d]

通常 \(T\le T_0\),但 token 数不是字数;同一字符可能占多个 UTF-8 bytes,同一常见片段也可能被一个 token 表示。

手算一个微型例子

lower初始 byte/char 视图
lowermerge: o + w
lowermerge: l + ow
lowermerge: e + r

这不是在发现语言学意义上的“词根”。它只是在给定语料与词表预算下,提高相邻 byte 片段的复用率。

LAB 01

BPE merge 工作台

每次点击只执行一条最频繁 pair merge,观察 token 数和规则表如何共同变化。

浏览器实验按 Unicode 字符演示规则;下载例子使用真实 UTF-8 bytes。

step0
last pair
vocab +0
compression1.00×

    尚未执行 merge。先预测哪个 pair 最常见,再点击验证。

    SOURCE MAP

    CS336 2026 Assignment 1,pp.3–12 要求实现 byte-level BPE、pre-tokenization、special token、encode/decode;Lecture 1 将 tokenizer 描述为 raw bytes 与 integer tokens 之间的工程接口。nanochat 则在 tokenizer.py 用 rustbpe 训练、tiktoken 推理,并在普通文本编码与 special token 注入之间保持明确边界。

    RUNNABLE EXAMPLE 01

    训练 8 条 merge rule,并检查 UTF-8 round-trip

    rules, vocab = train("low lower lowest low 你好", num_merges=8)
    ids = encode("low lower lowest low 你好", rules)
    text = decode(ids, vocab)
    
    assert text == "low lower lowest low 你好"
    assert len(ids) < len(text.encode("utf-8"))

    运行:python 01_bpe.py · 预期:round-trip: True

    特殊 token 不是普通字符串

    若把 <|bos|> 当普通文本编码,它应拆成多个普通 token;只有显式 special-token API 才把它映射为保留 id。混用两条路径会让训练模板和停止条件失配。

    压缩率不是语言能力

    更大词表通常缩短序列,却增加 embedding / LM head 参数,并可能让罕见片段学习不足。跨 tokenizer 比较 token-level perplexity 也会改变分母。

    这一章结束后,应能回答

    • BPE 学到的是按频率排序的压缩规则,而不是唯一正确的词边界。
    • byte-level 起点为何既能覆盖任意 UTF-8 文本,又会带来较长初始序列。
    • 为什么 encode/decode round-trip、special token 隔离和词表大小要分别测试。

    CHAPTER 02 · MODEL CORE

    一个 decoder block 如何计算

    从输入 embedding 开始追踪 residual stream:RMSNorm 控制尺度,causal attention 聚合可见历史,SwiGLU 做逐位置变换,RoPE 在 Q/K 中编码相对位置。

    [B,T]token ids
    [B,T,d]embedding
    RMSNorm → causal attention → residualRMSNorm → SwiGLU → residual× L blocks
    [B,T,V]logits

    Decoder-only 的“only”指什么

    它没有单独的 encoder memory。每层 self-attention 的 Q、K、V 都来自同一条 token residual stream,当前位置只能读取自己与左侧前缀。训练时所有位置并行计算,因果依赖由 mask 保证;生成时则只取最后位置的 logits 决定下一 token。

    \[ X_0=E[\mathrm{ids}],\quad X'_{\ell}=X_{\ell}+\operatorname{Attn}(\operatorname{RMSNorm}(X_{\ell})),\quad X_{\ell+1}=X'_{\ell}+\operatorname{SwiGLU}(\operatorname{RMSNorm}(X'_{\ell})). \]

    02.1 · CAUSAL MASK

    屏蔽发生在 softmax 之前

    \[ A=\operatorname{softmax}\left(\frac{QK^\top}{\sqrt{d_h}}+M\right),\qquad M_{ij}=\begin{cases}0,&j\le i\\-\infty,&j>i.\end{cases} \]

    把未来位置加 \(-\infty\) 后,softmax 概率严格为 0。若在 softmax 之后简单清零却不重新归一化,行和不再为 1;若完全不 mask,训练时就能偷看目标 token。

    02.2 · ROPE

    把位置变成 Q/K 的相位

    \[ \operatorname{RoPE}(x,m)_{2i:2i+2}=R(m\theta_i)x_{2i:2i+2},\qquad \langle R(m\theta)q,R(n\theta)k\rangle=\langle q,R((n-m)\theta)k\rangle. \]

    二维旋转保持向量范数;Q/K 的点积依赖位置差 \(n-m\)。V 不旋转,因为位置关系是在注意力相似度中注入,而不是改写被聚合的内容向量。

    02.3 · RMSNORM

    只校准均方根,不减均值

    \[ \operatorname{RMSNorm}(x)=g\odot\frac{x}{\sqrt{\frac{1}{d}\sum_{i=1}^{d}x_i^2+\epsilon}}. \]

    \(\epsilon\) 防止零向量除零;实现常先把平方与均值提升到 fp32,再转回计算 dtype。pre-norm 把归一化放在子层之前,让 residual path 保持直接。

    02.4 · SWIGLU

    用一条分支控制另一条分支

    \[ \operatorname{SwiGLU}(x)=W_2\left(\operatorname{SiLU}(W_1x)\odot W_3x\right). \]

    两次上投影分别产生 gate 与 value,逐元素相乘后再投回 \(d_{model}\)。参数量公平比较时常把 \(d_{ff}\) 设为约 \(8d/3\),而不是沿用普通 FFN 的 \(4d\)。

    LAB 02

    Decoder block 解剖台

    切换四个机制;同一控制区同时给出视觉变化、数值性质与实现边界。

    query=3 只能读取 key 0..3;未来 attention mass = 0。

    SOURCE CONTRAST

    CS336 2026 A1,pp.13–26 的主线是 pre-norm decoder-only Transformer、RMSNorm、RoPE、causal self-attention 与 SwiGLU;课程把这些视为现代 LM 的常见选择,而不是理论上唯一最优。nanochat 的 gpt.py 使用无可学习 scale 的 F.rms_norm 与 RoPE,attention 通过 causal=True 实现 mask;但其 MLP 使用 ReLU²,不是 SwiGLU。

    RUNNABLE EXAMPLE 02

    一个完整 pre-norm decoder block

    下载
    normalized = rms_norm(x)
    attended, _ = mha(normalized, normalized, normalized,
                      attn_mask=future)
    x = x + attended
    x = x + swiglu(rms_norm(x))
    assert x.shape == (batch, length, width)

    预期:output shape == input shape

    RUNNABLE EXAMPLE 03

    验证未来 attention 为零

    下载
    future = torch.triu(torch.ones(T, T, dtype=torch.bool), 1)
    scores = scores.masked_fill(future, float("-inf"))
    attention = scores.softmax(dim=-1)
    assert torch.all(attention[future] == 0)

    预期:lower-triangular attention,行和为 1

    RUNNABLE EXAMPLE 04

    RoPE 的相对位置性质

    下载
    dot_a = torch.dot(rope(q, 2), rope(k, 5))
    dot_b = torch.dot(rope(q, 9), rope(k, 12))
    assert torch.allclose(q.norm(), rope(q, 7).norm())
    assert torch.allclose(dot_a, dot_b, atol=1e-5)

    预期:共同平移位置后点积不变

    RUNNABLE EXAMPLE 05

    手写 RMSNorm 对齐 PyTorch

    下载
    inv_rms = torch.rsqrt(x.float().square().mean(-1, keepdim=True) + eps)
    manual = (x.float() * inv_rms * scale).to(x.dtype)
    reference = F.rms_norm(x, (x.size(-1),), scale, eps)
    assert torch.allclose(manual, reference)

    预期:输出 RMS 约为 1

    RUNNABLE EXAMPLE 06

    SwiGLU 门控不是 ReLU²

    下载
    y = (F.silu(x @ W_gate) * (x @ W_value)) @ W_out
    relu_squared = F.relu(x @ W_gate).square() @ W_out
    assert y.shape == x.shape
    assert not torch.allclose(y, relu_squared)

    预期:shape 相同,但两种 MLP 数值不同

    这一章结束后,应能回答

    • 训练时并行计算所有位置,为何仍不会看到未来 token。
    • RoPE、RMSNorm 与 SwiGLU 分别改写 attention 的位置关系、残差尺度和逐位置非线性。
    • 为什么“nanochat 是 decoder-only Transformer”不等于“它实现了所有课程默认组件”。

    CHAPTER 03 · LEARNING

    模型如何学会预测下一个 token

    训练只做一次 shift:真实序列的前 \(T-1\) 个 token 是输入,后 \(T-1\) 个 token 是目标;一次 causal forward 同时产生所有位置的 next-token loss。

    tokens<bos>喜欢机器学习
    inputs<bos>喜欢机器
    targets喜欢机器学习

    TEACHER FORCING

    训练前缀来自数据,而不是模型

    位置 \(t\) 的预测以真实 \(x_{<t}\) 为条件。这样所有位置可以并行监督,梯度稳定;推理时前缀来自模型自己的历史输出,一次错误会改变后续条件,这就是训练/推理分布差异的来源。

    \[\mathcal{L}(x)=-\sum_{t=1}^{T-1}\log p_\theta(x_{t+1}\mid x_{\le t}).\]

    PERPLEXITY

    平均 token NLL 的指数

    \[\operatorname{PPL}=\exp\left(\frac{1}{N}\sum_{i=1}^{N}-\log p_\theta(x_i\mid x_{<i})\right).\]

    直观上,它是模型在每个位置面对的“等效候选数”。但这个解释只在同一 tokenizer、同一数据与同一计数口径下可靠。

    perplexity 能说明perplexity 不能单独说明需要补充
    模型是否给评测文本较高概率回答是否事实正确任务级事实性评测
    同 tokenizer / 数据口径下的 LM 拟合跨 tokenizer 的绝对优劣bits per byte / byte-normalized loss
    平均 next-token calibration 的一部分指令遵循、安全与长程一致性生成式 benchmark 与人工审计
    常见 token 上的总体趋势关键答案 token 是否预测正确conditional / span-level 指标

    LAB 03

    Teacher forcing 对齐器

    拖动当前位置,查看输入、目标、目标概率、token NLL 与整段 perplexity。

    input
    target
    target p
    token NLL
    sequence PPL

    位置 2 用真实前缀预测下一个 token;所有位置可在一次 forward 中并行计算。

    SOURCE BOUNDARY

    CS336 2026 A1,pp.28–29 定义 next-token objective 与 perplexity;Lecture 12 明确指出 perplexity 仍常用但与真实任务需求可能错位。核验材料没有直接把 teacher forcing 作为独立术语或 scheduled sampling 专题讲授;本页用这个标准术语描述课程公式隐含的“真实前缀监督”机制。nanochat 的 core_eval.py 左移 target,loss_eval.py 则用 bits per byte 降低词表大小对评测口径的影响。

    RUNNABLE EXAMPLE 11

    检查 input / target 一位偏移

    下载
    inputs = tokens[:, :-1]
    targets = tokens[:, 1:]
    loss = F.cross_entropy(logits.reshape(-1, vocab), targets.reshape(-1))
    
    assert torch.equal(inputs[:, 1:], targets[:, :-1])
    assert correct_loss < unshifted_loss

    预期:正确 shift 的 loss 显著更低

    RUNNABLE EXAMPLE 12

    从 NLL 计算 perplexity 与 BPB

    下载
    ppl = math.exp(sum(token_nlls) / len(token_nlls))
    bpb = sum(token_nlls) / (math.log(2) * byte_count)
    
    assert ppl >= 1
    assert tokenization_a_ppl != tokenization_b_ppl
    assert math.isclose(tokenization_a_bpb, tokenization_b_bpb)

    预期:PPL 随 token 切分变化,BPB 按 byte 归一化

    这一章结束后,应能回答

    • teacher forcing 为什么既允许训练并行,又带来推理前缀分布差异。
    • perplexity 与 cross-entropy 的精确关系,以及为何它必须在相同 tokenizer/数据口径下比较。
    • 低 perplexity 为什么不等于事实正确、指令遵循或生成文本没有重复。

    CHAPTER 04 · INFERENCE

    生成为什么分成 prefill 和 decode

    同一模型有两种负载:prefill 一次处理完整 prompt,decode 每步只处理一个新 token。KV cache 连接两者,用显存换取前缀计算复用。

    PREFILL

    一次读完整 prompt

    • 输入 shape:[B,P]
    • 产生所有 prompt 位置的 K/V
    • token 维度可并行
    • 矩阵乘法规模大,常偏 compute-bound
    same weights
    different workload

    DECODE

    每步读一个新 token

    • 输入 shape:[B,1]
    • 追加一个位置的 K/V
    • 下一步依赖上一步 token
    • 权重与 cache 读取频繁,常偏 memory-bound

    KV cache 保存什么

    第 \(\ell\) 层对历史 token 已经计算过的 key 与 value。新 token 到来时,只计算它自己的 \(q_t,k_t,v_t\),把 \(k_t,v_t\) 追加进 cache,并让 \(q_t\) 读取所有可见历史 K/V。

    \[ K^{(\ell)}_{1:t}=[K^{(\ell)}_{1:t-1};k^{(\ell)}_t],\quad V^{(\ell)}_{1:t}=[V^{(\ell)}_{1:t-1};v^{(\ell)}_t],\quad y_t=\operatorname{softmax}\left(\frac{q_tK_{1:t}^{\top}}{\sqrt{d_h}}\right)V_{1:t}. \]

    MEMORY ACCOUNTING

    每个序列的 KV cache 字节数

    \[\text{bytes}=L\times 2\times B\times T\times H_{kv}\times d_h\times \text{bytes(dtype)}.\]

    因子 2 来自 K 与 V。GQA/MQA 通过减少 \(H_{kv}\) 降低 cache;量化或更短 context 也会降低显存。cache 并没有消除新 query 对历史 K/V 的读取,因此 decode 仍随 context 增长。

    LAB 04

    Prefill / decode 时间线

    改变 prompt 与生成长度,比较没有 cache 的重复前向和 cache 驱动的单 token decode。

    prefill tokens8
    decode forwards after g14
    naive token work
    cached token work

    prefill 可并行处理 8 个 prompt token,并从最后位置 logits 选出 g1;再做 4 次单-token decode forward,共生成 5 个 token。

    LAB 05

    KV cache 显存账本

    直接修改结构与上下文参数,计算每 token 和完整 batch 的 cache 占用。

    shape / layer[1,4096,8,128] × K,V

    bytes / token

    total cache

    attention history read

    改变 context 会线性改变 cache 容量,也会增加每次 decode 读取的历史 K/V。

    SOURCE MAP

    CS336 2026 Lecture 10 将 prefill 描述为接近训练的并行、compute-bound 阶段,将 generation/decode 描述为逐 token、memory-bound 阶段,并进一步讨论 KV cache。nanochat 的 KVCache 预分配 [layers,B,T,Hkv,D],随后在 generate 中先 batch=1 prefill、复制 cache,再逐 token decode。

    RUNNABLE EXAMPLE 07

    Prefill 多 query,decode 单 query

    下载
    prefill_output = causal_attention(prompt)       # [P,d]
    full_output = causal_attention(cat(prompt, new_token))
    decode_output = attention(q_new, K_all, V_all) # [1,d]
    
    assert torch.allclose(decode_output, full_output[-1:])

    预期:decode 输出与完整重算最后位置一致

    RUNNABLE EXAMPLE 08

    追加 KV 并核对显存

    下载
    cache.append(prompt_k, prompt_v)
    cache.append(new_k, new_v)
    bytes_used = T * layers * batch * heads * head_dim * 2 * 4
    
    assert cache.position == T
    assert torch.allclose(cached_attention, full_attention)

    预期:cache position 增长,输出数值一致

    cache 不是免费加速

    它显著减少重复计算,但会占用随 batch、context、layer 与 KV heads 线性增长的显存;长上下文 decode 仍需读取历史 K/V。

    位置偏移必须跟 cache 对齐

    decode 的 RoPE position 应从 cache 当前长度开始。如果每一步都从 position 0 旋转,shape 正确但相对位置语义已经损坏。

    这一章结束后,应能回答

    • prefill 与 decode 为什么使用同一模型,却有不同的并行度与性能瓶颈。
    • KV cache 精确保存哪些张量,为什么不能缓存 Q。
    • 如何从层数、KV heads、head dim、dtype、batch 和 context 推导 cache 显存。

    CHAPTER 05 · DECISION

    概率分布如何变成最终文本

    模型只输出 logits。greedy、temperature、top-p、repetition penalty 与 EOS 共同定义“如何选 token、怎样避免退化、何时停止”。

    Temperature 改变相对 logit 间距

    \[p_i(\tau)=\frac{\exp(z_i/\tau)}{\sum_j\exp(z_j/\tau)}.\]

    \(\tau<1\) 放大差距,分布更尖;\(\tau>1\) 压缩差距,分布更平。greedy 直接取 \(\arg\max_i z_i\),可视为 \(\tau\to0^+\) 的选择极限,但实现应显式走 argmax,避免除零。

    TOP-P / NUCLEUS

    候选集大小随分布变化

    先按概率降序,取累计概率首次达到 \(p\) 的最小集合,再在集合内重新归一化采样。模型很确定时集合可能只有一个 token;分布平坦时会保留更多候选。

    \[S_p=\min\left\{S:\sum_{i\in S}p_i\ge p\right\}.\]

    REPETITION

    修改已出现 token 的 logits

    常见 sign-aware penalty:正 logit 除以 \(r>1\),负 logit 乘以 \(r\)。它能抑制循环,但也可能损伤姓名、代码变量、引用或句法上必要的重复,因此不是单调增加的“质量旋钮”。

    EOS

    停止是状态机,不是删掉字符串

    采样到 EOS / assistant-end 时应停止当前序列;还要有 max tokens 作为硬边界。batch generation 中每一行可能在不同时间完成,已完成行不能继续追加普通 token。

    LAB 06

    采样与停止游乐场

    同一组 logits 下切换策略,观察概率、top-p 候选集、重复惩罚和 EOS 状态。

    candidate set:

    selected
    entropy
    staterunning

    先比较不同 temperature 下最大概率的变化,再切到 top-p 观察候选集大小。

    SOURCE BOUNDARY

    CS336 2026 A1,pp.37–42 明确覆盖 temperature、top-p 与 end-of-text stopping;核验材料没有直接系统讲授 repetition penalty,也没有把 greedy 作为独立重点术语。nanochat 的 sample_next_token 实现 temperature=0 的 argmax、temperature 与 top-k,当前没有 top-p 或通用 repetition penalty;生成状态机<|assistant_end|> 或 BOS 被采样时结束当前行。

    RUNNABLE EXAMPLE 09

    同一 logits 比较三种策略

    下载
    greedy = max(range(len(logits)), key=logits.__getitem__)
    warm_probs = softmax(logits, temperature=1.5)
    candidates = nucleus(warm_probs, p=0.8)
    sampled = sample(candidates, seed=9)
    
    assert sampled in {token for token, _ in candidates}

    预期:低 temperature 更尖,top-p 只在候选集采样

    RUNNABLE EXAMPLE 10

    重复惩罚与 EOS 停止

    下载
    adjusted = penalize(logits, set(generated), penalty=1.5)
    next_token = max(range(len(adjusted)), key=adjusted.__getitem__)
    if next_token == EOS:
        break
    generated.append(next_token)
    
    assert generated == [0, 1]

    预期:生成两个普通 token 后遇到 EOS 停止

    END-TO-END LOOP

    把五章重新接起来

    1. Tokenizer 把 prompt 编成 ids,并注入明确的 special token。
    2. 模型对 prompt 做 prefill,把每层 K/V 写入 cache。
    3. 取最后位置 logits,应用 repetition、temperature 与 top-p。
    4. 选择 next token;若是 EOS 则结束,否则追加到序列。
    5. 只对新 token 做 decode,追加新的 K/V,回到第 3 步。
    ids = tokenizer.encode(prompt)
    logits, cache = model.prefill(ids)
    while len(generated) < max_tokens:
        logits = apply_repetition_penalty(logits, generated)
        next_id = sample(logits, temperature, top_p)
        if next_id == eos_id:
            break
        generated.append(next_id)
        logits, cache = model.decode(next_id, cache)

    这一章结束后,应能回答

    • greedy、temperature 与 top-p 分别在哪一步改变 token 选择。
    • repetition penalty 为何可能缓解循环,也可能伤害必要重复。
    • EOS、max tokens 与 batch 中逐行完成应怎样组成停止状态机。

    SOURCES · VERIFIED 2026-07-22

    来源、实现差异与继续学习

    课程材料告诉我们教学主线,项目源码告诉我们具体取舍;两者不一致的地方,往往正是理解工程边界的入口。

    概念CS336 2026nanochat @ 92d63d4本页处理
    BPEA1 实现 byte-level BPErustbpe 训练、tiktoken 推理机制 + UTF-8 round-trip
    Decoder / mask / RoPE / RMSNormA1 现代 decoder 主线直接采用,并加入 GQA/FlashAttention 等系统取舍聚焦最小标准机制
    SwiGLUA1 明确要求使用 ReLU²并排比较,不抹平差异
    Prefill / decode / KV cacheLecture 10 机制与系统瓶颈真实 cache 与两阶段生成路径时间线 + 显存账本
    Temperature / top-pA1 明确覆盖 top-ptemperature + top-k独立实现 top-p
    Repetition没有直接系统讲授没有通用 penalty标准补充机制,明确风险
    Teacher forcingnext-token 公式隐含,术语未直接展开shifted targets 可直接追踪用标准术语解释机制
    Perplexity重要但不充分主推 bits per byte定义、局限与 BPB 对照

    真正理解,而不是只记住名词

    如果你能从一段文本开始,画出每一步的 shape;解释训练前缀与推理前缀的差异;写出 cache 显存公式;在同一 logits 上手算 temperature 和 top-p;最后说明 perplexity 为什么重要但不充分,那么这 12 个概念已经组成了一个可工作的心智模型。