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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.03 00:27:42