如何在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
相关产品推荐
相关产品推荐

