如何一次性更新Jax二维或多维数组的多个索引?
Jax二维数组批量更新多坐标的正确方法
你的问题出在索引方式不对——Jax的at索引中,多个坐标不能直接用逗号分隔的多个元组,这会被当成不同维度的索引,而你的数组是2维的,超过2个索引自然报错。
正确的批量更新多坐标的方式是把所有要更新的坐标整理成二维数组,然后按维度拆分索引:
import jax.numpy as jnp # 初始化数组 x = jnp.zeros((3,3)) # 定义要更新的坐标:每一行是一个(x,y)坐标 indices = jnp.array([[1, 0], [0, 1], [1, 1]]) # 对应要设置的值 values = jnp.array([1, 3, 6]) # 批量更新:按维度提取行和列索引 x = x.at[indices[:, 0], indices[:, 1]].set(values) print(x)
输出结果:
[[0. 3. 0.] [1. 6. 0.] [0. 0. 0.]]
另一种等价写法是把坐标数组转置后拆成元组:
x = x.at[tuple(indices.T)].set(values)
这种方式不管你要更新1个还是1000个坐标都适用,不会生成大量中间数组,完全适配训练场景的性能需求。
内容的提问来源于stack exchange,提问作者move37
相关产品推荐
相关产品推荐

