Tensor Shape

张量 shape 不是排版信息,而是模型中“谁和谁相乘、沿哪个轴归约、结果回到哪里”的接口契约。看不懂 shape,就无法判断一个 attention 实现到底是在 batch、head 还是 sequence 维度上做了错误广播。

统一符号与数据流

  • B:并行请求或样本数;T:序列长度;D:hidden size。
  • Hq / Hkv:Query 与 Key/Value head 数;Dh:每个 head 的维度;V:词表大小。
  • Tq / Tk:query 与 key 的长度。prefill 中二者通常都大于 1,decode 中 Tq=1Tk 包含历史。
1
2
3
4
5
6
token ids [B,T]
→ embedding/hidden states [B,T,D]
→ Q/K/V projection
→ attention scores [B,Hq,Tq,Tk]
→ merge heads [B,Tq,D]
→ block output → logits [B,T,V]

Attention shape ledger

张量 形状 轴的含义
hidden states [B,T,D] 每个 token 一个 D 维表示
Q [B,T,Hq×Dh] → [B,Hq,T,Dh] 每个 query head 的向量
K/V(MHA) [B,Hq,T,Dh] 每个 KV head 独占一份历史
K/V(GQA/MQA) [B,Hkv,T,Dh] 多个 Q head 共享 KV head
scores [B,Hq,Tq,Tk] 每个 query 对每个 key 的权重前 logits
attention output [B,Hq,Tq,Dh] 每个 head 聚合后的结果
merged heads [B,Tq,Hq,Dh] → [B,Tq,D] 拼回 hidden 维度
logits [B,T,V] 每个位置对词表的分数

最容易漏掉的是 scores 的两个序列轴:Tq 决定要计算多少个 query,Tk 决定每个 query 能读取多少历史。于是 attention 的计算和临时内存随 Tq×Tk 增长,而 KV Cache 的持久内存随历史 Tk 增长。

一个 decode 例子

假设 B=2Hq=8Hkv=2Dh=64,两个请求各自已经有不同长度的历史。当前一步只生成一个 token,因此:

1
2
3
4
Q: [2, 8, 1, 64]
K/V cache: [2, 2, Tk, 64]
scores: [2, 8, 1, Tk]
output: [2, 8, 1, 64] → [2, 1, 512]

这里不能把 Tk 当成 batch 维,也不能因为 Tq=1 就删除序列轴;scheduler 可能把不同长度请求 padding、packing 或用 block table 组织起来,逻辑轴仍必须清楚。

变形操作的语义

  • unsqueeze/squeeze 只改变轴数,不复制数据;错误地 squeeze 非 1 维会改变契约。
  • stack 新增轴,cat 沿已有轴拼接;二者都要求非拼接轴兼容。
  • transpose/permute 改变视图的步长,不保证连续内存。
  • view 依赖连续布局;transpose 后用 reshapecontiguous().view(...),否则可能报错或得到错误假设。

排错顺序与边界

看到 shape error 时先写出每个轴的语义,再检查矩阵乘法的收缩轴、mask 的广播轴和 device/dtype;不要先靠 reshape(-1, ...) 把错误压平。shape 正确也不代表语义正确:[B,T,D] 可以被错误地当成 [T,B,D],数值仍能运行但训练会悄悄失真。

它向上连接 Transformer Block,向下约束 Scaled Dot-Product AttentionKV Cache,横向对照 Multi-Head AttentionGQAMQA

广播最危险的地方

广播不是“形状差不多就能算”。例如 padding mask 可能是 [B,1,1,Tk],head mask 可能是 [B,Hq,1,1];二者相乘后虽然能得到 [B,Hq,1,Tk],但如果把 [B,T] 误 reshape 成 [B,1,T,1],错误会落在 query 轴上,代码仍然能运行。判断广播是否正确,必须逐轴写出它覆盖的是 batch、head、query 还是 key。

GQA 还会引入一个容易被忽略的语义:Hq/Hkv 是逻辑映射,不一定意味着把 K/V 在显存中真的复制多份。一个 kernel 可能在读取时根据 q_head // group_size 寻址;另一个朴素实现则先 repeat,结果数值相同但峰值显存和带宽完全不同。遇到性能回归时,shape ledger 要和 profiler 里的实际 stride、contiguous 状态一起看。

Ragged batch 与 packed sequence

真实 serving 中不同请求的 Tk 往往不同。padding 把它们补齐到同一长度,形状简单却浪费计算;packing 或 varlen attention 使用 cu_seqlens 把多个序列压成一个 token 平面,逻辑上仍有独立的 [Tq,Tk] 区域。此时不能只从一个二维 tensor 猜出 batch 边界,要同时检查 offsets、block table 和 mask。很多“第二个请求读到了第一个请求内容”的事故,本质是 packed 索引或 cache slot 的 shape 语义丢了。

用小尺寸守住契约

调试 attention 时先用 B=2, Hq=4, Hkv=2, Tq=1, Tk=3, Dh=2 的手算样例,给每个 batch/head/key 填入唯一编号,再验证:Q/K 点积收缩的是 Dh,softmax 只沿 Tk,不同 batch 不互相可见,GQA 的 Q head 0/1 只读 KV head 0。这个测试比随机输入更容易抓住错误广播,因为任何轴交换都会留下可追踪的编号。