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

如何基于另一数组的索引对jax.numpy数组做无循环向量化掩码修改

JAX按行指定起始位置掩码替换实现

需求说明

现有shape为(2,10)的JAX数组,定义方式:

arr = jnp.ones(shape=(2,10)) * 2

数组初始值全为2:

[[2. 2. 2. 2. 2. 2. 2. 2. 2. 2.]
 [2. 2. 2. 2. 2. 2. 2. 2. 2. 2.]]

给定索引数组[2,4],需要分别指定每一行从对应索引位置开始,将后续元素替换为掩码值-1,预期输出:

[[2. 2. -1. -1. -1. -1. -1. -1. -1. -1.]
 [2. 2. 2. 2.  -1. -1. -1. -1. -1. -1.]]

实现约束:

  • 仅可使用jax.numpy接口
  • 必须完全向量化,禁止使用任何形式的循环
  • 运行效率尽可能高,兼容JAX的JIT等变换

实现代码

核心通过广播生成位置布尔掩码,直接做条件赋值即可,全程无循环:

import jax.numpy as jnp

# 原数组定义
arr = jnp.ones(shape=(2,10)) * 2
# 每行的替换起始索引
start_idx = jnp.array([2, 4])

# 生成列方向索引序列 shape=(10,)
col_pos = jnp.arange(arr.shape[1])
# 广播比较得到布尔掩码:列索引>=当前行起始位置的位置标记为True
mask = col_pos >= start_idx[:, None]
# 按掩码替换值
result = jnp.where(mask, -1, arr)

逻辑说明

  • 列索引序列col_pos对应每一列的位置编号,和原数组列长度一致
  • 将起始索引数组通过[:, None]调整为(2,1)的二维形状,和形状为(10,)的列索引做广播运算,直接得到和原数组shape完全一致的布尔掩码矩阵,整个计算过程在数组底层完成,没有Python层面的循环开销
  • jnp.where按照掩码规则完成值替换:掩码为True的位置填-1,否则保留原数组的值
  • 该方案可直接扩展到任意行、列数的数组场景,只要起始索引数组长度和原数组行数匹配即可,所有操作均为JAX原生支持的数组操作,JIT编译后执行效率最优。

内容的提问来源于stack exchange,提问作者Valentin Macé

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.02 07:57:27