如何依据另一数组的索引将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
相关产品推荐
相关产品推荐

