【发布时间】:2021-11-24 12:23:42
【问题描述】:
我想使用 vmap 对这段代码进行矢量化以提高性能。
def matrix(dataA, dataB):
return jnp.array([[func(a, b) for b in dataB] for a in dataA])
matrix(data, data)
我试过这个:
def f(x, y):
return func(x, y)
mapped = jax.vmap(f)
mapped(data, data)
但这只会给出对角线条目。
基本上我有一个向量data = [1,2,3,4,5](示例),我想得到一个矩阵,使得矩阵的每个条目(i, j) 是f(data[i], data[j])。因此,生成的矩阵形状将是(len(data), len(data))。
【问题讨论】:
标签: python performance vectorization jit jax