直接调用torch.nn.multiheadattention不等于手动实现,因其内部封装了投影、mask处理、dropout、多头拼接等逻辑,参数不可见、前向过程黑盒化;而手写需从q/k/v线性投影起步,严格按线性投影→缩放点积→(可选)mask→加权求和四步可控执行,才能真正理解计算流、梯度流向及缩放因子作用。

为什么直接用 torch.nn.MultiheadAttention 不等于“手动实现”
很多人以为调用 torch.nn.MultiheadAttention 就是用了注意力机制,但它的内部封装了投影、mask处理、dropout、多头拼接等逻辑,参数不可见,前向过程黑盒化。真要理解 Self-Attention 的计算流、梯度流向、缩放因子作用,必须从 Q、K、V 三张权重矩阵开始手写。
关键点在于:不是“能不能跑”,而是“每一步是否可控”。比如你无法在 MultiheadAttention 中单独修改 softmax 前的缩放系数,也无法插入自定义的稀疏 mask 或替换 softmax 为 softmax - dropout 的变体。
手写 SelfAttention 层的四个核心步骤
一个标准的单头 Self-Attention(不含多头、不带 LayerNorm)只需四步,顺序不能错:
- 线性投影:用三个独立的
nn.Linear分别将输入x映射为Q、K、V,维度必须一致(如embed_dim=512) - 缩放点积:计算
Q @ K.T / sqrt(d_k),其中d_k是每个 head 的维度(单头时 =embed_dim),这步防止 softmax 输入过大导致梯度消失 - 可选 mask:若做 causal attention(如 GPT),需用
torch.tril构造下三角 mask,并把上三角位置设为-inf;注意要用masked_fill_而非加法,否则 NaN 风险高 - 加权求和:
softmax(...)后与V相乘,输出形状与输入x一致
示例片段(无 mask 版):
class SelfAttention(nn.Module):
def __init__(self, embed_dim):
super().__init__()
self.q_proj = nn.Linear(embed_dim, embed_dim)
self.k_proj = nn.Linear(embed_dim, embed_dim)
self.v_proj = nn.Linear(embed_dim, embed_dim)
self.scale = embed_dim ** -0.5
<pre class="brush:python;toolbar:false;">def forward(self, x): # x: [batch, seq_len, embed_dim]
Q = self.q_proj(x) # [b, s, d]
K = self.k_proj(x) # [b, s, d]
V = self.v_proj(x) # [b, s, d]
attn = Q @ K.transpose(-2, -1) * self.scale # [b, s, s]
attn = F.softmax(attn, dim=-1)
return attn @ V # [b, s, d]
forward 中最容易崩的三个地方
手写时出错往往不是逻辑错,而是 shape 或 dtype 不匹配:
-
Q @ K.T前没确保最后两维可乘:如果Q是[b, s, d],K.T必须是[b, d, s],用K.transpose(-2, -1)安全,别用K.permute(0, 2, 1)硬写——容易在 batch=1 时维度错位 - softmax 前填
-inf时用了float('-inf')但输入是half(FP16):会变成nan,必须用torch.finfo(Q.dtype).min - 没处理
seq_len变长情况:训练时pad过的序列,若没传key_padding_mask,padding 位置也会参与 attention 计算——必须显式 mask,不能依赖模型自动忽略
扩展到多头时,nn.Linear 的 out_features 怎么设
常见错误是把 embed_dim 直接当每个 head 的维度。正确做法是:先定 head 数 num_heads,再算单头维度 head_dim = embed_dim // num_heads,然后所有投影层的 out_features 设为 embed_dim(不是 head_dim)——因为你要一次性投出全部 heads 的 Q/K/V,再用 view 拆分。
例如 embed_dim=768、num_heads=12,则 head_dim=64;q_proj = nn.Linear(768, 768),之后 reshape 为 [b, s, 12, 64] 再 transpose 成 [b, 12, s, 64]。
漏掉这个 reshape + transpose 步骤,或者 transpose 维度写错(比如把 1 和 2 写反),就会导致 attention score 形状错乱,后续 @ 报 RuntimeError: mat1 and mat2 shapes cannot be multiplied。
真正难的从来不是公式,是 tensor 的 layout 和 view 的时机——PyTorch 不报错,只悄悄给你错的结果。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











