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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.24 12:44:58