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)

Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