PyTorch中是否存在与numpy.random.choice功能相同的方法?
PyTorch中对应numpy.random.choice的实现方法
PyTorch里有多种方式实现和numpy.random.choice一致的功能,根据需求和PyTorch版本可以选择不同方法:
1. 等概率随机选择
有放回采样
用torch.randint()直接生成指定范围的随机整数,对应numpy.random.choice(n, k, replace=True):
# 从0-9中随机选5个元素(有放回) torch.randint(low=0, high=10, size=(5,))
无放回采样
用torch.randperm()生成随机排列后取前k个,对应numpy.random.choice(n, k, replace=False):
# 从0-9中随机选5个不重复元素(无放回) torch.randperm(10)[:5]
2. 带权重的随机选择
如果需要按指定概率采样(对应numpy.random.choice的p参数),用torch.multinomial()是最常用的方案:
# 定义权重张量(需非负且和为1,和numpy的p参数逻辑一致) weights = torch.tensor([0.1, 0.2, 0.3, 0.4]) # 有放回采样3个元素 samples_with_replace = torch.multinomial(weights, num_samples=3, replacement=True) # 无放回采样2个元素 samples_no_replace = torch.multinomial(weights, num_samples=2, replacement=False)
3. PyTorch 1.13+ 原生torch.random.choice()
从PyTorch 1.13版本开始,官方提供了和numpy.random.choiceAPI高度对齐的torch.random.choice(),直接支持大部分参数:
# 从给定张量中随机选元素 arr = torch.tensor([10, 20, 30, 40]) samples = torch.random.choice(arr, 2) # 指定采样权重 samples_weighted = torch.random.choice(arr, 2, p=torch.tensor([0.1, 0.2, 0.3, 0.4])) # 无放回采样 samples_unique = torch.random.choice(arr, 2, replace=False)
内容的提问来源于stack exchange,提问作者ke qi
相关产品推荐
相关产品推荐

