QuanZhou's Wiki

M2-B 工作表:无 Cache MHA Oracle

状态:Archived · 未填写模板。历史题目与当时的评分方法见 M2-B Assignment;当前任务从阶段 2 路线页进入。

本文件是网站公开的空白模板。下载者可以在自己的副本中填写;本站维护者必须填写私人作答仓库中的同名文件,不得把答案写回这里。

学习合同

Question: 给定 Q/K/V=[B,H,T,D_head],怎样证明批量矩阵乘法没有混合 Batch 或 Head,并且每个位置只能读取历史 Token?

Prediction: 运行前填写 Score、Weight、Output Shape,Softmax Axis,以及 Future Token、Head、Batch 三种修改各会影响哪个切片。

Action: 先写独立 Oracle 与隔离性测试,再实现 H_q=H_kv 的无 Cache MHA。

Artifact: src/inference_lab/multi_head_attention.py
          tests/test_multi_head_attention.py
          grader_tests/test_m2_b_causal_attention.py
          checkoffs/m2-b-causal-attention.md

Acceptance: Shape、独立标量 Oracle、Causal Mask、归一化、隔离性、非法 Shape 与完整回归全部通过;最后关闭代码重做推导。

Feedback: unittest 的第一次失败、手算非零坐标、完整回归和闭卷解释。

Next decision: 全部通过后进入 M2-C Cached MHA;否则只修正 M2-B。

本轮边界

固定输入与运行前预测

固定互不相等的关键维度:

B=2, H=3, T=4, D_head=5
Q/K/V Shape = [2,3,4,5]

先关闭 multi_head_attention.py,填写:

K 转置最后两轴后的 Shape:
Score = Q @ K^T 的 Shape:
除以 sqrt(D_head) 改变 Shape 吗:
Causal Mask 的 Shape:
Softmax 应沿哪个 Axis:
Weight 的 Shape:
Output = Weight @ V 的 Shape:

修改 K/V[b,h,j] 且 j>i,对 Output[b,h,i] 的预测:
修改 Q/K/V[b,h],对其他 Head 与其他 Batch 的预测:

再选择两个非零坐标,写出求和范围;至少一个坐标必须满足 i > 0:

Score[b=?,h=?,i=?,j=?] =
Output[b=?,h=?,i=?,d=?] =

15 分钟前置阅读

只读本轮第一段实现会使用的内容,每处留一句确认或修正:

Annotated Transformer 确认或修正了:
matmul 确认或修正了:
triu 确认或修正了:

先写会失败的测试

在 tests/test_multi_head_attention.py 新增测试,Oracle 不得调用待测函数或生产 softmax:

第一次失败必须原样保存在“运行后记录”,不能在实现通过后补写一个更好看的错误。

最小实现边界

先实现一个可观察的教育接口:

def multi_head_causal_attention(
    q: np.ndarray,
    k: np.ndarray,
    v: np.ndarray,
) -> tuple[np.ndarray, np.ndarray]:
    """Return per-head output and attention weights without a KV cache."""

核心 Shape 路径固定为:

Q/K/V              [B,H,T,D_head]
swapaxes(K,-1,-2)  [B,H,D_head,T]
Score/Weight        [B,H,T,T]
Output              [B,H,T,D_head]

必须使用数值稳定 Softmax:先减去最后一维最大值,再 exp 和归一化。不要加入 Cached Decode、GQA 映射、merge_heads 或 Output Projection。

120 分钟执行顺序

验证命令

在 M2-B 当时的 Starter 中,先运行该练习的公开评分:

$ make m2-b

通过后再运行整份作业和已完成任务的回归:

$ make grade
$ make test

make m2-b、./grade-lab m2-b 与 make GRADEFLAGS=m2-b grade 等价。公开评分器固定 BLAS 线程并从 grader_tests/test_m2_b_causal_attention.py 验证 Shape、标量 Oracle、Mask、归一化、隔离性和输入契约。

运行后记录

实际 Score/Weight/Output Shape:
两个独立坐标 Oracle:
Future Token 是否不可见:
Head/Batch 是否隔离:
触发的非法配置:
第一次失败信息:
错误属于 Score / Mask / Softmax Axis / Value / 输入契约中的哪一层:
修正后的规则:
定向测试:
完整回归:
闭卷解释:为什么除以 sqrt(D_head),Softmax 为什么沿最后一个 Axis:

M2-B 验收

待完成

全部通过后进入 M2-C Cached MHA;任一项未通过时继续修正 M2-B,不提前实现 Cache 或 GQA。