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

如何在PyTorch中高效为布尔张量每行随机采样2个True值索引

问题

给定N×K(K≥2)的PyTorch布尔张量,每行至少包含2个True值,需要为每行随机选取2个True值对应的列索引。

示例

输入张量

tensor([[False,  True, False, False,  True],
        [ True,  True,  True, False,  True],
        [ True,  True, False,  True, False],
        [False,  True, False, False,  True],
        [False,  True, False, False,  True]])

可能的输出

tensor([[1, 4],
        [4, 2],
        [0, 3],
        [4, 1],
        [4, 1]])

现有方案的问题

当前实现依赖NumPy的np.random.choice,需要列表推导遍历每行,还要在CPU和GPU间传输数据:

available_trues = [t.nonzero(as_tuple=False).flatten() for t in input_tensor]
only_2_trues = [np.random.choice(t.cpu(), size=(2,), replace=False) for t in available_trues]
only_2_trues = torch.from_numpy(np.stack(only_2_trues)).cuda()

该方案未矢量化,处理大矩阵时性能严重下降,需要无需列表推导、无需跨设备传输的高效实现。


高效矢量化实现

以下是纯PyTorch的矢量化方案,全程在张量所在设备(GPU/CPU)运行,无跨设备数据传输:

import torch

def sample_two_true_indices(input_tensor):
    # 获取所有True值的行、列索引,以及每行True值的数量
    row_indices, col_indices = input_tensor.nonzero(as_tuple=True)
    row_true_counts = input_tensor.sum(dim=1)
    
    # 为每行生成两个不重复的随机偏移量(基于该行True值的位置列表)
    # 生成第一个偏移量,确保不超过该行True值数量范围
    offset1 = torch.randint(0, row_true_counts.max(), row_true_counts.shape, device=input_tensor.device)
    offset1 = torch.clamp(offset1, 0, row_true_counts - 1)
    
    # 生成第二个偏移量,确保与offset1不重复
    offset2 = torch.randint(0, row_true_counts.max() - 1, row_true_counts.shape, device=input_tensor.device)
    offset2 = torch.where(offset2 >= offset1, offset2 + 1, offset2)
    offset2 = torch.clamp(offset2, 0, row_true_counts - 1)
    
    # 计算每行True值在全局col_indices中的起始位置
    cum_start_indices = torch.cat([torch.tensor([0], device=input_tensor.device), row_true_counts.cumsum(dim=0)[:-1]])
    
    # 根据偏移量获取最终的列索引并堆叠结果
    idx1 = cum_start_indices + offset1
    idx2 = cum_start_indices + offset2
    result = torch.stack([col_indices[idx1], col_indices[idx2]], dim=1)
    
    return result

代码说明

  1. 提取True值信息:通过nonzero拿到所有True值的行列索引,sum(dim=1)统计每行True值的数量。
  2. 生成无重复随机偏移:
    • 第一个偏移量offset1限制在该行True值的数量范围内;
    • 第二个偏移量offset2通过调整避免与offset1重复,保证选取的两个索引不同。
  3. 映射全局索引:用累积和计算每行True值在全局列索引列表中的起始位置,加上偏移量得到最终的列索引,最后堆叠成N×2的结果张量。

测试示例

input_tensor = torch.tensor([[False, True, False, False, True],
                            [True, True, True, False, True],
                            [True, True, False, True, False],
                            [False, True, False, False, True],
                            [False, True, False, False, True]], dtype=torch.bool)

output = sample_two_true_indices(input_tensor)
print(output)

输出示例(随机结果):

tensor([[1, 4],
        [2, 0],
        [3, 1],
        [1, 4],
        [4, 1]])

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.16 10:48:06