跳到正文

Multi-Head Attention 维度变化

手动实现 QKV 变换、缩放点积、多头拼接及输出投影

原题:请手动实现一个多头注意力(Multi-Head Attention)模块,要求包括QKV线性变换、缩放点积注意力、多头拼接及输出投影,并说明各部分的作用和维度变化过程。

模型架构 · 字节真题

回答与解析

import math
import torch
from torch import nn

class MultiHeadAttention(nn.Module):
    def __init__(self, d_model: int, num_heads: int, dropout: float = 0.0):
        super().__init__()
        if d_model % num_heads != 0:
            raise ValueError("d_model must be divisible by num_heads")
        self.num_heads = num_heads
        self.head_dim = d_model // num_heads
        self.qkv = nn.Linear(d_model, 3 * d_model)
        self.out_proj = nn.Linear(d_model, d_model)
        self.dropout = nn.Dropout(dropout)

    def forward(self, x: torch.Tensor, mask: torch.Tensor | None = None):
        # x: [B, L, D];mask为bool,可广播到[B, H, L, L],True表示可见
        b, l, d = x.shape
        qkv = self.qkv(x).view(b, l, 3, self.num_heads, self.head_dim)
        q, k, v = qkv.permute(2, 0, 3, 1, 4).unbind(0)  # [B,H,L,Dh]

        scores = q @ k.transpose(-2, -1) / math.sqrt(self.head_dim)
        if mask is not None:
            mask = mask.to(torch.bool)
            scores = scores.masked_fill(~mask, torch.finfo(scores.dtype).min)

        weights = torch.softmax(scores, dim=-1)
        if mask is not None:
            # 避免整行均被mask时产生无意义的均匀权重
            weights = weights * mask.to(weights.dtype)
            weights = weights / weights.sum(-1, keepdim=True).clamp_min(1e-9)
        weights = self.dropout(weights)
        context = weights @ v                         # [B,H,L,Dh]
        context = context.transpose(1, 2).contiguous().view(b, l, d)
        return self.out_proj(context), weights        # [B,L,D]

维度与作用

[B,L,D] -> Q/K/V [B,H,L,Dh] -> scores [B,H,L,L] -> context [B,H,L,Dh] -> [B,L,D],其中 Dh=D/H。除以 sqrt(Dh) 是为了控制点积方差,避免softmax过早饱和。多头让不同投影子空间并行建模关系;out_proj负责跨头线性混合和恢复模型维度,它本身不增加非线性。

生产实现还要把因果mask与padding mask组合,并确认广播方向;长序列的主要瓶颈是 [L,L] 分数矩阵的计算和存储,可换用框架提供的scaled-dot-product attention或FlashAttention kernel。

口语版讲法(约4分钟)

  • 先约定输入和头维
  • 依次讲QKV、分头、注意力、拼接和输出投影
  • 解释缩放因子
  • 说明mask与全mask行
  • 指出长序列优化边界

我会先约定输入x的形状是B乘L乘D,分别表示batch、序列长度和模型维度。头数是H,每个头的维度Dh等于D除以H,所以初始化时必须检查D能被H整除。Q、K、V可以用三个线性层,也可以像代码里一样用一个输出三倍维度的融合线性层;两者在数学上等价,融合只是常见工程实现。

做完投影后,我把结果从B、L、三倍D reshape成B、L、3、H、Dh,再调整轴得到Q、K、V各自的B、H、L、Dh。接着计算Q乘K转置,分数形状是B、H、L、L,再除以根号Dh。这个缩放不是为了改变表达能力,而是控制点积的方差,避免维度较大时softmax输入幅值过大、梯度过早饱和。

如果有mask,我会先明确语义:布尔值True表示这个位置可见,并要求它能广播到B、H、L、L。因果mask负责挡住未来token,padding mask负责挡住补齐位置,实际使用时要组合。直接用负无穷mask时,如果某一整行全部不可见,softmax会产生NaN;示例代码用有限最小值,再把mask位置乘零并重新归一化,保证全mask行输出为零。生产代码也可以使用框架已经处理好这些边界的算子。

softmax后的权重乘V,得到B、H、L、Dh。然后把头维转回来,先变成B、L、H、Dh,再连续化并reshape成B、L、D。最后经过输出投影W O。这里要特别纠正一个常见说法:W O是线性层,它做的是跨头线性混合并映射回模型维度,本身不会增加非线性。模型的非线性来自MLP激活等其他组件。

多头的价值也不应该说成单头忙不过来。更准确地说,不同头拥有不同的QKV投影,可以在不同表示子空间里并行形成不同关系;但是否形成语法、位置或实体等可解释模式,需要分析而不能预设。attention权重上可以加dropout,模块外通常还有残差和归一化。

最后说复杂度。这个朴素实现会显式构造L乘L分数矩阵,计算和辅助存储在长序列下是主要瓶颈。工程上我会优先调用PyTorch的scaled dot product attention,让后端选择合适kernel,或者使用FlashAttention类实现。优化不会改变多头注意力的数学定义,重点是减少中间矩阵的HBM读写,并仍然正确处理mask和精度。

如果扩展到cross-attention,Q来自目标序列,K和V来自源序列,此时分数形状应是B、H、L q、L k,mask也要按这两个长度广播。代码评审时我会把self-attention和cross-attention测试分开,并检查非连续张量在view前是否已contiguous。这些边界往往比主公式更容易造成线上错误。

混合精度下还要关注scores的数值范围,优先使用经过验证的框架算子,并在不同dtype下检查输出和梯度容差,不能只在FP32小样本上验证。

关键一句:全mask行和因果mask、padding mask组合时如何避免NaN。

核验来源

  1. Attention Is All You Need
  2. PyTorch scaled_dot_product_attention

面试官还可能这样问

  1. 问法 1 · 场景切入

    假设输入 x 的形状为 [B,L,D],D 能被头数 H 整除。请实现 QKV 投影、拆头、缩放点积、mask、拼头和输出投影,并逐步标出形状。

  2. 问法 2 · 层层追问

    Q、K、V 怎样从 [B,L,D] 变成 [B,H,L,D/H]?……QK 转置后分数矩阵是什么形状?……为什么除以 sqrt(d_k)?……mask 与整行遮蔽怎样处理?……最后为什么需要 out projection?

  3. 问法 3 · 直球技术

    手写多头注意力模块并解释所有维度变化,重点说明缩放控制点积方差、避免 softmax 饱和,而不是笼统归因于梯度消失。

同模块相关题目