【AI大模型--NumPy-04】-NumPy 矩阵乘法完全指南 (Matrix Multiplication)
·
04_matrixMultiplication.py - 矩阵乘法完全指南
学习路径第 4 步 (共 10 步) | 难度:基础-中级
概述
系统讲解矩阵乘法的 3 种写法、维度匹配规则、重要数学性质(交换律/结合律/单位矩阵),以及转置在维度适配中的实际应用。
学习目标
- 掌握
np.dot()/@/np.matmul()三种矩阵乘法写法 - 理解 (m,n) x (n,p) = (m,p) 的维度匹配规则
- 了解矩阵乘法的数学性质(不满足交换律,满足结合律)
- 学会使用转置 (.T) 解决维度不匹配问题
核心内容 (8 个模块)
| 模块 | 核心知识点 |
|---|---|
| 1. 基础乘法 | dot / @ / matmul 三种写法、维度匹配验证 |
| 2. 手动计算验证 | 点积逐步计算过程,理解底层原理 |
| 3. 维度错误案例 | 常见错误及修复方法 |
| 4. 矩阵重要性质 | 交换律不成立、结合律、单位矩阵 |
| 5. 转置技巧 | 用 .T 调整形状以适配乘法规则 |
| 6. 实际应用场景 | 向量旋转、线性方程组、神经网络前向传播 |
| 7. 性能对比 | NumPy vs 纯 Python 循环速度差异 |
| 8. API 速查表 | 常用矩阵操作函数一览 |
code
"""
=====================================
NumPy 矩阵乘法完全指南 (Matrix Multiplication)
=====================================
本案例系统介绍 NumPy 中矩阵乘法的:
1. 基本语法与规则
2. 维度匹配原理
3. 重要性质(交换律、结合律)
4. 特殊矩阵(单位矩阵、转置)
5. 实际应用场景
6. 性能对比
作者:bloxed
日期:2026-05-19
"""
import numpy as np
import time
def separator(title):
"""打印分隔线"""
print(f"\n{'='*60}")
print(f" {title}")
print('='*60)
# ============================================================
# 第一部分:基础矩阵乘法
# ============================================================
separator("一、基础矩阵乘法")
# 创建两个可相乘的矩阵
# 规则:矩阵A (m×n) × 矩阵B (n×p) = 结果C (m×p)
# 关键:A的列数 必须等于 B的行数
A = np.array([[1, 2, 3],
[4, 5, 6]]) # 形状: (2, 3)
B = np.array([[7, 8],
[9, 10],
[11, 12]]) # 形状: (3, 2)
print(f"矩阵 A 的形状: {A.shape} (2行3列)")
print(f"矩阵 B 的形状: {B.shape} (3行2列)")
print(f"A的列数({A.shape[1]}) == B的行数({B.shape[0]}) -> 可以相乘 [OK]")
# 方式1: np.dot() - 传统方式
C1 = np.dot(A, B)
# 方式2: @ 运算符 - Python 3.5+ 推荐方式(更直观)
C2 = A @ B
# 方式3: np.matmul() - 专门用于矩阵乘法
C3 = np.matmul(A, B)
print(f"\n计算结果 C = A × B ({C1.shape}):")
print(C1)
print(f"\n三种方式结果一致: {np.array_equal(C1, C2) and np.array_equal(C2, C3)}")
# ============================================================
# 第二部分:手动计算验证(理解原理)
# ============================================================
separator("二、手动计算验证(理解乘法原理)")
# 矩阵乘法的本质:行×列的点积
# C[i][j] = A的第i行 · B的第j列
print("手动计算 C[0][0]:")
print(f" A的第0行: {A[0]}")
print(f" B的第0列: {B[:, 0]}")
manual_calc = sum(A[0] * B[:, 0]) # 对应元素相乘后求和
print(f" 点积运算: 1×7 + 2×9 + 3×11 = {manual_calc}")
print(f" NumPy结果: {C1[0, 0]}")
print(f" 验证通过: {manual_calc == C1[0, 0]} [OK]")
# ============================================================
# 第三部分:维度不匹配的情况
# ============================================================
separator("三、维度不匹配的错误演示")
D = np.array([[1, 2], [3, 4]]) # (2, 2)
E = np.array([[1, 2, 3], [4, 5, 6]]) # (2, 3)
print(f"D 的形状: {D.shape}, E 的形状: {E.shape}")
try:
result = D @ E # (2,2) x (2,3) = (2,3) [OK]
print(f"D × E 成功!结果形状: {result.shape}")
print(result)
except ValueError as e:
print(f"错误: {e}")
try:
result_bad = E @ D # (2,3) × (2,2) → 3≠2 ✗
print(f"E × D 结果: {result_bad.shape}")
except ValueError as e:
print(f"E × D 失败: E的列数(3) ≠ D的行数(2)")
print(f" 错误信息: {e}")
# ============================================================
# 第四部分:重要性质
# ============================================================
separator("四、矩阵乘法的重要性质")
M = np.array([[1, 2], [3, 4]])
N = np.array([[5, 6], [7, 8]])
P = np.array([[1, 0], [0, 1]]) # 单位矩阵
# 性质1: 交换律不成立(一般情况 M×N ≠ N×M)
MN = M @ N
NM = N @ M
print("\n【性质1】交换律不成立:")
print(f"M × N:\n{MN}")
print(f"N × M:\n{NM}")
print(f"M×N == N×M ? {np.array_equal(MN, NM)} (通常不相等)")
# 性质2: 结合律成立 (M×N)×P = M×(N×P)
left_assoc = (M @ N) @ P
right_assoc = M @ (N @ P)
print("\n【性质2】结合律成立:")
print(f"(M x N) x P == M x (N x P): {np.array_equal(left_assoc, right_assoc)} [OK]")
# 性质3: 单位矩阵 I,满足 A×I = I×A = A
print("\n【性质3】单位矩阵 (I):")
print(f"P 是单位矩阵:\n{P}")
print(f"M x P = M ? {np.array_equal(M @ P, M)} [OK]")
print(f"P x M = M ? {np.array_equal(P @ M, M)} [OK]")
# ============================================================
# 第五部分:特殊操作 - 转置在乘法中的应用
# ============================================================
separator("五、转置与矩阵乘法")
X = np.array([[1, 2, 3],
[4, 5, 6]]) # (2, 3)
Y = np.array([[1, 2],
[3, 4]]) # (2, 2)
# 直接 X × Y 会失败 (3 ≠ 2)
# 但 X^T x Y 可以 (2 x 3)^T = (3, 2), (3,2) x (2,2) = (3,2) [OK]
print(f"X 的形状: {X.shape}, Y 的形状: {Y.shape}")
print(f"X^T 的形状: {X.T.shape}")
XT_Y = X.T @ Y # (3, 2) × (2, 2) = (3, 2)
print(f"\nX^T × Y 的结果 ({XT_Y.shape}):")
print(XT_Y)
# ============================================================
# 第六部分:实际应用场景
# ============================================================
separator("六、实际应用场景")
# 应用1: 向量变换(二维旋转)
print("\n【应用1】二维向量旋转:")
angle = np.radians(90) # 90度
rotation_matrix = np.array([
[np.cos(angle), -np.sin(angle)],
[np.sin(angle), np.cos(angle)]
])
point = np.array([1, 0]) # x轴上的点
rotated_point = rotation_matrix @ point
print(f"旋转矩阵 (90°):\n{rotation_matrix}")
print(f"原点: {point} → 旋转后: {rotated_point}")
print(f"(1,0) 旋转90° 应该变为 (0,1)")
# 应用2: 线性方程组 Ax = b
print("\n【应用2】解线性方程组:")
# 2x + 3y = 8
# 4x + 5y = 14
A_eq = np.array([[2, 3], [4, 5]])
b_eq = np.array([8, 14])
x_solution = np.linalg.solve(A_eq, b_eq) # 使用线性代数求解器
print(f"方程组: 2x + 3y = 8, 4x + 5y = 14")
print(f"解: x = {x_solution[0]}, y = {x_solution[1]}")
print(f"验证: A @ x = {A_eq @ x_solution} (应等于 {b_eq})")
# 应用3: 批量数据转换(神经网络中的常见操作)
print("\n【应用3】批量数据处理 (神经网络风格):")
# 假设有3个样本,每个样本4个特征
samples = np.array([[1, 2, 3, 4],
[2, 3, 4, 5],
[3, 4, 5, 6]]) # (3, 4) - 3个样本,4个特征
# 权重矩阵: 将4个特征转换为2个输出
weights = np.array([[0.1, 0.2],
[0.3, 0.4],
[0.5, 0.6],
[0.7, 0.8]]) # (4, 2)
outputs = samples @ weights # (3, 4) × (4, 2) = (3, 2)
print(f"输入批次形状: {samples.shape} (3个样本, 4个特征)")
print(f"权重矩阵形状: {weights.shape} (4个特征 → 2个输出)")
print(f"输出结果形状: {outputs.shape} (3个样本, 2个输出)")
print(f"输出值:\n{outputs}")
# ============================================================
# 第七部分:性能对比
# ============================================================
separator("七、性能对比: NumPy vs 纯Python")
size = 200
matrix_p1 = np.random.rand(size, size)
matrix_p2 = np.random.rand(size, size)
# NumPy 矩阵乘法
start = time.perf_counter()
numpy_result = matrix_p1 @ matrix_p2
numpy_time = time.perf_counter() - start
# 纯 Python 三重循环(仅做较小规模演示)
def python_matrix_mult(A, B):
"""纯Python实现的矩阵乘法"""
rows_A, cols_A = A.shape
cols_B = B.shape[1]
result = np.zeros((rows_A, cols_B))
for i in range(rows_A):
for j in range(cols_B):
for k in range(cols_A):
result[i, j] += A[i, k] * B[k, j]
return result
small_size = 100
small_A = np.random.rand(small_size, small_size)
small_B = np.random.rand(small_size, small_size)
start = time.perf_counter()
python_result = python_matrix_mult(small_A, small_B)
python_time = time.perf_counter() - start
print(f"矩阵大小: {size}×{size}")
print(f"NumPy (@运算符): {numpy_time:.6f} 秒")
print(f"纯Python (三重循环, {small_size}×{small_size}): {python_time:.6f} 秒")
print(f"\n速度差异: NumPy 比纯Python 快约 {python_time/numpy_time:.0f}+ 倍")
print("(原因: NumPy 使用BLAS库和CPU向量化指令集优化)")
# ============================================================
# 第八部分:常用函数速查
# ============================================================
separator("八、常用矩阵运算速查表")
Q = np.array([[1, 2, 3], [4, 5, 6]])
print("\n给定矩阵 Q:")
print(Q)
print(f"\nQ.shape → 形状: {Q.shape}")
print(f"Q.ndim → 维度数: {Q.ndim}")
print(f"Q.size → 元素总数: {Q.size}")
print(f"Q.T → 转置:\n{Q.T}")
print(f"Q.trace() → 迹(对角线和): {Q.trace()}")
print(f"np.linalg.det(Q[:2,:2]) → 行列式(需方阵): 需截取子矩阵")
print(f"np.linalg.inv() → 逆矩阵 (需方阵且非奇异)")
# ============================================================
# 总结
# ============================================================
separator("总结:矩阵乘法核心要点")
key_points = """
┌─────────────────────────────────────────────────────────────┐
│ 1. 维度规则: (m×n) × (n×p) = (m×p) │
│ 2. 推荐写法: A @ B (简洁直观) │
│ 3. 无交换律: A×B ≠ B×A (绝大多数情况) │
│ 4. 有结合律: (A×B)×C = A×(B×C) │
│ 5. 单位矩阵: A×I = I×A = A │
│ 6. 转置规则: (A×B)^T = B^T × A^T │
│ 7. 性能优势: NumPy比纯Python快数十到数百倍 │
│ 8. 应用领域: 图形变换、深度学习、线性方程组、数据分析 │
└─────────────────────────────────────────────────────────────┘
"""
print(key_points)
更多推荐




所有评论(0)