如何在NumPy/JAX中向量化实现批量多矩阵连乘?
在JAX中向量化实现3x3矩阵序列的连乘
核心问题分析
你需要对形状为(B, N, N, 3, 3)的数组X,沿第二个维度(N个元素)依次执行矩阵乘法,且避免显式for循环。einsum无法实现这一点,因为它的核心是张量收缩求和,而非累积式的连乘操作。JAX提供了专门的工具来处理这类累积操作,无需手动循环。
解决方案:使用JAX的扫描原语
JAX中有两个高效的扫描函数适合这个场景:jax.lax.scan(线性扫描)和jax.lax.associative_scan(分治扫描,效率更高),两者都支持自动微分且会被JAX编译器优化。
情况1:X形状为(B, N, 3, 3)(每个batch含N个3x3矩阵)
如果你的X第二个N是笔误,实际是每个batch包含N个3x3矩阵,目标是得到每个batch内所有矩阵的连乘结果(形状(B, 3, 3)):
import jax import jax.numpy as jnp # 生成测试数据 B, N = 2, 5 X = jnp.random.normal(size=(B, N, 3, 3)) # 方法1:用jax.lax.scan(线性扫描) def scan_step(carry, elem): # carry是当前累积的乘积,elem是当前待乘的矩阵 return carry @ elem, None # 初始值为batch维度的单位矩阵 initial_carry = jnp.eye(3)[None, ...] # 形状(1, 3, 3),自动广播到(B, 3, 3) final_product, _ = jax.lax.scan(scan_step, initial_carry, X, axis=1) # final_product形状:(B, 3, 3) # 方法2:用jax.lax.associative_scan(分治扫描,大N时更高效) def matmul_op(a, b): return a @ b # 对N维度执行关联扫描,结果的最后一个元素就是完整连乘结果 scan_results = jax.lax.associative_scan(matmul_op, X, axis=1) final_product_assoc = scan_results[:, -1, :, :] # final_product_assoc形状:(B, 3, 3)
情况2:X形状确实为(B, N, N, 3, 3)
如果每个batch内包含N组,每组有N个3x3矩阵,目标是对每组(第三个维度)的N个矩阵执行连乘,得到形状(B, N, 3, 3)的结果:
import jax import jax.numpy as jnp # 生成测试数据 B, N = 2, 5 X = jnp.random.normal(size=(B, N, N, 3, 3)) def matmul_op(a, b): return a @ b # 沿第二个N维度(axis=1)执行关联扫描 scan_results = jax.lax.associative_scan(matmul_op, X, axis=1) final_product = scan_results[:, -1, :, :, :] # final_product形状:(B, N, 3, 3),对应每组的完整连乘结果
关键说明
associative_scan依赖操作的结合律(矩阵乘法满足结合律),因此可以用分治法并行计算,时间复杂度为O(logN),远快于线性扫描的O(N),适合大N场景。- 两种方法都支持自动微分,完全兼容JAX的反向传播需求。
- 显式for循环在JAX中会被自动矢量化,但扫描原语的优化程度更高,尤其是
associative_scan。
内容的提问来源于stack exchange,提问作者rick
相关产品推荐
相关产品推荐

