大模型入门:面试必会 Multi-Head Attention,从 QKV、Mask 到 KV Cache

在这里插入图片描述

写在前面:会背公式,不等于真的懂 Attention

大模型面试里,Multi-Head Attention 几乎是绕不开的题。

面试官不一定会让你推完整篇 Transformer 论文,但很可能会连续追问:

  • Q、K、V 分别是什么,为什么要拆成三个矩阵?
  • Multi-Head 到底多在哪里,是参数变多还是维度切分?
  • Mask 是加在哪里,为什么不是 softmax 之后再置零?
  • KV Cache 为什么只缓存 K 和 V,不缓存 Q?
  • 训练阶段和推理阶段,Attention 的计算方式有什么不同?
  • MHA、MQA、GQA 的差别为什么都围绕 KV Cache 展开?

很多人第一问还能答,第二问开始背公式,第三问就容易乱。

原因也很直接:只看公式时,Attention 像一个抽象数学模块;一旦手写代码,你必须面对具体问题:

  • 输入张量到底是 [batch, seq_len, hidden] 还是 [seq_len, batch, hidden]
  • num_heads 拆完以后,head 这一维放在哪里?
  • Q @ K.transpose(-2, -1) 得到的矩阵形状是什么?
  • 上三角 Mask 的 True 到底表示保留还是屏蔽?
  • 推理时历史 token 的 K/V 拼到哪个维度?

所以这篇不从“公式解释公式”开始,而是从一个最小可运行的实现开始。

一句话理解:Multi-Head Attention 就是把每个 token 的隐藏向量投影成 Q/K/V,再按多个 head 分组计算“当前 token 应该看哪些 token”,最后把多个 head 的结果拼回 hidden 维度。

本文路线

在这里插入图片描述

我们按这条线走:

  1. 先看 Q/K/V 到底是什么。
  2. 再看单头 Attention 的矩阵乘法。
  3. 然后把 hidden 维度切成多个 head。
  4. 加上 causal mask,防止偷看未来 token。
  5. 最后加 KV Cache,理解推理为什么能少算很多历史 K/V。

这条路线比直接背公式更适合面试,因为它能把概念、代码和工程意义连起来。

1. Q/K/V:不是三个神秘概念,而是三个线性投影

假设输入是:

x.shape == [batch_size, seq_len, hidden_dim]

这里每个 token 已经变成一个 hidden_dim 维向量。

Attention 要解决的问题是:

对于当前位置的 token,我应该从前面哪些 token 里拿信息,每个 token 拿多少?

Q/K/V 可以用一个很朴素的方式理解:

名称 直觉 在计算里做什么
Query 当前 token 发出的查询 拿它去和所有 Key 做相似度
Key 每个 token 对外暴露的索引 被 Query 匹配,用来算注意力分数
Value 每个 token 真正提供的信息 按注意力权重加权求和

这三个向量都来自输入 x,只是用了三组不同的线性层:

Q = self.q_proj(x)
K = self.k_proj(x)
V = self.v_proj(x)

如果 hidden_dim = 768,那么投影前后通常仍是:

Q.shape == K.shape == V.shape == [batch_size, seq_len, 768]

为什么不直接用 x @ x.T 算相似度?

因为模型需要学会“用什么角度匹配”和“真正取什么信息”。Q/K 决定匹配关系,V 决定被聚合的信息。如果三者全混在一起,表达能力会弱很多。

2. 单头 Attention:核心只有三步

先不考虑多头,单头 Scaled Dot-Product Attention 的核心是:

scores  = Q @ K^T / sqrt(head_dim)
weights = softmax(scores)
output  = weights @ V

形状变化是理解 Attention 的关键。

假设:

Q.shape == [batch, seq_len, head_dim]
K.shape == [batch, seq_len, head_dim]
V.shape == [batch, seq_len, head_dim]

那么:

scores = Q @ K.transpose(-2, -1)
scores.shape == [batch, seq_len, seq_len]

