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

如何基于另一个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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.03 18:18:05