Attention Block 推理路径:从输入 Shape 到 Vocab Logits(单请求版,含数字例子)
本文只讨论**推理(inference)**路径,不涉及训练反向传播、参数更新、dropout。目标是把 input -> attention block -> logits over vocab 串成一条连续 shape 链路,并在每个步骤给出具体数字例子。全文默认单请求,不引入 batch 维。
1. 记号与统一数字配置
通用符号:
- :本次输入 token 数(prefill 常见 ,decode 常见 )
- :历史 KV cache 长度
- :当前可见总长度,
- :隐藏维度
- :attention 头数(标准 MHA 下,Q/K/V 头数相同)
- :每头维度,通常
- :词表大小
本文数字例子(decode 一步)统一使用:
- (单请求)
- ()
输入 token id:
这是一个长度为 的 token id 序列。
数字化 shape:。
2. 从 Token 到首层输入
token embedding + 位置编码(RoPE/ALiBi 等)可写成:
嵌入输出是二维隐藏表示,shape 为 。
数字化 shape:。
这一步可以理解为:把离散 token id 变成连续向量,并注入位置信息,供后续注意力计算。
3. 单个 Attention Block(Pre-Norm)
以第 层为例,输入:
的 shape 为 。
数字化 shape:。
3.1 LayerNorm + QKV 线性映射
LN 不改 shape:。 作用是先把数值尺度稳定下来,让 Q/K/V 投影更平稳。
更展开地看,对单个 token 向量 ,LayerNorm 在特征维上做:
其中 是可学习参数, 是数值稳定项。
直觉上,它先把激活“标准化”,再用 恢复可学习的缩放与平移能力。
其中:
先看 projection 后最后一维:
数字化:
reshape 成多头:
数字化:
- 直觉上,这是把同一个 token 的表示切成 32 个子空间并行建模,每个头各自关注不同关系。
拆分过程可写成(以 为例):
在本例里,,所以本质只是把最后一维 4096 重排为 32 x 128。
3.2 与 KV Cache 拼接(推理关键)
历史缓存:
数字化:
拼接后:
数字化:
- 因为是 decode 场景(),每一步只新增一个位置,历史部分复用 cache,不重复算旧 token 的 K/V。
标准 MHA 下,Q/K/V 头数一致,均为 。
3.3 注意力分数到权重
其中 是因果 mask。 它保证当前位置只能看见历史与当前,不能看见未来 token。
的 shape 为 。
数字化:。
softmax 后:
数字化:。 也就是每个头都会产出一条长度 129 的注意力分布。
3.4 权重加权 Value 得到上下文
的 shape 为 。
数字化:。
可看作“按注意力权重汇总后的上下文摘要”。
合并多头后, 的 shape 为 。
数字化:。
合并过程与拆分相反:
实现上通常先把维度从 调整到 ,再把后两维展平。
因为元素总数不变(本例每个位置仍是 ),所以这一步是无信息损失的重排。
输出投影:
数字化:。
残差连接:
数字化:。 残差连接让模型在“保留原信息”和“引入新信息”之间平衡,也让深层堆叠更稳定。
3.5 MLP 子层(推理前向)
数字化(关注输入输出主 shape):
- 可以把这部分理解成逐位置的非线性特征变换:attention 负责“找信息”,MLP 负责“加工信息”。
4. 多层堆叠到最终 Logits
经过 层后:
的 shape 为 。
数字化:。
最终归一化:
数字化:。
映射到词表:
其中:
数字化:。 中每个值都是对应词表 token 的未归一化分数(logit)。
推理采样通常取最后位置:
- 数字化:
5. 一条完整的 Shape 链(通用 + 数字)
通用链路可读作:
→ → 层 Block → → → → next token。
数字例子对应:
→ → 层 block → → → 。
6. Prefill 与 Decode 的区别(推理视角)
- Prefill: 较大,一次性处理多个 token 并构建初始 cache。
- Decode:通常 ,每步只算新 token 的 Q/K/V,再与历史 cache 拼接;(单请求省略 batch 维)代价主要在长度 的注意力读取。
两者公式一致,主要区别是 、 和 cache 状态。