这个 seq_len x seq_len 矩阵是什么意思?

i 行表示:第 i 个 token 看所有 token 的分数。

如果是 Decoder-only 大模型,在生成任务里,第 i 个 token 不能看第 i+1 个 token,因为那是未来信息。这就是后面 causal mask 要解决的问题。

为什么要除以 sqrt(head_dim)

如果向量维度很大,点积值会变大,softmax 容易变得特别尖,梯度也会变得不稳定。Transformer 原论文和 PyTorch 文档里的标准写法都会做这个缩放。

3. Multi-Head:不是复制模型,而是切分 hidden 维度

Multi-Head Attention 容易被误解成“并行跑多个 Attention 模型”。

更准确地说,它通常是把 hidden 维度切成多个 head,每个 head 在自己的子空间里算注意力。

比如:

hidden_dim = 768
num_heads = 12
head_dim = hidden_dim // num_heads  # 64

投影后先得到:

Q.shape == [batch, seq_len, hidden_dim]

然后 reshape 成:

Q.shape == [batch, num_heads, seq_len, head_dim]

这样每个 head 看到的是 64 维子空间,而不是完整 768 维。

在这里插入图片描述

PyTorch 里最容易写错的地方就在这里:

q = self.q_proj(x)
q = q.view(batch_size, seq_len, self.num_heads, self.head_dim)
q = q.transpose(1, 2)

transpose(1, 2) 之后,形状才是:

[batch_size, num_heads, seq_len, head_dim]

为什么要把 num_heads 放到前面?

因为后面要让每个 head 独立计算:

scores = q @ k.transpose(-2, -1)

这时结果就是:

scores.shape == [batch_size, num_heads, seq_len, seq_len]

也就是说,每个 batch、每个 head 都有一张自己的注意力图。

4. 从零手写一个最小 Multi-Head Attention

下面这份代码保留了面试最核心的部分:

  • Q/K/V 三个投影;
  • hidden 维度切成多个 head;
  • causal mask;
  • softmax;
  • 多头结果拼接;
  • 输出投影。
import math
import torch
from torch import nn


class MultiHeadAttention(nn.Module):
    def __init__(self, hidden_dim: int, num_heads: int, dropout: float = 0.0):
        super().__init__()
        assert hidden_dim % num_heads == 0

        self.hidden_dim = hidden_dim
        self.num_heads = num_heads
        self.head_dim = hidden_dim // num_heads

        self.q_proj = nn.Linear(hidden_dim, hidden_dim)
        self.k_proj = nn.Linear(hidden_dim, hidden_dim)
        self.v_proj = nn.Linear(hidden_dim, hidden_dim)
        self.o_proj = nn.Linear(hidden_dim, hidden_dim)
        self.dropout = nn.Dropout(dropout)

    def _split_heads(self, x: torch.Tensor) -> torch.Tensor:
        batch_size, seq_len, _ = x.shape
        x = x.view(batch_size, seq_len, self.num_heads, self.head_dim)
        return x.transpose(1, 2)  # [B, H, T, D]

    def _merge_heads(self, x: torch.Tensor) -> torch.Tensor:
        batch_size, num_heads, seq_len, head_dim = x.shape
        x = x.transpose(1, 2).contiguous()
        return x.view(batch_size, seq_len, num_heads * head_dim)

    def forward(self, x: torch.Tensor, attn_mask: torch.Tensor | None = None):
        q = self._split_heads(self.q_proj(x))
        k = self._split_heads(self.k_proj(x))
        v = self._split_heads(self.v_proj(x))

        scores = q @ k.transpose(-2, -1)
        scores = scores / math.sqrt(self.head_dim)

        if attn_mask is not None:
            scores = scores.masked_fill(attn_mask, float("-inf"))

        attn_weights = torch.softmax(scores, dim=-1)
        attn_weights = self.dropout(attn_weights)

        context = attn_weights @ v
        output = self._merge_heads(context)
        output = self.o_proj(output)
        return output, attn_weights

