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

如何依据另一数组的索引将JAX数组指定位置设为1

解决JAX数组按索引批量赋值的问题

给定形状为(3, 2, 3)的全0JAX数组X:

[[[0. 0. 0.]
 [0. 0. 0.]]

 [[0. 0. 0.]
 [0. 0. 0.]]

 [[0. 0. 0.]
 [0. 0. 0.]]]

以及形状为(3, 2, 2)的索引数组Y:

[[[1 2]
 [1 2]]

 [[0 2]
 [0 1]]

 [[1 0]
 [1 0]]]

需要将X中每个(i,j)位置对应的Y[i,j]索引处的值设为1,以下是两种可行的实现方式:

方法一:利用JAX的at索引批量赋值

import jax
import jax.numpy as jnp

# 定义原始数组
X = jnp.zeros((3, 2, 3))
Y = jnp.array([
    [[1, 2], [1, 2]],
    [[0, 2], [0, 1]],
    [[1, 0], [1, 0]]
])

# 生成i、j维度的网格索引
i = jnp.arange(X.shape[0])[:, None, None]
j = jnp.arange(X.shape[1])[None, :, None]

# 批量赋值
X_updated = X.at[i, j, Y].set(1.0)

# 输出结果
print(X_updated)

方法二:嵌套vmap处理每个子数组

import jax
import jax.numpy as jnp

X = jnp.zeros((3, 2, 3))
Y = jnp.array([
    [[1, 2], [1, 2]],
    [[0, 2], [0, 1]],
    [[1, 0], [1, 0]]
])

# 定义单个子数组的赋值逻辑
def set_single_row(x, y_indices):
    return x.at[y_indices].set(1.0)

# 嵌套vmap分别处理两个维度
X_updated = jax.vmap(jax.vmap(set_single_row))(X, Y)

print(X_updated)

两种方法都会输出期望的结果:

[[[0. 1. 1.]
  [0. 1. 1.]]

 [[1. 0. 1.]
  [1. 1. 0.]]

 [[1. 1. 0.]
  [1. 1. 0.]]]

内容的提问来源于stack exchange,提问作者jeffreyveon

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.02 09:35:17