【发布时间】:2020-11-04 18:00:52
【问题描述】:
我有两种方法对jnp = jax.numpy 中的矩阵求幂。一种
直截了当的:
jnp.exp(-X/reg)
还有一些额外的操作:
def exp_reg(X, reg):
K = jnp.empty_like(X)
K = jnp.divide(X, -reg)
return jnp.exp(K)
但是,当我测试它们时:
%timeit jnp.exp(-X/reg).block_until_ready()
%timeit exp_reg(X, reg).block_until_ready()
尽管表面上增加了一些额外开销,但第二种方法的表现却优于其他方法。我运行了一个%timeit,其矩阵大小为 2000 x 2000:
7.85 ms ± 567 µs per loop (mean ± std. dev. of 7 runs, 100 loops each)
5.19 ms ± 52.6 µs per loop (mean ± std. dev. of 7 runs, 100 loops each)
为什么会这样?
【问题讨论】:
标签: performance numpy jax