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

Numpy中如何对三维数组的每个子矩阵应用矩阵指数替代for循环

三维Numpy矩阵批量计算子矩阵指数的实现方案

你当前的for循环实现逻辑正确,以下是两种更简洁、部分场景下性能更优的实现方式:

方案1:使用带signature参数的np.vectorize封装

无需额外依赖,语法比手写for循环更简洁,出错概率更低,性能和手写for循环基本持平:

import numpy as np
from scipy.linalg import expm

X = np.random.rand(3, 3, 3)
# 指定输入输出为(m,m)维度的矩阵,实现批量适配
batch_expm = np.vectorize(expm, signature='(m,m)->(m,m)')
y = batch_expm(X)

方案2:块对角矩阵拼接法

将所有子矩阵拼接为大块对角矩阵,仅调用一次expm运算后再拆分,适合子矩阵尺寸小、数量适中的场景:

import numpy as np
from scipy.linalg import expm, block_diag

X = np.random.rand(3, 3, 3)
# 拼接为块对角矩阵
block_X = block_diag(*X)
block_y = expm(block_X)
# 拆分回原结构
y = np.stack([block_y[i*3:(i+1)*3, i*3:(i+1)*3] for i in range(X.shape[0])])

方案3:JAX vmap自动向量化(性能最优)

如果子矩阵数量大、需要更高性能,可使用JAX的自动向量化能力,支持CPU多核心/GPU加速,大规模运算下性能远超纯Python循环:

import jax
import jax.numpy as jnp
from jax.scipy.linalg import expm

X = jnp.array(np.random.rand(3, 3, 3))
# 对第一个维度批量映射expm运算
batch_expm = jax.vmap(expm)
y = batch_expm(X)
# 如需转回numpy数组,调用np.array(y)即可

注:该方案需要先安装JAX库,适合对性能有要求的生产/科研场景。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.04 22:42:03