You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何通用化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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.10 18:25:34