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

如何高效计算PyTorch张量每行中连续零的最大长度?

如何高效计算PyTorch张量每行中连续零的最大长度?

我明白你想要在PyTorch中实现和现有Numpy版本同样高效的“每行最长连续零”计算逻辑,先再明确下你的需求示例:

给定输入张量:

import torch
input = torch.tensor([[0, 1, 0, 0, 0, 1],[0, 0, 1, 0, 1, 0],[1, 0, 0, 0, 0, 0]])

你期望得到输出:tensor([3, 2, 5])

你已经有了高效的Numpy实现,但自己尝试的PyTorch版本性能不佳。其实我们可以直接将你的Numpy思路迁移到PyTorch上,用向量化的PyTorch操作实现,性能完全能和Numpy版本看齐。

迁移后的PyTorch高效实现

下面是和你的Numpy逻辑1:1对齐的PyTorch版本,全程使用向量化操作,避免显式循环以保证性能:

import torch

def max_consecutive_zeros_per_row(input_tensor):
    # 1. 在每行首尾填充1,对应Numpy的np.pad
    padded = torch.nn.functional.pad(input_tensor, (1, 1), mode='constant', value=1)
    
    # 2. 计算每行的相邻元素差分,对应np.diff(axis=1)
    diffs = torch.diff(padded, dim=1)
    
    # 3. 定位连续零段的起始(1→0,差分=-1)和结束(0→1,差分=1)位置
    start_rows, start_cols = torch.where(diffs == -1)
    end_rows, end_cols = torch.where(diffs == 1)
    
    # 4. 计算每个零段的长度
    run_lengths = end_cols - start_cols
    
    # 5. 对每行取零段的最大长度,对应Numpy的np.maximum.at
    # 使用scatter_reduce实现向量化的逐行取最大,PyTorch 1.12+支持
    max_zeros = torch.zeros(
        input_tensor.size(0), 
        dtype=torch.long, 
        device=input_tensor.device
    )
    max_zeros = max_zeros.scatter_reduce(
        0, start_rows, run_lengths, 
        reduce='amax', include_self=False
    )
    
    return max_zeros

# 测试示例
input_tensor = torch.tensor([[0, 1, 0, 0, 0, 1],[0, 0, 1, 0, 1, 0],[1, 0, 0, 0, 0, 0]])
print(max_consecutive_zeros_per_row(input_tensor))  # 输出: tensor([3, 2, 5])

代码细节解释

每一步都和你的Numpy实现逻辑完全对应,保证思路一致性:

  • 填充操作:用torch.nn.functional.pad替代np.pad,给每行首尾加1,确保行首/行尾的连续零段能被正确捕获(比如全零行、行尾全零的情况)
  • 差分计算:torch.diff(dim=1)和np.diff(axis=1)行为一致,通过差分结果快速定位零段的起止点
  • 起止点定位:torch.where替代np.where,获取所有零段的行号和列号索引
  • 逐行取最大:scatter_reduce是PyTorch的向量化聚合操作,替代np.maximum.at,避免循环,保证大规模数据下的效率

兼容性与边界情况

  • 旧版本PyTorch兼容:如果你的PyTorch版本低于1.12(scatter_reduce在1.12版本引入),可以用以下循环替代(行数较多时性能略降,但比逐元素循环好):
    # 旧版本PyTorch的替代方案
    max_zeros = torch.zeros(input_tensor.size(0), dtype=torch.long, device=input_tensor.device)
    for row_idx in torch.unique(start_rows):
        row_mask = start_rows == row_idx
        max_len = run_lengths[row_mask].max() if row_mask.any() else 0
        max_zeros[row_idx] = max_len
    
  • 边界情况测试:
    # 测试全零行和全一行
    test_tensor = torch.tensor([[0,0,0], [1,1,1], [0,1,0]])
    print(max_consecutive_zeros_per_row(test_tensor))  # 输出: tensor([3, 0, 1])
    

这个实现完全沿用了你验证过的高效逻辑,在GPU环境下性能会优于Numpy,CPU环境下也能和Numpy版本持平,应该能满足你的需求。

备注:内容来源于stack exchange,提问作者Paulo Nascimento

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.14 16:02:59