可以用一个简单输入验证形状:

x = torch.randn(2, 5, 768)
mha = MultiHeadAttention(hidden_dim=768, num_heads=12)

output, attn_weights = mha(x)

print(output.shape)        # torch.Size([2, 5, 768])
print(attn_weights.shape)  # torch.Size([2, 12, 5, 5])

这两个 shape 很重要。

output 回到了输入的 hidden 维度,说明它可以继续接残差、LayerNorm、FFN。

attn_weights[batch, heads, target_len, source_len],说明每个 head 都有自己的注意力分布。

5. Causal Mask:为什么要在 softmax 前加负无穷

训练 Decoder-only 大模型时,我们通常一次性把整段序列喂进去并行计算。

比如句子是:

我 / 喜欢 / 写 / 代码

训练时第 1 个 token 不能看到第 2、3、4 个 token;第 2 个 token 不能看到第 3、4 个 token。

所以需要一个上三角 mask:

def build_causal_mask(seq_len: int, device=None):
    return torch.triu(
        torch.ones(seq_len, seq_len, dtype=torch.bool, device=device),
        diagonal=1
    )

seq_len = 5,mask 大概是:

0 1 1 1 1
0 0 1 1 1
0 0 0 1 1
0 0 0 0 1
0 0 0 0 0

这里 1 / True 表示“不能看”。

实际传给前面的 MHA 时,要扩展成能广播到 [B, H, T, T] 的形状:

mask = build_causal_mask(seq_len=5, device=x.device)
mask = mask[None, None, :, :]  # [1, 1, T, T]
output, weights = mha(x, attn_mask=mask)

为什么不是 softmax 之后把未来 token 权重置零?

因为 softmax 会让一行概率和为 1。如果先 softmax 再置零,剩下的概率和就不再是 1,还要重新归一化。

标准做法是:先把未来位置的分数变成 -inf,再 softmax。

scores = scores.masked_fill(mask, float("-inf"))
weights = torch.softmax(scores, dim=-1)

这样未来位置经过 softmax 后天然变成 0,合法位置的概率仍然会重新归一化。

在这里插入图片描述

面试里这点很容易被追问:

Mask 是加在 attention score 上,不是加在 V 上,也不是 softmax 后随便置零。

6. KV Cache:推理阶段为什么不重复算历史 K/V

前面的 MHA 是训练视角:一次性处理整个序列。

但大模型推理是自回归的,通常一个 token 一个 token 生成。

假设已经生成了:

我 / 喜欢 / 写

下一步要预测“代码”。如果每一步都把整段历史重新送进模型,就会重复计算前面 token 的 K 和 V。

KV Cache 的思路很简单:

历史 token 的 K/V 在当前层里已经算过了,下一步生成时直接复用;新 token 只需要算自己的 Q/K/V,然后用新的 Q 去看历史 K/V。

Hugging Face 文档里也强调,KV cache 是推理阶段的优化:缓存过去 token 的 key/value,后续 token 直接复用,避免重复计算。

注意这里说的是“每一层”的 K/V。Transformer 有多少层,就有多少层自己的 KV Cache。

在这里插入图片描述

为什么不缓存 Q?

因为 Q 表示“当前 token 想查什么”。

生成第 t 个 token 时,真正需要拿去和历史 K 做匹配的是当前 token 的 Q。历史 token 的 Q 对预测下一个 token 没有用了。

但历史 token 的 K/V 还会继续被后面的 token 查询,所以值得缓存。

一句话:

Q 是一次性的查询,K/V 是可复用的索引和内容。

7. 手写带 KV Cache 的 MHA

下面在前面的版本上加 past_key_value

约定:

past_key_value = (past_k, past_v)

past_k.shape == [batch, heads, past_len, head_dim]
past_v.shape == [batch, heads, past_len, head_dim]

