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

如何使用jax vmap实现嵌套循环向量化,生成全配对计算矩阵?

问题根因

单次调用vmap默认会将两个输入参数的对应位置元素配对计算,也就是只会计算f(data[i], data[i]),最终输出一维数组,对应矩阵的对角线元素,和预期不符。

正确实现方案

要生成所有(i,j)配对的二维矩阵,需要嵌套两层vmap,分别控制行、列维度的映射逻辑,两种常用写法如下:

写法1:lambda嵌套写法(简洁版)

import jax
import jax.numpy as jnp

# 替换为你自己的func实现
def func(x, y):
    return x * y + x  # 示例函数

data = jnp.array([1,2,3,4,5])
# 外层vmap遍历行维度i,内层vmap遍历列维度j
matrix_result = jax.vmap(lambda a: jax.vmap(lambda b: func(a,b))(data))(data)

print(matrix_result.shape) # 输出 (5,5)
# 验证:matrix_result[i,j] 等于 func(data[i], data[j])

写法2:显式指定映射轴(可读性更高)

通过in_axes参数明确指定每个参数是否参与当前维度的映射,None表示该参数在当前vmap层级保持不变:

# 定义二元函数
f = lambda x, y: func(x, y)
# 第一层vmap:固定第一个参数x,遍历第二个参数y的所有元素,生成单行
vmap_row = jax.vmap(f, in_axes=(None, 0))
# 第二层vmap:遍历第一个参数x的所有元素,每行对应一个x
vmap_mat = jax.vmap(vmap_row, in_axes=(0, None))

# 调用生成矩阵
matrix_result = vmap_mat(data, data)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.02 00:54:00