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

基于范围索引的Numpy向量化填充实现(标签编码)

向量化实现Class ID映射的稠密张量生成

你手头有一个形状为(batch_size, class_id, range_indices)(示例维度是(4, 3, 2))的int64张量,需要按规则生成稠密表示:Class ID 0填充1,ID 1填充2,ID 2填充3,其余情况填0。要实现最符合Numpythonic风格的方案,核心是用索引映射数组完成向量化替换,完全避开显式循环,充分利用numpy的底层优化能力。

实现思路

  1. 先定义一个映射数组:数组的索引对应Class ID,索引位置的值就是要填充的目标值
  2. 利用numpy的布尔掩码或直接索引,批量完成值的替换,无需逐个元素遍历

示例代码

import numpy as np

# 构造示例输入张量 (4,3,2)
input_tensor = np.zeros((4, 3, 2), dtype=np.int64)
# 模拟不同位置的Class ID
input_tensor[0, 0, :] = 0  # Class ID 0
input_tensor[0, 1, :] = 1  # Class ID 1
input_tensor[0, 2, :] = 2  # Class ID 2
input_tensor[1, 1, :] = 1
input_tensor[2, 2, :] = 2
input_tensor[3, 0, :] = 0
# 加入一个超出范围的Class ID测试默认填充
input_tensor[3, 1, :] = 3

# 定义映射规则:索引=Class ID,值=填充值
class_map = np.array([1, 2, 3], dtype=np.int64)

# 生成结果:先创建全0数组,再替换符合条件的位置
result = np.zeros_like(input_tensor)
# 筛选出Class ID在0-2范围内的位置
valid_mask = (input_tensor >= 0) & (input_tensor <= 2)
# 用映射数组批量替换有效值
result[valid_mask] = class_map[input_tensor[valid_mask]]

print("输入张量:")
print(input_tensor)
print("\n输出稠密表示:")
print(result)

方案优势

  • 高效性:完全基于numpy的底层C语言运算,没有Python层面的循环,处理大张量时速度提升明显
  • 简洁性:映射关系一目了然,代码逻辑清晰,符合Numpythonic"用数组操作替代循环"的设计哲学
  • 可维护性:后续修改映射规则只需调整class_map数组,核心逻辑无需改动

对比非向量化实现

如果是类似以下的循环实现(比如常见的非优化版本):

# 非Numpythonic的循环实现示例
result = np.zeros_like(input_tensor)
for batch in range(input_tensor.shape[0]):
    for cls in range(input_tensor.shape[1]):
        for idx in range(input_tensor.shape[2]):
            cid = input_tensor[batch, cls, idx]
            if cid == 0:
                result[batch, cls, idx] = 1
            elif cid == 1:
                result[batch, cls, idx] = 2
            elif cid == 2:
                result[batch, cls, idx] = 3
            else:
                result[batch, cls, idx] = 0

这种写法不仅代码冗余,而且当张量规模扩大(比如batch_size到1000+)时,运算效率会远低于向量化方案——numpy的数组操作是批量并行处理,而Python循环是逐元素解释执行,性能差距会非常显著。

内容的提问来源于stack exchange,提问作者Muhammad Ikhwan Perwira

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.21 04:35:10