如何使用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
相关产品推荐
相关产品推荐

