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

JAX中如何使用多层vmap替代嵌套for循环批量计算梯度

你可以通过嵌套vmap实现笛卡尔积维度的自动向量化,完全避免嵌套循环,代码运行效率也会更高,具体实现如下:

import jax.numpy as jnp
from jax import grad, vmap

def fn(xt, yt, zt, xs, ys, zs):
    return jnp.sqrt((xt - xs) ** 2 + (yt - ys) ** 2 + (zt - zs) ** 2)

# 输入数据
xt = jnp.array([1., 2., 3., 4.])
yt = jnp.array([1., 2., 3., 4.])
zt = jnp.array([1., 2., 3., 4.])
xs = jnp.array([1., 2., 3.])
ys = jnp.array([3., 3., 3.])
zs = jnp.array([1., 1., 1.])

# 1. 定义梯度函数,先处理xs/ys/zs的批量维度
grad_fn = grad(fn, argnums=(0,1,2))
vmap_grad = vmap(grad_fn, in_axes=(None, None, None, 0, 0, 0))

# 2. 嵌套vmap实现xt、yt、zt的笛卡尔积批量计算
vmap_zt = vmap(vmap_grad, in_axes=(None, None, 0, None, None, None))
vmap_yt = vmap(vmap_zt, in_axes=(None, 0, None, None, None, None))
vmap_xt = vmap(vmap_yt, in_axes=(0, None, None, None, None, None))

# 3. 一次性计算所有结果
res = vmap_xt(xt, yt, zt, xs, ys, zs)
# 调整形状和原循环版本一致:(4,4,4,3,3) -> (64,3,3)
res = jnp.array(res).reshape(64,3,3)

print(res.shape) # 输出 (64,3,3)

这里三层嵌套的vmap分别对应你原来的三重循环,每层vmap的in_axes参数指定当前要批量处理的输入轴,其余轴设为None就会自动广播,该实现和你原循环版本的计算结果完全等价,且在JAX的JIT编译下运行速度远快于Python层的嵌套循环。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.29 17:24:03