PyTorch如何按张量元素自身概率采样并获取采样值与对应索引
PyTorch 按概率值采样张量元素实现方案
直接使用PyTorch内置的torch.multinomial算子即可满足全部需求,无需使用WeightedRandomSampler,全程保持张量类型,且可直接获取采样元素在原张量中的索引位置。
核心实现逻辑
torch.multinomial接收权重张量作为输入,按照每个位置的权重值作为概率采样对应位置,直接返回索引张量,全程不脱离PyTorch张量体系。
- 输入的概率张量无需额外处理,已归一化的场景下会直接按照元素值作为采样概率,完全匹配「元素被选中概率等于自身数值」的要求
- 返回值为
torch.Tensor类型,存储的是采样元素在原张量中的索引,可直接定位到原张量的具体位置 - 支持GPU运算、自动微分,兼容所有PyTorch张量操作
代码示例
单次采样
import torch # 输入归一化概率张量 A = torch.tensor([0.0316, 0.2338, 0.2338, 0.2338, 0.0316, 0.0316, 0.0860, 0.0316, 0.0860]) # 采样1个元素,replacement=True表示有放回采样 sample_idx = torch.multinomial(A, num_samples=1, replacement=True) # 取对应采样值 sample_val = A[sample_idx]
运行后返回结果示例:
sample_idx: tensor([2]) # 对应原张量索引为2的位置 sample_val: tensor([0.2338]) # 对应位置的元素值,保持张量类型
批量采样
如果需要一次采样多个元素,仅需修改num_samples参数即可:
# 一次采样5个元素 batch_idx = torch.multinomial(A, num_samples=5, replacement=True) batch_val = A[batch_idx]
注意事项
- 无放回采样(参数设置
replacement=False)时,num_samples的取值不能超过输入张量的长度 - 不推荐使用
WeightedRandomSampler实现该需求:该工具是为PyTorch DataLoader设计的批量采样器,返回值为Python原生整数类型,会脱离张量体系,无法满足返回值保持张量类型的要求
内容的提问来源于stack exchange,提问作者Teodorico Levoff
相关产品推荐
相关产品推荐