代码如下:

class MultiHeadAttentionWithKVCache(MultiHeadAttention):
    def forward(
        self,
        x: torch.Tensor,
        attn_mask: torch.Tensor | None = None,
        past_key_value: tuple[torch.Tensor, torch.Tensor] | None = None,
        use_cache: bool = False,
    ):
        q = self._split_heads(self.q_proj(x))
        k = self._split_heads(self.k_proj(x))
        v = self._split_heads(self.v_proj(x))

        if past_key_value is not None:
            past_k, past_v = past_key_value
            k = torch.cat([past_k, k], dim=2)
            v = torch.cat([past_v, v], dim=2)

        present_key_value = (k, v) if use_cache else None

        scores = q @ k.transpose(-2, -1)
        scores = scores / math.sqrt(self.head_dim)

        if attn_mask is not None:
            scores = scores.masked_fill(attn_mask, float("-inf"))

        attn_weights = torch.softmax(scores, dim=-1)
        attn_weights = self.dropout(attn_weights)

        context = attn_weights @ v
        output = self._merge_heads(context)
        output = self.o_proj(output)

        return output, attn_weights, present_key_value

最关键的一行是:

k = torch.cat([past_k, k], dim=2)
v = torch.cat([past_v, v], dim=2)

为什么是 dim=2

因为我们的 K/V 形状是:

[batch, heads, seq_len, head_dim]

第 2 维才是序列长度维度。

一个简化的推理过程可以这样写:

mha = MultiHeadAttentionWithKVCache(768, 12)
past = None

for step in range(10):
    # 每次只输入当前 token 的 hidden state
    x_t = torch.randn(2, 1, 768)

    out, weights, past = mha(
        x_t,
        past_key_value=past,
        use_cache=True,
    )

    print(past[0].shape)  # [2, 12, step + 1, 64]

每生成一步,cache 的 seq_len 增加 1。

8. 训练、Prefill、Decode:三个阶段不要混

很多 KV Cache 问题答错,是因为把训练和推理混在一起。

可以这样区分:

阶段 输入 是否并行 是否用 KV Cache 重点
训练 完整序列 用 causal mask 防止看未来
Prefill prompt 完整序列 写入 cache 一次性算出 prompt 的 K/V
Decode 当前新 token 否,逐 token 读写 cache 新 Q 查询历史 K/V

训练阶段通常不需要 KV Cache,因为整段序列并行计算更高效。

推理阶段分两段:

  1. Prefill:用户输入的 prompt 一次性过模型,得到第一批 KV Cache。
  2. Decode:后面每生成一个 token,只处理当前 token,并复用历史 K/V。

Hugging Face 的 KV Cache 文档也提醒,缓存主要用于推理;如果训练时打开缓存,可能引入不符合预期的问题。

9. KV Cache 不是免费午餐:它省计算,但吃显存

KV Cache 是典型的空间换时间。

它减少了重复计算,但会随着这些因素线性增长:

  • batch size;
  • 序列长度;
  • 层数;
  • KV head 数量;
  • head_dim;
  • 数据类型字节数。

一个常见估算公式是:

KV Cache bytes
= batch_size
  * seq_len
  * num_layers
  * num_kv_heads
  * head_dim
  * 2
  * bytes_per_element

这里的 2 表示 K 和 V 两份缓存。

所以后续 MQA、GQA、MLA 这些注意力变体,很多都在围绕一个问题做优化:

如何减少 KV Cache 的存储和读取压力,同时尽量保住模型效果?

这也是为什么面试官问完 MHA,经常继续问 GQA。

因为 GQA 的核心不是“名字更高级”,而是减少 KV head 数量,让多个 Q head 共享一组 K/V head,从而降低 KV Cache 成本。

10. 面试时怎么讲清楚

如果面试官让你“手写 MHA”,不要一上来背完整 Transformer。

可以按这个顺序答:

