You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Numpy matmul与einsum较MATLAB慢6-7倍,求性能优化方案

Python NumPy 移植MATLAB代码性能优化问题

性能差异原因

  1. 数组维度与存储不匹配:你的Python代码中X维度为(M,n,N),而MATLAB中是(n,N,M),且MATLAB采用Fortran连续存储(列优先),NumPy默认是C连续存储(行优先)。维度错位导致中间计算生成更大的张量,同时非连续内存访问降低了缓存命中率,额外增加计算开销。
  2. 批量计算优化不足:MATLAB的pagemtimes是专门针对页维度的批量矩阵乘法做了深度优化,而NumPy的einsum在处理高维张量时,无法像MATLAB那样高效处理批量操作,会生成冗余的中间大数组(比如你的einsum版本生成了(M,M,n,N)的张量,内存占用远高于MATLAB的单页计算)。
  3. 循环优化差距:MATLAB的JIT编译器对循环的优化能力强于原生Python,即使是Python的NumPy循环,中间数组的生成与调度也没有MATLAB的JIT高效。

优化建议

1. 对齐数组维度与存储顺序

先将X和w的维度调整为与MATLAB一致,并采用Fortran连续存储,提升内存访问效率:

import numpy as np

n = 4
N = 200
M = 100
# 维度与MATLAB对齐,并用Fortran连续存储
X = np.asfortranarray(0.1 * np.random.rand(n, N, M))
w = np.asfortranarray(0.1 * np.random.rand(N, 1, M))

2. 用Numba JIT编译循环(最推荐)

Numba可以将Python循环编译为机器码,接近MATLAB JIT的优化效果,并行化后能大幅提速:

from numba import njit, prange

@njit(parallel=True)
def compute_G(X, w):
    n, N, M = X.shape
    G = np.zeros((M, M))
    for i in prange(M):
        X_i = X[:, :, i]  # 提取第i页的X
        w_i = w[:, 0, i]   # 提取第i页的w
        for k in range(M):
            X_k = X[:, :, k]
            # 计算核心矩阵操作
            mat = X_i.T @ X_k
            G[k, i] = w_i @ np.exp(mat) @ w[:, 0, k]
    return G

G = compute_G(X, w)

3. 利用矩阵对称性减少计算量

注意到G[k,i] = w_i^T @ exp(X_i^T @ X_k) @ w_k,而G[i,k]是其转置(标量转置等于自身),因此G是对称矩阵。可以只计算上三角部分再复制,减少一半计算量:

@njit(parallel=True)
def compute_G(X, w):
    n, N, M = X.shape
    G = np.zeros((M, M))
    for i in prange(M):
        X_i = X[:, :, i]
        w_i = w[:, 0, i]
        for k in range(i, M):
            X_k = X[:, :, k]
            mat = X_i.T @ X_k
            val = w_i @ np.exp(mat) @ w[:, 0, k]
            G[k, i] = val
            G[i, k] = val
    return G

4. 优化原生NumPy循环(无需额外库)

如果不想用Numba,调整维度后结合tensordot减少中间数组开销:

G = np.zeros((M, M))
# 将w压缩为(N,M),简化操作
w_squeeze = np.squeeze(w, axis=1)
for i in range(M):
    X_i_T = X[:, :, i].T  # (N, n)
    exp_mat = np.exp(X_i_T @ X)  # (N, N, M)
    # 批量计算w_i^T @ exp_mat @ w_k
    G[:, i] = np.tensordot(w_squeeze[:, i].T @ exp_mat, w_squeeze, axes=([1], [0]))

内容的提问来源于stack exchange,提问作者Rushi

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.17 23:49:53