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

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. 先定位所有值为1的位置:
ones_pos = torch.where(idx == 1)[0]
  1. 初始化全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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.20 23:46:12