PyTorch:如何无循环实现二维张量按mask索引得到列表结果
无循环实现按行分组的布尔索引结果
可以通过计算每行有效元素个数结合torch.split实现,完全避免Python循环,利用PyTorch的底层优化提升效率,步骤如下:
- 统计布尔mask每行中为
True的元素数量,得到计数张量counts - 提取所有mask为
True的元素得到扁平化张量 - 用
counts作为拆分依据,将扁平化张量拆分为对应每行的子张量列表
代码示例
import torch # 示例二维张量 R = torch.tensor([[1, 2, 3], [4, 5, 6], [7, 8, 9]]) # 对应形状的布尔mask mask = torch.tensor([[True, False, False], [True, True, True], [False, True, False]]) # 1. 计算每行有效元素个数 counts = mask.sum(dim=1) # 输出: tensor([1, 3, 1]) # 2. 提取所有有效元素 flattened_valid = R[mask] # 输出: tensor([1, 4, 5, 6, 8]) # 3. 按行拆分得到分组结果 grouped_result = list(torch.split(flattened_valid, counts.tolist())) print(grouped_result) # 输出: [tensor([1]), tensor([4, 5, 6]), tensor([8])]
处理空行情况
如果存在某行mask全为False(对应counts中为0),torch.split会生成空张量,可通过列表推导式过滤:
filtered_result = [t for t in grouped_result if t.numel() > 0]
这个方法的核心是利用PyTorch的内置算子完成拆分,避免了Python循环带来的性能损耗,适合大规模张量场景。
内容的提问来源于stack exchange,提问作者Josemi
相关产品推荐
相关产品推荐

