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

如何在PyTorch中实现指定规则的Multi-hot向量?

在PyTorch中生成指定Multi-hot向量

可以通过两种高效方式实现需求,避免低效的逐元素循环,以下是具体方案:

方案一:广播索引+Scatter填充

import torch

max_num = 4
multi_hot_num = 3
x = torch.tensor([0, 2, 1])

# 1. 创建目标形状的零张量
batch_size = len(x)
total_dim = multi_hot_num * max_num
multi_hot = torch.zeros(batch_size, total_dim, dtype=torch.int)

# 2. 计算所有需要置1的索引
start_idx = x * multi_hot_num  # 每个样本对应区间的起始位置
offsets = torch.arange(multi_hot_num)  # 区间内的偏移量(0到multi_hot_num-1)
indices = start_idx.unsqueeze(1) + offsets  # 广播得到所有目标索引

# 3. 填充1到对应位置
multi_hot.scatter_(1, indices, 1)

print(multi_hot)

代码解释

  • 初始化零张量:根据样本数len(x)和总维度multi_hot_num * max_num创建全零张量,指定int类型匹配预期输出。
  • 生成目标索引:通过广播机制批量计算每个样本需要置1的连续索引区间,避免循环遍历。
  • 原地填充:使用scatter_方法在指定列索引位置填充1,完成多热向量构建。

方案二:One-hot+重复插值(更简洁)

import torch

max_num = 4
multi_hot_num = 3
x = torch.tensor([0, 2, 1])

# 先生成普通one-hot向量,再将每个1的位置重复multi_hot_num次
one_hot_base = torch.nn.functional.one_hot(x, num_classes=max_num)
multi_hot = one_hot_base.repeat_interleave(multi_hot_num, dim=1)

print(multi_hot)

输出验证

两种方案的输出均与预期完全一致:

tensor([[1, 1, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0],
        [0, 0, 0, 0, 0, 0, 1, 1, 1, 0, 0, 0, 0, 0, 0],
        [0, 0, 0, 1, 1, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0]], dtype=torch.int32)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.13 17:17:31