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

基于多位置独热编码的数组非原地更新技术求助

用One-Hot编码实现多索引的非原地数组更新

我需要根据给定的单个或多个索引,对数组进行非原地更新:保留指定索引位置的元素,其余置0。单索引场景已实现,但多索引场景无法得到正确结果,相关代码如下:

import numpy as np

def one_hot(x, depth):
    return np.take(np.eye(depth), x, axis=0)

arr1 = np.random.rand(2, 5, 3, 8)
ohe_depth = 5

# 单索引场景:正常工作
curr_idx = np.array([1])
curr_idx_ohe = one_hot(curr_idx, ohe_depth).reshape(1, ohe_depth, 1, 1)
out = arr1 * curr_idx_ohe  # 输出正确

# 多索引场景:需要实现等价于下方原地更新的非原地操作
curr_idx = np.arange(3)
curr_idx_ohe = one_hot(curr_idx, ohe_depth)

# 目标等价操作(禁止原地更新):
# out2 = np.zeros_like(arr1)
# out2[:, curr_idx, ...] = arr1[:, curr_idx, ...]

解决方案

你只需要将多索引生成的one-hot数组求和后,调整形状为可广播的维度,再和原数组相乘即可:

# 多索引场景的正确处理
curr_idx = np.arange(3)
# 生成one-hot并沿索引维度求和,得到目标轴的掩码
curr_idx_ohe = one_hot(curr_idx, ohe_depth).sum(axis=0).reshape(1, ohe_depth, 1, 1)
out2 = arr1 * curr_idx_ohe

原理说明

  1. one_hot(curr_idx, ohe_depth) 生成形状为 (3, 5) 的数组,每一行对应一个索引的one-hot向量;
  2. sum(axis=0) 将多个one-hot向量求和,得到形状为 (5,) 的掩码数组,其中curr_idx对应的位置值为1,其余为0;
  3. reshape(1, ohe_depth, 1, 1) 将掩码调整为和原数组arr1(形状(2,5,3,8))可广播的维度,确保乘法操作能正确作用于目标轴;
  4. 最终相乘后,原数组中目标索引位置的元素被保留,其余位置被置0,完全等价于你标注的原地更新逻辑,且全程为非原地操作。

验证

可以通过以下代码确认结果一致性:

# 手动生成目标结果
out2_manual = np.zeros_like(arr1)
out2_manual[:, curr_idx, ...] = arr1[:, curr_idx, ...]

# 验证是否等价
print(np.allclose(out2, out2_manual))  # 输出True

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 23:12:18