【发布时间】:2021-10-17 14:07:45
【问题描述】:
这对我来说可能是一件很简单的事情,但我想知道如何在下面的示例中执行映射。
假设我们有一个函数要计算关于xt、yt 和zt 的导数,但它还需要额外的参数xs、ys 和zs。
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.])
为了评估xt、yt 和zt 中每对数据点的梯度,我必须执行以下操作:
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