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

PyTorch含预抽取样本的多维无放回随机抽取高效实现

咱们来搞定这个PyTorch张量索引扩展的性能问题——你原来的逐行循环方法在大Z或者大N场景下确实会慢,毕竟Python循环+逐行索引操作没办法充分利用GPU的并行能力。下面给你两个针对性的高效实现方案,分别适配不同的场景:

方案1:批量掩码采样(适合Z中等规模场景)

当Z的大小在显存可承受范围内(比如Z≤1e4,N*Z的矩阵不会占太多GPU内存),这个方案全程用GPU批量操作,性能拉满:

实现思路

  1. 构建一个(N, Z+1)的布尔掩码矩阵,标记每行中未被原索引占用的位置(1表示可采样,0表示已存在)
  2. 对每行生成0~Z的随机排列,通过掩码过滤出有效索引,取前X个
  3. 拼接原索引和新采样的索引

代码实现

import torch

def extend_indices_batch(foo, X, Z):
    device = foo.device
    N, I = foo.shape
    
    # 1. 创建全局掩码:标记每行已存在的索引
    mask = torch.ones((N, Z+1), dtype=torch.bool, device=device)
    # 利用高级索引将原索引对应的位置设为False
    mask[torch.arange(N, device=device).unsqueeze(1), foo] = False
    
    # 2. 生成每行的随机排列(基于随机值排序实现)
    random_perm = torch.rand((N, Z+1), device=device).argsort(dim=1)
    
    # 3. 筛选出有效索引并取前X个
    # 把mask为True的索引按随机排列顺序取出,整理成(N, Z+1-I)的矩阵
    valid_indices = random_perm[mask].view(N, Z+1 - I)
    selected_new = valid_indices[:, :X]
    
    # 4. 拼接原索引和新索引
    return torch.cat([foo, selected_new], dim=1)

优势

  • 完全没有Python循环,所有操作都是GPU并行执行,速度远快于逐行处理
  • 逻辑清晰,能保证新采样的索引绝对不重复,且是无放回均匀分布

方案2:拒绝采样(适合超大Z场景)

如果Z特别大(比如Z≥1e5),(N, Z+1)的掩码矩阵会占用巨量显存,这时候拒绝采样是更高效的选择——因为我们只需要X个新索引,相对于庞大的Z来说,重复的概率极低,几次采样就能凑够数量。

实现思路

  1. 批量生成比需要数量多一些的随机候选索引(比如X*2个,减少循环次数)
  2. 检查每个候选是否出现在该行的原索引中,过滤掉重复的
  3. 重复上述步骤,直到凑够X个有效索引

代码实现

def extend_indices_large_Z(foo, X, Z):
    device = foo.device
    N, I = foo.shape
    result_parts = [foo]
    remaining = X
    
    while remaining > 0:
        # 生成候选索引,数量是剩余需求的2倍,减少循环次数
        candidates = torch.randint(0, Z+1, (N, remaining*2), device=device)
        # 广播对比:检查每个候选是否在该行原索引中
        is_duplicate = (candidates.unsqueeze(2) == foo.unsqueeze(1)).any(dim=2)
        # 筛选出不重复的候选
        valid_candidates = candidates[~is_duplicate].view(N, -1)
        # 取需要的数量
        take_count = min(valid_candidates.shape[1], remaining)
        result_parts.append(valid_candidates[:, :take_count])
        remaining -= take_count
    
    return torch.cat(result_parts, dim=1)

优势

  • 显存占用极低,不需要创建大矩阵,完全适配Z超大的场景
  • 当X远小于Z时,重复概率极低,循环次数极少,性能接近批量方案

适用场景对比

场景推荐方案核心原因
Z较小(N*Z显存可承受)批量掩码采样全程并行,速度最快
Z超大(N*Z显存超出限制)拒绝采样显存友好,效率足够

测试验证

用你给出的示例张量测试:

foo = torch.tensor([[83, 5, 85], [ 7, 60, 66], [89, 25, 63], [58, 67, 47], [12, 46, 40]], device='cuda:0')
Z = 90  # 示例中最大索引是89,所以Z设为90
X = 2

# 用批量方案测试
result = extend_indices_batch(foo, X, Z)
# 验证每行索引唯一
for row in result:
    assert len(row.unique()) == len(row)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.08 08:32:37