Val*_*acé 3 python arrays jit shapes jax
重要提示:我需要这里的所有内容都与 jit 兼容,否则我的问题是微不足道的:)
我有一个 jax numpy 数组,例如:
a = jnp.array([1,5,3,4,5,6,7,2,9])
Run Code Online (Sandbox Code Playgroud)
首先,我根据一个值对其进行过滤,假设我只保留 < 5 的值
a = jnp.where((a < 5), x=a, y=jnp.nan)
# a is now [ 1. nan 3. 4. nan nan nan 2. nan]
Run Code Online (Sandbox Code Playgroud)
我只想保留非 nan 值:[ 1. 3. 4. 2.]然后我将使用该数组进行其他操作。
但更重要的是,在我的程序执行期间,该代码将被执行多次,并且阈值会发生变化(即它不会总是 5)。
因此,最终数组的形状也会改变。这是我的 jit 编译问题,我不知道如何使其与 jit 兼容,因为形状取决于有多少元素符合阈值条件。
JAX 的 JIT 目前与动态(数据相关)形状的数组不兼容,因此无法执行您的问题所要求的操作。
目前正在进行一些关于在 JAX 转换(如 JIT)中处理动态形状的实验性工作(请参阅https://github.com/google/jax/pull/9335),但我不确定何时可以使用。
通常的解决方法是用具有填充值的静态形状数组重新表达计算;例如,你可以使用这样的东西:
a = jnp.where((a < 5), size=len(a), fill_value=np.nan)
Run Code Online (Sandbox Code Playgroud)
这将创建一个与 长度相同的数组a,前面带有非 nan 值,nan末尾填充有值。
| 归档时间: |
|
| 查看次数: |
1970 次 |
| 最近记录: |