如何通用化JAX数组任意维度的索引更新逻辑?
通用JAX数组动态维度索引更新方案
要实现任意维度JAX数组的通用索引更新,核心是将调整后的索引数组转换为元组传入.at[]。JAX的索引API要求.at接收元组形式的索引参数,直接用*update_indices无法被正确解析,显式转为元组即可适配任意维度。
修正后的通用代码
import jax.numpy as jnp # 以2D数组为例 b = jnp.arange(16).reshape([4,4]) print("原数组:") print(b) update_indices = jnp.array([[1,1], [3,2], [0,3]]) # 将索引的最后一维移到最前面,得到(数组维度, 更新点数量)形状的数组 update_indices = jnp.moveaxis(update_indices, -1, 0) # 关键:将索引数组转为元组传入.at b = b.at[tuple(update_indices)].set([333, 444, 555]) print("\n更新后数组:") print(b)
3D数组适配示例
以下是3D数组的测试代码,索引逻辑完全复用,无需额外修改:
# 3D数组测试 b_3d = jnp.arange(27).reshape([3,3,3]) print("\n原3D数组:") print(b_3d) # 3D索引:每个元素是(z,y,x)坐标 update_indices_3d = jnp.array([[0,1,2], [2,0,1], [1,1,1]]) update_indices_3d = jnp.moveaxis(update_indices_3d, -1, 0) b_3d = b_3d.at[tuple(update_indices_3d)].set([999, 888, 777]) print("\n更新后3D数组:") print(b_3d)
原理说明
- 调用
jnp.moveaxis(update_indices, -1, 0)后,索引数组形状变为(数组维度, 更新点数量),比如2D时是(2,3),3D时是(3,3)。 - JAX的
.at接口要求索引参数是长度等于数组维度的元组,每个元素对应一维的索引数组。将调整后的索引数组转为元组后,完全符合接口要求,自然适配任意维度的数组。
内容的提问来源于stack exchange,提问作者zephyrus
相关产品推荐
相关产品推荐

