M23:从零实现 Scaled Causal Attention
Scaled causal attention 用下三角可见性约束对 QKᵀ/√d_k 做 softmax,再以所得权重聚合 V。
内容类型:预习教材(不代表已完成)
日期:2026-09-15
阶段:P1 · AI Model Engineering 90
周次:W4 · Transformer 第一性原理
节奏:周二最小实现
状态:教材已备;学习未完成
标签:Attention、Tensor Shape、Causal Mask、Softmax、最小实现
一句话定义
Scaled causal attention 用下三角可见性约束对 QKᵀ/√d_k 做 softmax,再以所得权重聚合 V。
学习目标
- 能逐步写出单头、单 batch 的最小 attention。
- 对每一步标注 shape,避免广播产生静默错误。
- 正确区分布尔可见性 mask 与加性负无穷 mask。
- 设计三个不依赖训练的结构检查。
核心知识
最小输入 X 的 shape 可设为 [T,C]。线性投影后 Q、K、V 为 [T,d]。QKᵀ 得 [T,T],第 i 行表示位置 i 对所有 key 的分数。除以 √d 后,在未来列上填充负无穷,按最后一维 softmax,得到每行和为 1 的权重 A。A@[T,d] 的 V 后输出 [T,d]。
扩展到 batch 和 heads 时常见 shape 为 [B,H,T,D]。分数 [B,H,T,T],mask 要广播到 batch/head 而不误遮其他轴。transpose/permute 改的是轴顺序,view/reshape 只改变解释方式;张量不连续时直接 view 可能失败或产生错误理解,需要明确 contiguous 与内存布局。
softmax 应在 key 轴执行,因为每个 query 要在可见 keys 上分配权重。若在 query 轴 softmax,列和为 1,语义已改变但代码可能仍能运行。结构检查比看“输出有数字”更可靠。
机制与实现顺序
建议先不写类:
- 准备一个小 X 和固定投影矩阵。
- 计算 Q、K、V,并 assert shape。
- scores = QKᵀ/√d。
- 创建 T×T 下三角 mask;未来位置填极小值。
- weights = softmax(scores, dim=-1)。
- output = weights V。
三个不变量:每行权重和接近 1;上三角未来权重为 0;改变未来 token 不应改变更早位置输出。第三个是行为级 causal 检查,比仅查看 mask 数组更强。
最小练习
- 先实现 T=3、d=2 的单头版本,不加 dropout、batch 或模块封装。
- 打印或记录每一步预期 shape,实际运行后再核对。
- 用固定输入检查三条不变量。
- 把最后一个 token 改成极端值,比较位置 0、1 的输出是否不变。
- 故意把 softmax 轴改错,只预测现象,不需要保留错误代码。
常见误区
- mask 后填 0;softmax 仍会给未来位置非零权重。
- softmax 维度错误,但因为 shape 正确而未被发现。
- 忘记 scale,短例子仍运行就误以为无影响。
- 只看未来权重接近零,不检查未来输入是否能影响过去输出。
- 一开始实现 multi-head、dropout、cache,掩盖核心机制。
文档场景连接
decoder 在生成文档摘要的第 t 个 token 时只能使用已有 prompt 和先前输出。若 mask 有漏洞,训练 loss 会因偷看未来而异常乐观,部署生成却无法复现这种信息条件。
自检问题
- scores 与 weights 的 shape 分别是什么?
- 为什么 mask 通常在 softmax 前加负无穷?
- 如何用输入扰动检查 causal 属性?
- softmax 应沿哪个轴,语义是什么?
专业课程对齐
- Stanford CS336:对应 Transformer 语言模型从零实现,重点参考 attention 的投影、shape、mask 和数值稳定处理。
- Stanford CS224N:对应 scaled dot-product attention 与 Transformer,补齐实现背后的寻址解释。
- PyTorch Tutorials:对应张量矩阵乘、广播和 softmax 等官方用法,用于核对最小实现的框架语义。
深入学习提示
按“CS224N 公式 → CS336 实现 → PyTorch API”顺序学习。对输入 [B,T,C] 逐步标注 Q/K/V、scores [B,H,T,T] 与 output 的形状,并在 softmax 前施加上三角负无穷 mask。观察 1/√d_k 是否让 score 方差随 head dimension 稳定;代码检查每行 attention 权重和为 1、未来位置权重为 0、改变未来 token 不影响过去输出。反例包括 mask 方向反了、在 softmax 后再 mask、用有限大负数在低精度下泄漏、transpose 轴错但广播仍能运行。先用固定小张量和单头验证中间值,再扩 batch/head;不要以“loss 能下降”替代结构检查,因为带未来泄漏的模型往往下降得更快。
学后填写区
- 最小实现位置:____
- Shape 表:____
- 三条结构检查结果:____
- 尚未解释的差异:____
- 实际学习日期与用时:____