Hang Zhengyang
English ↗

Attention Block 推理路径:从输入 Shape 到 Vocab Logits(单请求版,含数字例子)

本文只讨论**推理(inference)**路径,不涉及训练反向传播、参数更新、dropout。目标是把 input -> attention block -> logits over vocab 串成一条连续 shape 链路,并在每个步骤给出具体数字例子。全文默认单请求,不引入 batch 维。

1. 记号与统一数字配置

通用符号:

  • S:本次输入 token 数(prefill 常见 S>1,decode 常见 S=1)
  • Thist:历史 KV cache 长度
  • T:当前可见总长度,T=Thist+S
  • dmodel:隐藏维度
  • H:attention 头数(标准 MHA 下,Q/K/V 头数相同)
  • dhead:每头维度,通常 dmodel=H·dhead
  • V:词表大小

本文数字例子(decode 一步)统一使用:

  • S=1(单请求)
  • Thist=128
  • T=129
  • dmodel=4096
  • H=32
  • dhead=128(32·128=4096)
  • V=32000

输入 token id:

这是一个长度为 S 的 token id 序列。
数字化 shape:𝐈:1。

2. 从 Token 到首层输入

token embedding + 位置编码(RoPE/ALiBi 等)可写成:

𝐗0=Embed(𝐈)+PosEnc

嵌入输出是二维隐藏表示,shape 为 S×dmodel。
数字化 shape:𝐗0:1×4096。 这一步可以理解为:把离散 token id 变成连续向量,并注入位置信息,供后续注意力计算。

3. 单个 Attention Block(Pre-Norm)

以第 ℓ 层为例,输入:

𝐗ℓ 的 shape 为 S×dmodel。
数字化 shape:𝐗ℓ:1×4096。

3.1 LayerNorm + QKV 线性映射

𝐗^ℓ=LN(𝐗ℓ)

LN 不改 shape:1×4096。 作用是先把数值尺度稳定下来,让 Q/K/V 投影更平稳。

更展开地看,对单个 token 向量 𝐱∈ℝdmodel,LayerNorm 在特征维上做:

μ=1dmodel∑i=1dmodelxi,σ2=1dmodel∑i=1dmodel(xi−μ)2
LN(xi)=γixi−μσ2+ϵ+βi

其中 γ,β∈ℝdmodel 是可学习参数,ϵ 是数值稳定项。
直觉上,它先把激活“标准化”,再用 γ,β 恢复可学习的缩放与平移能力。

𝐐=𝐗^ℓ𝐖Q,𝐊=𝐗^ℓ𝐖K,𝐕=𝐗^ℓ𝐖V

其中:

  • 𝐖Q,𝐖K,𝐖V∈ℝdmodel×(Hdhead)

先看 projection 后最后一维:

  • 𝐐raw:S×(Hdhead)
  • 𝐊raw:S×(Hdhead)
  • 𝐕raw:S×(Hdhead)

数字化:

  • 𝐐raw:1×4096
  • 𝐊raw:1×4096
  • 𝐕raw:1×4096

reshape 成多头:

  • 𝐐,𝐊new,𝐕new:H×S×dhead

数字化:

  • 𝐐:32×1×128
  • 𝐊new:32×1×128
  • 𝐕new:32×1×128 直觉上,这是把同一个 token 的表示切成 32 个子空间并行建模,每个头各自关注不同关系。

拆分过程可写成(以 𝐐raw 为例):

𝐐raw∈ℝS×(Hdhead)→reshape𝐐∈ℝH×S×dhead

在本例里,Hdhead=32×128=4096,所以本质只是把最后一维 4096 重排为 32 x 128。

3.2 与 KV Cache 拼接(推理关键)

历史缓存:

  • 𝐊cache:H×Thist×dhead
  • 𝐕cache:H×Thist×dhead

数字化:

  • 𝐊cache:32×128×128
  • 𝐕cache:32×128×128

拼接后:

𝐊all=Concat(𝐊cache,𝐊new)∈ℝH×T×dhead
𝐕all=Concat(𝐕cache,𝐕new)∈ℝH×T×dhead

