如何基于另一个Torch张量对目标张量按0分隔区间求和
PyTorch 按标记张量分段求和的向量化实现
需求说明
给定两个等长Torch张量:
- 张量
a存储浮点型数据 - 张量
b仅包含0和1两类值
要求以b中0的出现位置作为分割点,对相邻两个0之间对应的a的元素进行求和,返回结果张量。
实现思路
通过前缀和运算配合索引切片实现纯向量化计算,无任何显式循环,适配CPU/GPU运算,性能远高于逐元素遍历实现。
完整实现代码
import torch def segment_sum(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor: # 输入校验 assert a.shape == b.shape, "两个输入张量长度必须匹配" # 获取b中所有0值的位置索引 zero_pos = torch.where(b == 0)[0] # 处理开头无0的边界:把张量起始位置0加入索引列表 if len(zero_pos) == 0 or zero_pos[0] != 0: zero_pos = torch.cat([torch.tensor([0], device=zero_pos.device), zero_pos]) # 处理结尾边界:把张量总长度加入索引列表 zero_pos = torch.cat([zero_pos, torch.tensor([a.shape[0]], device=zero_pos.device)]) # 计算a的前缀和,头部补0简化区间差计算 cumsum_a = torch.cat([torch.tensor([0.], device=a.device), torch.cumsum(a, dim=0)]) # 切片求各区间和 return cumsum_a[zero_pos[1:]] - cumsum_a[zero_pos[:-1]]
示例测试
# 示例输入 a = torch.tensor([1., 1., 1., 1., 1., 1., 1., 1., 1., 1.]) b = torch.tensor([0., 1., 1., 1., 0., 1., 1., 1., 1., 0.]) # 调用函数 c = segment_sum(a, b) print(c) # 输出结果:tensor([4., 5., 1.])
实现特性
- 所有操作均为Torch原生向量化运算,自动适配张量所在设备(CPU/GPU)
- 兼容各类边界场景:b开头无0、结尾无0、全0、全1等特殊输入
- 时间复杂度为O(n),空间复杂度为O(n),适合大规模张量运算
内容的提问来源于stack exchange,提问作者Papoudakis Giorgos
相关产品推荐
相关产品推荐

