如何在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
代码说明
- 提取True值信息:通过
nonzero拿到所有True值的行列索引,sum(dim=1)统计每行True值的数量。 - 生成无重复随机偏移:
- 第一个偏移量
offset1限制在该行True值的数量范围内; - 第二个偏移量
offset2通过调整避免与offset1重复,保证选取的两个索引不同。
- 第一个偏移量
- 映射全局索引:用累积和计算每行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
相关产品推荐
相关产品推荐