数字化:

  • 𝐊all:32×129×128
  • 𝐕all:32×129×128 因为是 decode 场景(S=1),每一步只新增一个位置,历史部分复用 cache,不重复算旧 token 的 K/V。

标准 MHA 下,Q/K/V 头数一致,均为 H。

3.3 注意力分数到权重

𝐒=𝐐𝐊all⊤dhead+𝐌

其中 𝐌 是因果 mask。 它保证当前位置只能看见历史与当前,不能看见未来 token。

𝐒 的 shape 为 H×S×T。
数字化:𝐒:32×1×129。

softmax 后:

𝐏=softmax(𝐒,last dim)

数字化:𝐏:32×1×129。 也就是每个头都会产出一条长度 129 的注意力分布。

3.4 权重加权 Value 得到上下文

𝐂=𝐏𝐕all

𝐂 的 shape 为 H×S×dhead。
数字化:𝐂:32×1×128。 𝐂 可看作“按注意力权重汇总后的上下文摘要”。

合并多头后,𝐂merge 的 shape 为 S×(Hdhead)。
数字化:𝐂merge:1×4096。

合并过程与拆分相反:

𝐂∈ℝH×S×dhead→transpose + reshape𝐂merge∈ℝS×(Hdhead)

实现上通常先把维度从 (H,S,dhead) 调整到 (S,H,dhead),再把后两维展平。
因为元素总数不变(本例每个位置仍是 32×128=4096),所以这一步是无信息损失的重排。

输出投影:

𝐎attn=𝐂merge𝐖O,𝐖O∈ℝ(Hdhead)×dmodel

数字化:𝐎attn:1×4096。

残差连接:

𝐘ℓ=𝐗ℓ+𝐎attn

数字化:𝐘ℓ:1×4096。 残差连接让模型在“保留原信息”和“引入新信息”之间平衡,也让深层堆叠更稳定。

3.5 MLP 子层(推理前向)

𝐘^ℓ=LN(𝐘ℓ)
𝐔=ϕ(𝐘^ℓ𝐖up)⊙(𝐘^ℓ𝐖gate)
𝐎mlp=𝐔𝐖down
𝐗ℓ+1=𝐘ℓ+𝐎mlp

数字化(关注输入输出主 shape):

  • LN(𝐘ℓ):1×4096
  • 𝐎mlp:1×4096
  • 𝐗ℓ+1:1×4096 可以把这部分理解成逐位置的非线性特征变换:attention 负责“找信息”,MLP 负责“加工信息”。

4. 多层堆叠到最终 Logits

经过 L 层后:

𝐗L 的 shape 为 S×dmodel。
数字化:𝐗L:1×4096。

最终归一化:

𝐇=LNfinal(𝐗L)

数字化:𝐇:1×4096。

映射到词表:

𝐙=𝐇𝐖lm⊤+𝐛lm

其中:

  • 𝐖lm∈ℝV×dmodel
  • 𝐙∈ℝS×V

数字化:𝐙:1×32000。 𝐙 中每个值都是对应词表 token 的未归一化分数(logit)。

推理采样通常取最后位置:

  • 𝐙last=𝐙[:,−1,:]
  • 数字化:𝐙last:32000

5. 一条完整的 Shape 链(通用 + 数字)

通用链路可读作:
𝐈(S) → 𝐗0(S,dmodel) → L 层 Block → 𝐇(S,dmodel) → 𝐙(S,V) → 𝐙last(V) → next token。

数字例子对应:
(1) → (1,4096) → L 层 block → (1,4096) → (1,32000) → (32000)。

6. Prefill 与 Decode 的区别(推理视角)

  • Prefill:S 较大,一次性处理多个 token 并构建初始 cache。
  • Decode:通常 S=1,每步只算新 token 的 Q/K/V,再与历史 cache 拼接;(单请求省略 batch 维)代价主要在长度 T 的注意力读取。

两者公式一致,主要区别是 S、T 和 cache 状态。