Tensor Shape
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=1而Tk包含历史。
1 | token ids [B,T] |
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=2、Hq=8、Hkv=2、Dh=64,两个请求各自已经有不同长度的历史。当前一步只生成一个 token,因此:
1 | Q: [2, 8, 1, 64] |
这里不能把 Tk 当成 batch 维,也不能因为 Tq=1 就删除序列轴;scheduler 可能把不同长度请求 padding、packing 或用 block table 组织起来,逻辑轴仍必须清楚。
变形操作的语义
unsqueeze/squeeze只改变轴数,不复制数据;错误地 squeeze 非 1 维会改变契约。stack新增轴,cat沿已有轴拼接;二者都要求非拼接轴兼容。transpose/permute改变视图的步长,不保证连续内存。view依赖连续布局;transpose 后用reshape或contiguous().view(...),否则可能报错或得到错误假设。
排错顺序与边界
看到 shape error 时先写出每个轴的语义,再检查矩阵乘法的收缩轴、mask 的广播轴和 device/dtype;不要先靠 reshape(-1, ...) 把错误压平。shape 正确也不代表语义正确:[B,T,D] 可以被错误地当成 [T,B,D],数值仍能运行但训练会悄悄失真。
它向上连接 Transformer Block,向下约束 Scaled Dot-Product Attention 与 KV Cache,横向对照 Multi-Head Attention、GQA 和 MQA。
广播最危险的地方
广播不是“形状差不多就能算”。例如 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。这个测试比随机输入更容易抓住错误广播,因为任何轴交换都会留下可追踪的编号。