【问题标题】:JAX vmap behaviourJAX vmap 行为
【发布时间】:2021-06-07 11:03:03
【问题描述】:

我试图了解 JAX vmap 的行为,所以我编写了以下代码:

import jax.numpy as jnp
from jax import vmap

def what(a,b,c):
  z = jnp.dot(a,b)
  return z + c

v_what = vmap(what, in_axes=(None,0,None))

a = jnp.array([1,1,3])
b = jnp.array([2,2])
c = 1.0

v_what(a,b,c)

输出是:

DeviceArray([[3., 3., 7.],
             [3., 3., 7.]], dtype=float32)

我知道唯一被改变的输入是b,但是有人能解释一下为什么会这样吗?以及在我对函数进行矢量化后点积的行为如何?

【问题讨论】:

    标签: python vectorization jax


    【解决方案1】:

    您已指定转换后的函数应映射到b 的第一个轴上,而不是映射到a 或c 的任何轴上。大致来说,您已经创建了一个映射函数来执行此操作:

    def v_what(a, b, c):
      return jnp.stack([what(a, b_i, c) for b_i in b], axis=0)
    

    对于您的输入,每一行中的点积看起来像jnp.dot(a, 2),结果相当于a * 2。

    【讨论】:

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