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
相关产品推荐
相关产品推荐

