【问题标题】:Multiple `vmap` in JAX?JAX中的多个`vmap`?
【发布时间】:2021-10-17 14:07:45
【问题描述】:

这对我来说可能是一件很简单的事情,但我想知道如何在下面的示例中执行映射。

假设我们有一个函数要计算关于xtytzt 的导数,但它还需要额外的参数xsyszs

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.])

为了评估xtytzt 中每对数据点的梯度,我必须执行以下操作:

fn_prime = vmap(grad(fn, argnums=(0, 1, 2)), in_axes=(None, None, None, 0, 0, 0))

a = []
for _xt in xt:
    for _yt in yt:
        for _zt in zt:
            a.append(fn_prime(_xt, _yt, _zt, xs, ys, zs))

它会产生一个元组列表。 一旦列表转换为jnp.array,它的形状如下:

a = jnp.array(a)
print(f`shape = {a.shape}')
shape = (64, 3, 3)

我的问题是: 有没有办法避免这种 for 循环并在同一扫描中评估所有梯度?

【问题讨论】:

    标签: python vectorization jax


    【解决方案1】:

    对于这种情况,一个好的经验法则是,每个嵌套的 for 循环都转换为嵌套的 vmap 覆盖适当的 in_axis。考虑到这一点,您可以用这种方式重新表达您的计算:

    def f_loops(xt, yt, zt, xs, ys, zs):
      a = []
      for _xt in xt:
        for _yt in yt:
          for _zt in zt:
            a.append(fn_prime(_xt, _yt, _zt, xs, ys, zs))
      return jnp.array(a)
    
    def f_vmap(xt, yt, zt, xs, ys, zs):
      f_z = vmap(fn_prime, in_axes=(None, None, 0, None, None, None))
      f_yz = vmap(f_z, in_axes=(None, 0, None, None, None, None))
      f_xyz = vmap(f_yz, in_axes=(0, None, None, None, None, None))
      return jnp.stack(f_xyz(xt, yt, zt, xs, ys, zs), axis=3).reshape(64, 3, 3)
    
    out_loops = f_loops(xt, yt, zt, xs, ys, zs)
    out_vmap = f_vmap(xt, yt, zt, xs, ys, zs)
    
    np.testing.assert_allclose(out_loops, out_vmap)  # passes
    

    【讨论】:

      猜你喜欢
      • 1970-01-01
      • 2021-06-07
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      • 2021-11-05
      • 2021-08-11
      • 1970-01-01
      • 2022-11-21
      相关资源
      最近更新 更多