- SignalDesk2小时前
推理时只需要最后一个 token 的 Q,配合完整的历史 K 和 V,就能生成下一个 token!当然还有一个前提——大语言模型得是生成式的——类似transformer中encoder的结构——注意力的计算受因果关系约束——这种特性确保之前的q的对应输出不会被新的q、k、v影响!这就是为什么只需要对每层的KV进行缓存!会混淆的原因就是因为被训练时的矩阵计算迷惑了,但请仔细想想,训练时整个Q矩阵都参与计算是为了并行训练,而训练时也是使用上一个token对应的最终输出来预测下一个token!至于为什么每层都能缓存,因为一个transformer块里最后要经过的FFN是逐位置/token计算的——token之间不会相互影响,一个块最终输出的形状与序列刚经过embedding的形状是一致的,会被送往下一个块,后续就与第一块一致了! tip:同样根据transformer中去掉cross-attention的encoder的原理可知,用户的最初的输入是要全部送进去准备缓存的,而且QKV全部都要计算,因为计算肯定是要全部计算的,缓存只是把之前计算过的重复利用而已,不能只计算最后一个q,因为后面还有好几个相同的transformer块,它们计算序列中最后一个token对应的输出也是与第一层一样需要全部的k、v——完整的K、V矩阵,而后续的自生成过程中,因为有之前的K、V缓存,只需对新生成的token执行对应的计算(计算q、k、v和后续计算)并追加进缓存即可,而生成的token在块中的计算并不需要之前的Q! 如在新预测得到的token进行追加计算时,在一个transformer块中的计算过程—— 让我们再次推演一遍矩阵乘法,但这一次,我们已经缓存了前 4个 token 的 K 和 V 矩阵,并且只传入单个 token 的嵌入(embeddings)。 计算新的 Q 只会输出单行结果。 W_Q 与之前相同,没有改变。 \underbrace{\begin{bmatrix} 0.20 & -0.10 & 0.70 \end{bmatrix}}_{\text{嵌入}} \times \underbrace{\begin{bmatrix} -0.74 & 0.91 & -0.21 \\ -0.51 & 0.43 & 0.18 \\ -0.56 & -0.21 & 0.76 \end{bmatrix}}_{W_Q} = \underbrace{\begin{bmatrix} -0.49 & -0.01 & 0.48 \end{bmatrix}}_{Q(\text{新})} 随后,计算新的 K 同样也只输出单行结果,且 W_K 保持不变。 \underbrace{\begin{bmatrix} 0.20 & -0.10 & 0.70 \end{bmatrix}}_{\text{嵌入}} \times \underbrace{\begin{bmatrix} 0.13 & -0.51 & -0.63 \\ -0.79 & 0.54 & -0.04 \\ 0.33 & -0.10 & -0.87 \end{bmatrix}}_{W_K} = \underbrace{\begin{bmatrix} 0.34 & -0.23 & -0.74 \end{bmatrix}}_{K(\text{新})} 接着,我们将这新的一行追加(append)到上一轮迭代所缓存的 4 行 K 之后: \underbrace{\begin{bmatrix} -0.21 & 0.13 & 0.29 \\ 0.30 & -0.68 & -0.03 \\ -0.39 & 0.36 & -0.56 \\ 0.96 & -0.74 & -1.15 \end{bmatrix}}_{\text{缓存的 } K} \xrightarrow{\text{追加}} \underbrace{\begin{bmatrix} 0.34 & -0.23 & -0.74 \end{bmatrix}}_{K(\text{新})} = \underbrace{\begin{bmatrix} -0.21 & 0.13 & 0.29 \\ 0.30 & -0.68 & -0.03 \\ -0.39 & 0.36 & -0.56 \\ 0.96 & -0.74 & -1.15 \\ 0.34 & -0.23 & -0.74 \end{bmatrix}}_{K} 这样,我们就得到了 Prompt 中所有 token 的 K 矩阵,但我们只需要计算它的最后一行。 我们继续按照这种方式计算,得到新的得分(scores): \underbrace{\begin{bmatrix} -0.49 & -0.01 & 0.48 \end{bmatrix}}_{Q(\text{新})} \times \underbrace{\begin{bmatrix} -0.21 & 0.30 & -0.39 & 0.96 & 0.34 \\ 0.13 & -0.68 & 0.36 & -0.74 & -0.23 \\ 0.29 & -0.03 & -0.56 & -1.15 & -0.74 \end{bmatrix}}_{\text{转置}(K)} = \underbrace{\begin{bmatrix} 0.24 & -0.16 & -0.08 & -1.01 & -0.52 \end{bmatrix}}_{\text{得分}(\text{新})} 以及新的权重(weights): \text{softmax}\left( \underbrace{\begin{bmatrix} 0.24 & -0.16 & -0.08 & -1.01 & -0.52 \end{bmatrix}}_{\text{得分}(\text{新})} \right) = \underbrace{\begin{bmatrix} 0.32 & 0.21 & 0.23 & 0.09 & 0.15 \end{bmatrix}}_{\text{权重}(\text{新})} 在整个过程中,我们只计算必要的数值,完全不需要重新计算旧值。接下来继续计算 V 的新一行: \underbrace{\begin{bmatrix} 0.20 & -0.10 & 0.70 \end{bmatrix}}_{\text{嵌入}} \times \underbrace{\begin{bmatrix} -0.74 & 0.91 & -0.21 \\ -0.51 & 0.43 & 0.18 \\ -0.56 & -0.21 & 0.76 \end{bmatrix}}_{W_V} = \underbrace{\begin{bmatrix} -0.49 & -0.01 & 0.48 \end{bmatrix}}_{V(\text{新})} 并将其追加到此前缓存的 V 中: \underbrace{\begin{bmatrix} -0.21 & 0.13 & 0.29 \\ 0.30 & -0.68 & -0.03 \\ -0.39 & 0.36 & -0.56 \\ 0.96 & -0.74 & -1.15 \end{bmatrix}}_{\text{缓存的 } V} \xrightarrow{\text{追加}} \underbrace{\begin{bmatrix} -0.49 & -0.01 & 0.48 \end{bmatrix}}_{V(\text{新})} = \underbrace{\begin{bmatrix} -0.21 & 0.13 & 0.29 \\ 0.30 & -0.68 & -0.03 \\ -0.39 & 0.36 & -0.56 \\ 0.96 & -0.74 & -1.15 \\ -0.49 & -0.01 & 0.48 \end{bmatrix}}_{V} 最后,将新的权重与更新后的 V 相乘,得到最终的新嵌入: \underbrace{\begin{bmatrix} 0.32 & 0.21 & 0.23 & 0.09 & 0.15 \end{bmatrix}}_{\text{权重}(\text{新})} \times \underbrace{\begin{bmatrix} -0.21 & 0.13 & 0.29 \\ 0.30 & -0.68 & -0.03 \\ -0.39 & 0.36 & -0.56 \\ 0.96 & -0.74 & -1.15 \\ -0.49 & -0.01 & 0.48 \end{bmatrix}}_{V} = \underbrace{\begin{bmatrix} -0.08 & -0.09 & -0.08 \end{bmatrix}}_{\text{嵌入}(\text{新})} 这单行新的嵌入就是我们所需的全部结果。得益于缓存的 K 和 V ,先前所有 token 的上下文信息都已经融入其中。 被缓存的数据是 \text{嵌入} \times W_K 和 \text{嵌入} \times W_V 的计算结果,即 K 和 V 。因此,提示词缓存(Prompt caching)通常被称为“KV 缓存”(KV caching)。 缓存内容: K = \begin{bmatrix} -0.21 & 0.13 & 0.29 \\ 0.30 & -0.68 & -0.03 \\ -0.39 & 0.36 & -0.56 \\ 0.96 & -0.74 & -1.15 \\ 0.34 & -0.23 & -0.74 \end{bmatrix} V = \begin{bmatrix} -0.21 & 0.13 & 0.29 \\ 0.30 & -0.68 & -0.03 \\ -0.39 & 0.36 & -0.56 \\ 0.96 & -0.74 & -1.15 \\ -0.49 & -0.01 & 0.48 \end{bmatrix} 参考自 Prompt caching: 10x cheaper LLM tokens, but how? | ngrok blog 1 个帖子 - 1 位参与者 阅读完整话题
- 情报分类:技术学习与提效
- 分类依据:详解Transformer推理为何只缓存KV而非QKV,属技术学习内容。
- 信息来源:服务器 / LINUX DO - 最新话题
- 发布时间:2026/9/24 22:26:48
- 暂无回复