扩展 - bmm() 函数简介
1 课程概览
本课讲解 bmm() 函数。bmm() 是批量矩阵乘法函数,专门用于处理序列数据(如 RNN、Transformer),效率比 matmul() 高。bmm() 不支持广播机制。
2 核心概念与定义
- bmm():批量矩阵乘法函数(Batch Matrix Multiplication)。
- matmul():通用矩阵乘法函数(Matrix Multiplication)。
- 广播机制:matmul() 支持,bmm() 不支持。
3 模型与算法详解
bmm() vs matmul()
| 函数 | 全称 | 应用场景 | 广播机制 | 效率 |
|---|---|---|---|---|
| matmul() | Matrix Multiply | 通用场景 | 支持 | 一般 |
| bmm() | Batch Matrix Multiply | 批量矩阵乘法 | 不支持 | 高 |
matmul() 应用场景
通用场景,支持广播。
# 示例:10×3×4 和 1×4×5
# 广播后:10×3×4 和 10×4×5 → 10×3×5
result = torch.matmul(a, b)
bmm() 应用场景
批量矩阵乘法,不支持广播。
# 示例:10×3×4 和 10×4×5
# 必须维度一致
result = torch.bmm(a, b)
bmm() 优势
处理序列数据时效率高。
- 并行计算
- 不用一组一组手动算
- 适用于 RNN、Transformer 等框架
4 数学原理与推导
矩阵乘法
$$\text{C} = \text{A} \times \text{B}$$
其中:
- $\text{A}$ 是 $m \times n$ 矩阵
- $\text{B}$ 是 $n \times p$ 矩阵
- $\text{C}$ 是 $m \times p$ 矩阵
批量矩阵乘法
$$\text{C}_i = \text{A}_i \times \text{B}_i, \quad i = 1, 2, ..., N$$
其中:
- $\text{A}$ 是 $N \times m \times n$ 张量
- $\text{B}$ 是 $N \times n \times p$ 张量
- $\text{C}$ 是 $N \times m \times p$ 张量
5 代码示例
import torch
def bmm_demo():
"""演示 bmm() 函数"""
print("=== bmm() 函数简介 ===")
# 1. matmul() 支持广播
print("\n1. matmul() 支持广播:")
a = torch.randn(10, 3, 4) # 10×3×4
b = torch.randn(1, 4, 5) # 1×4×5
result = torch.matmul(a, b)
print(f" a 形状: {a.shape}")
print(f" b 形状: {b.shape}")
print(f" 结果形状: {result.shape}") # 10×3×5
# 2. bmm() 不支持广播
print("\n2. bmm() 不支持广播:")
a = torch.randn(10, 3, 4) # 10×3×4
b = torch.randn(10, 4, 5) # 10×4×5
result = torch.bmm(a, b)
print(f" a 形状: {a.shape}")
print(f" b 形状: {b.shape}")
print(f" 结果形状: {result.shape}") # 10×3×5
# 3. bmm() 报错示例
print("\n3. bmm() 报错示例(维度不匹配):")
try:
a = torch.randn(10, 3, 4)
b = torch.randn(1, 4, 5) # 第一维不是 10
result = torch.bmm(a, b)
except RuntimeError as e:
print(f" 错误: {e}")
# 4. 注意力机制中的 bmm()
print("\n4. 注意力机制中的 bmm():")
batch_size = 2
seq_len = 10
d_k = 32
d_v = 64
Q = torch.randn(batch_size, seq_len, d_k) # 2×10×32
K = torch.randn(batch_size, seq_len, d_k) # 2×10×32
V = torch.randn(batch_size, seq_len, d_v) # 2×10×64
# attention = softmax(Q · K^T / sqrt(d_k)) · V
attn_weights = torch.softmax(torch.bmm(Q, K.transpose(1, 2)) / torch.sqrt(torch.tensor(d_k, dtype=torch.float32)), dim=-1)
attention = torch.bmm(attn_weights, V)
print(f" Q 形状: {Q.shape}")
print(f" K 形状: {K.shape}")
print(f" V 形状: {V.shape}")
print(f" attention 形状: {attention.shape}") # 2×10×64
# 测试
if __name__ == "__main__":
bmm_demo()
代码说明
| 代码 | 说明 |
|---|---|
torch.matmul(a, b) | 通用矩阵乘法,支持广播 |
torch.bmm(a, b) | 批量矩阵乘法,不支持广播 |
K.transpose(1, 2) | 转置第 1 和第 2 维 |
torch.sqrt(...) | 平方根 |
6 重难点与易错提醒
- ❗重点:bmm() 是批量矩阵乘法,不支持广播。
- ❗重点:matmul() 是通用矩阵乘法,支持广播。
- ❗重点:bmm() 在处理序列数据时效率比 matmul() 高。
- ⚠️易错:bmm() 要求两个张量的第一维(batch_size)相同。
7 课堂问答精选
Q1:bmm() 和 matmul() 的区别是什么?
A:bmm() 是批量矩阵乘法,不支持广播,效率高;matmul() 是通用矩阵乘法,支持广播,效率一般。
Q2:bmm() 适用于什么场景?
A:bmm() 适用于处理序列数据,如 RNN、Transformer 等框架,效率比 matmul() 高。
Q3:bmm() 是否支持广播?
A:不支持。bmm() 要求两个张量的第一维(batch_size)相同。
8 本课小结
- bmm() 是批量矩阵乘法函数,不支持广播。
- matmul() 是通用矩阵乘法函数,支持广播。
- bmm() 在处理序列数据时效率比 matmul() 高。
- 适用于 RNN、Transformer 等框架。
9 延伸思考
- 如何用 bmm() 实现注意力机制?
- 如何自定义注意力机制模块?