【发布时间】:2021-09-19 00:11:09
【问题描述】:
我有一个函数compute(x),其中x 是jnp.ndarray。现在,我想使用vmap 将其转换为一个函数,该函数接受一批数组x[i],然后jit 对其进行加速。 compute(x) 类似于:
def compute(x):
# ... some code
y = very_expensive_function(x)
return y
但是,每个数组x[i] 具有不同的长度。我可以通过用尾随零填充数组来轻松解决这个问题,这样它们都具有相同的长度 N 和 vmap(compute) 可以应用于形状为 (batch_size, N) 的批次。
但是,这样做会导致在每个数组 x[i] 的尾随零上也调用 very_expensive_function()。有没有办法修改compute(),使得very_expensive_function() 只在x 的一部分上调用,而不干扰vmap 和jit?
【问题讨论】:
-
显而易见的解决方案是将每个 x[i] 的实际长度也传递给计算,然后对该 x[i] 进行切片,但这可能不受 jax 支持。看看这个:github.com/google/jax/issues/1007。也许传递一个面具是你可以做的。
-
this 回答有用吗?