如何高效计算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
相关产品推荐
相关产品推荐

