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
相关产品推荐
相关产品推荐

