注意力机制 - 代码实现
1 课程概览
本课讲解注意力机制的代码实现。通过 Seq2Seq 架构完成翻译任务,自定义 MyAttention 类,实现 Q、K、V 的注意力计算。初始化参数包括 query_size、key_size、value_size、output_size,使用线性层和 bmm() 函数。
2 核心概念与定义
- MyAttention 类:自定义注意力机制模块。
- query_size:查询张量维度。
- key_size:键张量维度。
- value_size:值张量维度。
- output_size:输出维度。
- bmm():批量矩阵乘法。
3 模型与算法详解
任务描述
Seq2Seq 翻译任务。
- V 是内容(32 个单词,每个 4 个特征)
- K 是 32 个单词的索引
- Q 是查询张量
- 任务:计算 Q 的注意力权重分布和结果表示
MyAttention 类
自定义注意力机制模块。
class MyAttention(nn.Module):
def __init__(self, query_size, key_size, value_size, output_size):
...
def forward(self, Q, K, V):
...
实现思路
1. __init__: 初始化线性层
2. forward:
- 计算 Q 的注意力权重分布
- 用 bmm() 计算 Q 的注意力结果表示
- 按指定维度输出
- 返回注意力结果表示
注意力计算流程
1. Q · K^T → 匹配分
2. softmax → 概率分布
3. 概率分布 · V → attention_q
4. 线性层 → output
4 数学原理与推导
注意力计算
$$\text{attn_weights} = \text{softmax}(\text{Q} \cdot \text{K}^T)$$
$$\text{attention_q} = \text{bmm}(\text{attn_weights}, \text{V})$$
$$\text{output} = \text{Linear}(\text{attention_q})$$
其中:
- $\text{Q}$ 是查询张量
- $\text{K}$ 是键张量
- $\text{V}$ 是值张量
- $\text{bmm}$ 是批量矩阵乘法
5 代码示例
import torch
import torch.nn as nn
import torch.nn.functional as F
class MyAttention(nn.Module):
"""自定义注意力机制模块"""
def __init__(self, query_size, key_size, value_size, output_size):
"""
初始化函数
Args:
query_size: 查询张量维度
key_size: 键张量维度
value_size: 值张量维度
output_size: 输出维度
"""
super(MyAttention, self).__init__()
# 1. 线性层
self.query_layer = nn.Linear(query_size, key_size)
# 2. 注意力权重分布
self.attn_weights = None
# 3. 输出层
self.output_layer = nn.Linear(value_size, output_size)
def forward(self, Q, K, V):
"""
前向传播
Args:
Q: 查询张量 [batch_size, 1, query_size]
K: 键张量 [batch_size, seq_len, key_size]
V: 值张量 [batch_size, seq_len, value_size]
Returns:
output: 注意力结果表示 [batch_size, 1, output_size]
attn_weights: 注意力权重分布 [batch_size, 1, seq_len]
"""
# 1. 计算查询张量 Q 的注意力权重分布
# Q: [batch_size, 1, query_size] -> [batch_size, 1, key_size]
Q = self.query_layer(Q)
# attn_weights: [batch_size, 1, seq_len]
attn_weights = F.softmax(torch.matmul(Q, K.transpose(-2, -1)), dim=-1)
self.attn_weights = attn_weights
# 2. 用 bmm() 计算查询张量 Q 的注意力结果表示
# attention_q: [batch_size, 1, value_size]
attention_q = torch.bmm(attn_weights, V)
# 3. 按指定维度输出
# output: [batch_size, 1, output_size]
output = self.output_layer(attention_q)
# 4. 返回注意力结果表示和权重分布
return output, attn_weights
# 测试
if __name__ == "__main__":
print("=== 注意力机制 - 代码实现 ===")
# 1. 参数设置
batch_size = 1
seq_len = 32 # 32 个单词
query_size = 4 # 查询张量维度
key_size = 4 # 键张量维度
value_size = 4 # 值张量维度
output_size = 64 # 输出维度
# 2. 创建 Q、K、V
Q = torch.randn(batch_size, 1, query_size) # [1, 1, 4]
K = torch.randn(batch_size, seq_len, key_size) # [1, 32, 4]
V = torch.randn(batch_size, seq_len, value_size) # [1, 32, 4]
print(f"Q 形状: {Q.shape}")
print(f"K 形状: {K.shape}")
print(f"V 形状: {V.shape}")
# 3. 创建注意力机制模块
attention = MyAttention(query_size, key_size, value_size, output_size)
# 4. 前向传播
output, attn_weights = attention(Q, K, V)
print(f"\noutput 形状: {output.shape}") # [1, 1, 64]
print(f"attn_weights 形状: {attn_weights.shape}") # [1, 1, 32]
print(f"attn_weights: {attn_weights}")
代码说明
| 代码 | 说明 |
|---|---|
nn.Linear(query_size, key_size) | 线性层 |
torch.matmul(Q, K.transpose(-2, -1)) | Q · K^T |
F.softmax(..., dim=-1) | softmax 归一化 |
torch.bmm(attn_weights, V) | 批量矩阵乘法 |
self.output_layer(attention_q) | 输出层 |
6 重难点与易错提醒
- ❗重点:自定义 MyAttention 类,继承 nn.Module。
- ❗重点:使用 bmm() 计算注意力结果表示。
- ❗重点:forward 方法接收 Q、K、V 三个参数。
- ⚠️易错:Q 的维度需要通过线性层转换为与 K 相同的维度。
7 课堂问答精选
Q1:如何实现注意力机制?
A:自定义 MyAttention 类,继承 nn.Module,实现 __init__ 和 forward 方法。在 forward 中计算注意力权重分布和结果表示。
Q2:MyAttention 类的初始化参数有哪些?
A:query_size(查询张量维度)、key_size(键张量维度)、value_size(值张量维度)、output_size(输出维度)。
Q3:forward 方法接收哪些参数?
A:接收 Q、K、V 三个参数,分别表示查询张量、键张量、值张量。
8 本课小结
- 自定义 MyAttention 类,继承 nn.Module。
- 初始化参数:query_size、key_size、value_size、output_size。
- forward 方法接收 Q、K、V,返回注意力结果表示和权重分布。
- 使用 bmm() 计算注意力结果表示。
9 延伸思考
- 如何测试注意力机制?
- 注意力机制的参数如何解释?