jax.numpy.delete中assume_unique_indices参数报错求助
解决JAX 0.4.8中jnp.delete的assume_unique_indices参数报错问题
问题原因
assume_unique_indices参数是JAX后续版本新增的特性,你使用的0.4.8版本还未引入这个参数,所以调用时会触发「意外关键字参数」的TypeError。官方文档默认展示的是最新版API,和你当前使用的旧版本不兼容。
解决方案
1. 升级JAX版本
把JAX升级到支持该参数的版本(比如0.4.13及以上),升级命令:
pip install --upgrade jax jaxlib
升级后就能正常运行你的示例代码:
import jax import jax.numpy as jnp arr = jnp.array([1, 2, 3, 4, 5]) idx = jnp.array([0, 2, 4]) print(jax.__version__) # 输出≥0.4.13即可 out = jnp.delete(arr, idx, assume_unique_indices=True) print(out) # [2 4]
2. 不升级版本,手动实现JIT兼容的唯一索引删除
如果无法升级,又需要JIT编译删除逻辑,可以用布尔掩码的方式实现——因为已知索引唯一,直接构造掩码即可:
import jax import jax.numpy as jnp from jax import jit arr = jnp.array([1, 2, 3, 4, 5]) idx = jnp.array([0, 2, 4]) @jit def delete_unique_indices(arr, idx): mask = jnp.ones(arr.shape[0], dtype=bool) mask = mask.at[idx].set(False) return arr[mask] out = delete_unique_indices(arr, idx) print(out) # [2 4]
这个实现支持JIT编译,效果和带assume_unique_indices=True的jnp.delete一致。
内容的提问来源于stack exchange,提问作者Chenxi
相关产品推荐
相关产品推荐

