【发布时间】:2020-10-06 08:07:43
【问题描述】:
我想从一个函数构造一个二维数组,以便我可以利用jax.jit。
我通常使用numpy 执行此操作的方式是创建一个空数组,然后就地填充该数组。
xx = jnp.empty((num_a, num_b))
yy = jnp.empty((num_a, num_b))
zz = jnp.empty((num_a, num_b))
for ii_a in range(num_a):
for ii_b in range(num_b):
a = aa[ii_a, ii_b]
b = bb[ii_a, ii_b]
xyz = self.get_coord(a, b)
xx[ii_a, ii_b] = xyz[0]
yy[ii_a, ii_b] = xyz[1]
zz[ii_a, ii_b] = xyz[2]
为了在jax 内完成这项工作,我尝试使用jax.opt.index_update。
xx = xx.at[ii_a, ii_b].set(xyz[0])
yy = yy.at[ii_a, ii_b].set(xyz[1])
zz = zz.at[ii_a, ii_b].set(xyz[2])
这运行没有错误,但当我尝试使用 @jax.jit 装饰器时非常慢(至少比纯 python/numpy 版本慢一个数量级)。
使用jax 从函数填充多维数组的最佳方法是什么?
【问题讨论】: