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

PyTorch中如何基于分组部分和为一维张量添加零填充

PyTorch实现分组累加阈值填充

实现逻辑

核心思路是利用PyTorch内置的前缀和、二分查找算子替代Python级别的元素遍历,全程以张量运算为主,支持GPU加速,性能远高于逐元素循环的纯Python实现:

  • 预先计算输入一维张量的全局前缀和,避免重复计算分组累加值
  • 逐组通过二分查找快速定位当前组取数的截止位置:找到第一个让分组累加和达到阈值的元素位置
  • 按取数长度把对应元素拷贝到输出张量的分组区间内,剩余位置保留初始化的0值
  • 自动处理输入元素耗尽、组内元素累加未达阈值、单元素超过阈值等边界情况

可运行代码

import torch

# 入参定义
groupsize = 4
totalgroups = 3
partialsum = 15
d1tensor = torch.tensor([3, 12, 5, 5, 5, 4, 11])

# 初始化全0输出张量
output_len = groupsize * totalgroups
output = torch.zeros(output_len, dtype=d1tensor.dtype, device=d1tensor.device)
# 预计算输入张量的前缀和
input_cumsum = d1tensor.cumsum(dim=0)
input_len = len(d1tensor)

current_input_pos = 0  # 下一个待取的输入元素索引
current_output_pos = 0 # 当前分组的输出起始位置

for _ in range(totalgroups):
    if current_input_pos >= input_len:
        current_output_pos += groupsize
        continue
    # 计算当前组需要达到的累加目标值
    base_sum = input_cumsum[current_input_pos - 1] if current_input_pos > 0 else 0
    target_cumsum = base_sum + partialsum
    # 二分查找第一个满足累加和达标的元素位置
    end_pos = torch.searchsorted(input_cumsum, target_cumsum)
    # 计算实际取数长度:不超过分组大小、不超过剩余输入长度
    take_count = min(end_pos - current_input_pos + 1, groupsize, input_len - current_input_pos)
    # 填充对应元素
    output[current_output_pos : current_output_pos + take_count] = d1tensor[current_input_pos : current_input_pos + take_count]
    # 更新游标
    current_input_pos += take_count
    current_output_pos += groupsize

# 输出结果验证
print(output)
# 运行输出:tensor([ 3, 12,  0,  0,  5,  5,  5,  0,  4, 11,  0,  0]),和预期结果一致

说明

  • 代码仅在分组维度做循环,没有逐元素的Python层遍历,分组数远小于元素总数时性能损耗可忽略
  • 算子全部为PyTorch原生实现,可直接运行在CUDA设备上,适配大张量批量计算场景
  • 自动兼容边界场景:输入元素不足时剩余位置自动补0、组内元素总和未达阈值时全量填充后补0、单元素值大于等于阈值时仅填充该元素后补0

内容的提问来源于stack exchange,提问作者dinaa123

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.03 04:31:04