【问题标题】:Efficiently fill an array from a function从函数中有效地填充数组
【发布时间】: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 从函数填充多维数组的最佳方法是什么?

【问题讨论】:

    标签: python numpy jax


    【解决方案1】:

    JAX 有一个专门为此类应用程序设计的vmap transform

    只要您的get_coords 函数与 JAX 兼容(即是一个没有副作用的纯函数),您就可以在一行中完成:

    from jax import vmap
    xx, yy, zz = vmap(vmap(get_coord))(aa, bb)
    

    【讨论】:

    • 为了让它工作,似乎我需要像这样显式添加in_axesjax.vmap(jax.vmap(get_coord, (None,0)), (0, None))
    【解决方案2】:

    这可以通过使用jax.vmapjax.numpy.vectorize 函数有效地实现。

    一个使用vectorize的例子:

    import jax.numpy as jnp
    
    def get_coord(a, b):
        return jnp.array([a, b, a+b])
    
    f0 = jnp.vectorize(get_coord, signature='(),()->(i)')
    f1 = jnp.vectorize(f0, excluded=(1,), signature='()->(i,j)')
    
    xyz = f1(a,b)
    

    vectorize 函数在底层使用vmap,所以这应该完全等同于:

    f0 = jax.vmap(get_coord, (None, 0))
    f1 = jax.vmap(f0, (0, None)) 
    

    使用vectorize的好处是代码仍然可以在标准的numpy中运行。缺点是代码不太简洁,并且可能由于包装器而产生少量开销。

    【讨论】:

      猜你喜欢
      • 2013-07-26
      • 2014-09-21
      • 1970-01-01
      • 2021-06-24
      • 2011-07-26
      • 1970-01-01
      • 1970-01-01
      • 2021-06-03
      • 2021-06-25
      相关资源
      最近更新 更多