🎯 课程主题
Attention 函数的代码实现——缩放点积注意力的 PyTorch 实现。
📝 核心知识点
1. Attention 函数的输入输出
- 概念说明:实现注意力机制的核心计算。
- 关键细节:
- 输入:$Q, K, V$ 矩阵 + mask + dropout
- 输出:信息融合结果 + 注意力得分矩阵
- 返回两个值
2. 计算流程
- 概念说明:按公式顺序执行 QKV 计算。
- 关键细节:
- 获取 $d_k$ 维度(Q 的最后一维)
- $Q \times K^T$ → 注意力得分
- 除以 $\sqrt{d_k}$(缩放)
- 如需 mask,填充负无穷(用大负数代替)
- Softmax 归一化
- 可选 Dropout
- 注意力得分 × $V$ → 输出
3. Mask 的处理
- 概念说明:通过填充负无穷实现掩码。
- 关键细节:
- 代码中用很大的负数代替 $-\infty$
- 填充掩码和因果掩码都通过此机制实现
- 经 Softmax 后负无穷变为 0
- 通过参数控制是否需要 mask
4. 函数的通用性
- 概念说明:此函数适用于三种注意力机制。
- 关键细节:
- 自注意力:Q=K=V 来自同一输入
- 因果掩码注意力:Q=K=V + 因果 mask
- 交叉注意力:Q 来自解码器,K/V 来自编码器
- 核心计算逻辑完全相同
🧮 核心公式与推导
Attention 函数公式:
$$\text{Attention}(Q, K, V) = \text{Softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right) V$$
带 Mask 的版本:
$$\text{Attention}(Q, K, V, M) = \text{Softmax}\left(\frac{QK^T}{\sqrt{d_k}} + M\right) V$$
其中 $M$ 中屏蔽位置为 $-\infty$,有效位置为 0。
维度变化:
- $Q \in \mathbb{R}^{L_q \times d_k}$
- $K \in \mathbb{R}^{L_k \times d_k}$
- $V \in \mathbb{R}^{L_v \times d_v}$(通常 $L_k = L_v$)
- $QK^T \in \mathbb{R}^{L_q \times L_k}$
- 输出 $\in \mathbb{R}^{L_q \times d_v}$
🏗️ 模型架构与数据流向
💻 代码实战
import torch
import torch.nn as nn
import torch.nn.functional as F
import math
def attention(Q, K, V, mask=None, dropout=None):
"""
缩放点积注意力函数
参数:
Q: [batch, n_heads, L_q, d_k] 查询矩阵
K: [batch, n_heads, L_k, d_k] 键矩阵
V: [batch, n_heads, L_v, d_v] 值矩阵 (通常 L_k = L_v)
mask: [batch, 1, L_q, L_k] 或 None, 屏蔽位置为0
dropout: dropout层或None
返回:
output: [batch, n_heads, L_q, d_v] 注意力输出
attn_weights: [batch, n_heads, L_q, L_k] 注意力得分矩阵
"""
d_k = Q.size(-1)
# 1. 计算注意力得分: Q @ K^T / sqrt(d_k)
# [batch, heads, L_q, d_k] @ [batch, heads, d_k, L_k] → [batch, heads, L_q, L_k]
scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(d_k)
# 2. 应用Mask (填充掩码或因果掩码)
if mask is not None:
# mask为0的位置填充负无穷(用-1e9代替, 避免数值问题)
scores = scores.masked_fill(mask == 0, -1e9)
# 3. Softmax归一化 (最后一维, 即对每个查询位置归一化)
attn_weights = F.softmax(scores, dim=-1)
# 4. 可选Dropout
if dropout is not None:
attn_weights = dropout(attn_weights)
# 5. 信息融合: 注意力得分 @ V
# [batch, heads, L_q, L_k] @ [batch, heads, L_v, d_v] → [batch, heads, L_q, d_v]
output = torch.matmul(attn_weights, V)
return output, attn_weights
# 测试示例
if __name__ == "__main__":
batch_size = 2
n_heads = 8
L = 5 # 序列长度
d_k = 64
# 随机生成Q, K, V
Q = torch.randn(batch_size, n_heads, L, d_k)
K = torch.randn(batch_size, n_heads, L, d_k)
V = torch.randn(batch_size, n_heads, L, d_k)
# 1. 无Mask的注意力
output, attn = attention(Q, K, V)
print(f"输出形状: {output.shape}") # [2, 8, 5, 64]
print(f"注意力得分形状: {attn.shape}") # [2, 8, 5, 5]
print(f"注意力得分行和: {attn[0, 0, 0].sum()}") # ≈1.0
# 2. 带因果Mask的注意力 (下三角可见)
causal_mask = torch.tril(torch.ones(L, L)).expand(
batch_size, n_heads, L, L
) # [batch, heads, L, L] 下三角为1
output_masked, attn_masked = attention(Q, K, V, mask=causal_mask)
print(f"因果掩码后注意力:\n{attn_masked[0, 0]}")
# 上三角为0, 下三角有值
# 3. 带填充Mask的注意力
padding_mask = torch.tensor([
[1, 1, 1, 1, 1], # 样本1: 全有效
[1, 1, 1, 0, 0] # 样本2: 后2位是padding
]).unsqueeze(1).unsqueeze(2).expand(batch_size, n_heads, L, L)
output_pad, attn_pad = attention(Q, K, V, mask=padding_mask)
⚠️ 常见问题与避坑指南
- 代码中用
-1e9代替 $-\infty$,避免数值计算问题 - K 需转置(
K.transpose(-2, -1))才能与 Q 相乘 - Softmax 在最后一维(
dim=-1)做归一化 - mask 为 0 的位置被屏蔽,非 0 位置保留
- 此函数是通用的,适用于自注意力、因果掩码注意力、交叉注意力
💡 个人总结与延伸
Attention 函数是 Transformer 的核心,实现了 $\text{Softmax}(\frac{QK^T}{\sqrt{d_k}})V$ 的计算。代码实现简洁但功能强大,通过 mask 参数支持多种掩码场景。现代大模型中,Flash Attention 等优化技术通过融合计算减少内存访问,大幅提升注意力计算效率。