PyTorch中如何高效生成满足特定区间条件的布尔张量?
问题
我有一个名为idx的tensor,其元素为0到27的整数,示例如下:
idx = torch.tensor([9, 0, 6, 5, 2, 1, 10, 18, 26, 0, 11])
我希望生成另一个布尔类型的tensor(flag),当元素严格位于左侧的1与右侧的0之间,或位于1之后到tensor末尾时,对应值为True,示例如下:
flag = torch.tensor([False, False, False, False, False, False, True, True, True, False, False])
需要注意的是,如果1之后没有出现0,那么从1开始到tensor末尾的元素都对应True,示例如下:
idx = torch.tensor([9, 0, 6, 5, 2, 1, 10, 18]) flag = torch.tensor([False, False, False, False, False, False, True, True])
请问是否存在无需昂贵for循环的快速实现方式?若必须使用循环,最高效的实现方法是什么?
重要说明:我的idx中0和1总是按固定间隔交替出现,具体来说,1与0、0与1之间始终恰好有3个元素。
高效实现方案
一、利用固定间隔特性的轻量实现
因为你明确给出0和1的间隔规则,完全可以基于这个规则快速生成flag,几乎无额外开销:
- 先定位所有值为1的位置:
ones_pos = torch.where(idx == 1)[0]
- 初始化全False的
flag,根据每个1的位置标记对应True区间:每个1的位置p,对应的True范围是p+1到p+3(因为1和下一个0间隔3个元素);如果p+4超出tensor长度,就把范围延伸到末尾:
flag = torch.zeros_like(idx, dtype=torch.bool) for p in ones_pos: start = p + 1 end = min(p + 4, len(idx)) flag[start:end] = True
这里的循环仅遍历1的位置,数量远少于tensor总元素数,实际开销可以忽略。
二、不依赖间隔规则的通用无循环实现
如果后续规则变动,也可以用纯张量内置操作实现,完全避免显式循环:
# 标记所有处于1之后的位置(含1本身) after_one = torch.cumsum((idx == 1).int(), dim=0) > 0 # 标记所有处于0之前的位置(含0本身) before_zero = torch.flip(torch.cumsum(torch.flip((idx == 0).int(), dims=(0,)), dim=0) > 0, dims=(0,)) # 计算最终flag:同时满足在1之后、0之前,且本身不是1或0 flag = after_one & before_zero & (idx != 1) & (idx != 0)
该方案利用张量的累积操作,时间复杂度为O(n),效率远高于遍历整个tensor的循环。
三、循环实现的最优方案
如果必须用循环,最优方式是单次遍历,记录当前是否处于需要标记True的状态:
flag = torch.zeros_like(idx, dtype=torch.bool) in_true_range = False for i in range(len(idx)): if idx[i] == 1: in_true_range = True elif idx[i] == 0: in_true_range = False # 仅当处于True区间且当前元素不是1时标记 if in_true_range and idx[i] != 1: flag[i] = True
该循环仅遍历一次tensor,时间复杂度为O(n),是循环实现里的最优选择。
内容的提问来源于stack exchange,提问作者Physics_Student
相关产品推荐
相关产品推荐

