Transformer解码器推理阶段注意力层token使用与矩阵形状匹配问题
问题解答
1. 仅用单个token做投影的原因
你观察到的解码器自注意力QKV投影、编解码注意力Q投影仅用到单个token的现象,是Transformer推理阶段普遍采用的**KV缓存(KV Cache)**优化导致的,属于工程优化的正常结果,并不违背Transformer的基础逻辑。
原生无优化的推理流程中,每生成第k个token时都需要将前k个历史token全部输入解码器,重新计算所有位置的QKV,计算量会随生成步数线性增长,效率极低。KV缓存的核心思路就是把每一步计算得到的K、V张量存储下来,下一步推理时仅输入新生成的单个token,只计算该token对应的Q、K、V,再将新的K、V追加到历史缓存中,直接用新token的Q和全部历史KV做注意力计算即可,完全不需要重复计算历史token的投影,因此你观察到的投影层输入永远只有1个token。
2. 可变长度下的张量形状匹配逻辑
我们先统一约定维度符号:
B:批次大小d_model:模型隐层维度h:注意力头数,单头维度d_k = d_model / hS:编码器输出序列长度(推理阶段固定,因为源输入是固定的)k:当前已生成的解码器序列长度(每生成一个token就+1,是可变值)
解码器自注意力形状匹配
分两种实现场景说明:
- 无KV缓存的原生实现
输入为前k个历史token嵌入,形状为(B, k, d_model)- Q/K/V投影:
(B, k, d_model) × (d_model, d_model) = (B, k, d_model),拆分为多头后形状为(B, h, k, d_k) - 注意力计算:
Q × K^T得到形状(B, h, k, k),掩码后乘V得到(B, h, k, d_k),合并头后输出形状为(B, k, d_model)
投影层权重是固定大小的,仅要求输入最后一维和权重第一维匹配即可,序列维度k可变不会影响矩阵乘的合法性。
- Q/K/V投影:
- 带KV缓存的优化实现
输入仅为当前新生成的单个token嵌入,形状为(B, 1, d_model)- Q投影:
(B, 1, d_model) × (d_model, d_model) = (B, 1, d_model),拆分为多头后形状为(B, h, 1, d_k) - K投影:同上得到
(B, h, 1, d_k),追加到历史缓存K(形状(B, h, k-1, d_k))后,总K形状为(B, h, k, d_k) - V投影:同上得到
(B, h, 1, d_k),追加到历史缓存V(形状(B, h, k-1, d_k))后,总V形状为(B, h, k, d_k) - 注意力计算:
Q × K^T得到形状(B, h, 1, k),乘V得到(B, h, 1, d_k),合并头后输出形状为(B, 1, d_model)
你观察到的自注意力环节可变形状的BGemms,就是这里Q和总K、计算结果和总V的矩阵乘操作,因为K/V的序列长度随步数k增长,所以形状是可变的。
- Q投影:
编解码注意力形状匹配
编码器输出在推理阶段是固定的,因此编解码注意力的K、V可以在推理开始前一次性计算并缓存下来,全程不需要重复计算,形状固定为(B, h, S, d_k)。
编解码注意力的Q全部来自解码器自注意力的输出:
- 无KV缓存时,Q形状为
(B, k, d_model),投影后为(B, h, k, d_k),和固定形状的K/V计算注意力得到(B, h, k, d_k) - 带KV缓存时,Q形状为
(B, 1, d_model),投影后为(B, h, 1, d_k),和固定形状的K/V计算注意力得到(B, h, 1, d_k)
因为K/V的形状全程固定,所以你观察到编解码注意力环节的Gemm都是固定大小的,完全符合你的观测结果。
内容的提问来源于stack exchange,提问作者Joe Black
相关产品推荐
相关产品推荐