输入是 [B, T, C]。先通过三个线性层得到 Q、K、V,然后 reshape 成 [B, H, T, D],其中 D=C/H。每个 head 内部计算 QK^T / sqrt(D) 得到 [B, H, T, T] 的注意力分数。如果是 Decoder 自回归任务,需要在 softmax 前加 causal mask,把未来 token 的 score 置为 -inf。softmax 后乘以 V 得到每个 head 的 context,再把 head 维度拼回 [B, T, C],最后过输出投影。

如果继续问 KV Cache,可以接:

训练时通常不需要 KV Cache,因为完整序列可以并行算;推理时是逐 token decode,历史 token 的 K/V 在每一层已经算过,后续 token 会反复用到,所以缓存 K/V。Q 是当前 token 的查询,每步都不同,历史 Q 对预测下一个 token 没有复用价值,所以一般不缓存 Q。KV Cache 降低重复计算,但显存会随层数、序列长度、KV head 数线性增长。

这段回答基本能覆盖 QKV、Mask、KV Cache 和工程权衡。

11. 常见坑

在这里插入图片描述

坑 1:head_dim 算错

hidden_dim 必须能被 num_heads 整除:

assert hidden_dim % num_heads == 0

否则 reshape 会直接错。

坑 2:transpose 后忘了 contiguous

transpose 会改变张量视图,后面如果直接 view,可能报错或行为不符合预期。

x = x.transpose(1, 2).contiguous()

坑 3:mask 的含义写反

有些实现里 True 表示屏蔽,有些接口里 1 表示保留。自己手写时要统一约定。

本文代码里:

True = 不能看
False = 可以看

坑 4:softmax 维度写错

Attention 权重应该沿着 source token 维度归一化:

torch.softmax(scores, dim=-1)

坑 5:KV Cache 拼接维度错

K/V 形状是:

[batch, heads, seq_len, head_dim]

所以拼历史 cache 时是:

torch.cat([past_k, k], dim=2)

坑 6:把训练和推理混在一起

训练靠 causal mask 并行计算。

推理靠 KV Cache 复用历史 K/V。

这两件事都和“不能看未来 token”有关,但不是同一个问题。

12. 一张面试速记表

在这里插入图片描述

问题 关键回答
Q/K/V 是什么? 输入 hidden state 的三组线性投影,Q 负责查询,K 负责匹配,V 负责提供内容
为什么要多头? 把 hidden 维度拆到多个子空间,每个 head 学不同关系
attention score 形状? [B, H, T_q, T_k]
mask 加在哪里? softmax 前,加在 score 上
causal mask 做什么? 防止当前位置看到未来 token
为什么缩放? 降低大维度点积导致的 softmax 饱和风险
KV Cache 存什么? 每一层历史 token 的 K 和 V
为什么不存 Q? Q 是当前 token 的查询,历史 Q 后续没有复用价值
KV Cache 代价? 省计算但占显存,随长度、层数、KV head 数线性增长
GQA 解决什么? 让多个 Q head 共享较少的 K/V head,降低 KV Cache 成本

总结

Multi-Head Attention 不应该只靠背公式。

要抓住的是四件事:

  1. Q/K/V 是从输入 hidden state 投影出来的三种视角。
  2. Multi-Head 是把 hidden 维度切成多个子空间并行算注意力。
  3. Causal Mask 是在 softmax 前屏蔽未来 token。
  4. KV Cache 是推理阶段复用历史 K/V 的工程优化。

参考资料

  • Vaswani et al.:Attention Is All You Need
    https://arxiv.org/abs/1706.03762
  • PyTorch:torch.nn.MultiheadAttention 文档
    https://docs.pytorch.org/docs/2.12/generated/torch.nn.MultiheadAttention.html
  • Hugging Face Transformers:Caching explained
    https://huggingface.co/docs/transformers/main/cache_explanation
  • Hugging Face Transformers:Cache strategies
    https://huggingface.co/docs/transformers/main/kv_cache
Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