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

如何用PyTorch高效实现张量的滚动运算?

基于PyTorch实现高效滚动求和与最大值函数

核心实现思路

避免低效的手动循环,利用PyTorch内置的向量/张量操作实现高效滚动运算:

  • 滚动求和:通过前缀和(torch.cumsum)结合切片操作实现,时间复杂度O(n),性能拉满。
  • 滚动最大值:借助一维最大池化(torch.nn.functional.max_pool1d)配合填充逻辑,适配前向窗口逐步扩大的需求。

完整可运行代码

import numpy as np
import torch
import torch.nn.functional as F

X = np.array([1,2,3,4,5,4,3,2,1]).reshape(-1, 1, 1)
X = torch.tensor(X, dtype=torch.float32)

def RollingSum(X, window_size, axis):
    '''Return sum of a rolling window on the tensor over a specified axis'''
    # 统一将目标轴转置到0轴处理,完成后还原
    if axis != 0:
        X = X.transpose(axis, 0)
    
    # 计算前缀和数组
    cumsum = torch.cumsum(X, dim=0)
    roll_sum = torch.zeros_like(X)
    
    # 前window_size个元素直接取前缀和(窗口从1逐步扩大到指定大小)
    roll_sum[:window_size] = cumsum[:window_size]
    # 后续元素用当前前缀和减去window_size步前的前缀和,得到滑动窗口和
    roll_sum[window_size:] = cumsum[window_size:] - cumsum[:-window_size]
    
    # 还原轴顺序
    if axis != 0:
        roll_sum = roll_sum.transpose(axis, 0)
    
    return roll_sum

def RollingMax(X, window_size, axis):
    '''Return max of a rolling window on the tensor over a specified axis'''
    # 统一将目标轴转置到最后一维,适配max_pool1d输入格式
    original_shape = X.shape
    if axis != 0:
        X = X.transpose(axis, -1)
    
    # 调整为(batch, channels, seq_len)格式,兼容池化操作
    batch_channels = X.shape[:-1]
    seq_len = X.shape[-1]
    X_reshaped = X.reshape(-1, 1, seq_len)
    
    # 在前部填充window_size-1个负无穷,让前window_size-1个位置的池化只包含已有元素
    pad = (window_size - 1, 0)
    X_padded = F.pad(X_reshaped, pad, mode='constant', value=-torch.inf)
    
    # 滑动窗口最大池化,步长为1
    max_pooled = F.max_pool1d(X_padded, kernel_size=window_size, stride=1)
    
    # 还原原张量形状与轴顺序
    max_pooled = max_pooled.reshape(*batch_channels, seq_len)
    if axis != 0:
        max_pooled = max_pooled.transpose(axis, -1)
    
    return max_pooled

# 测试滚动求和
Xroll_sum = RollingSum(X, window_size=3, axis=0)
print("滚动求和结果:")
print(Xroll_sum)

# 测试滚动最大值
Xroll_max = RollingMax(X, window_size=3, axis=0)
print("\n滚动最大值结果:")
print(Xroll_max)

代码细节说明

RollingSum函数

  1. 轴适配:将目标轴转置到0轴统一处理,避免为不同轴编写重复逻辑。
  2. 前缀和计算:torch.cumsum是PyTorch优化后的向量操作,比循环快几个数量级。
  3. 窗口求和逻辑:
    • 前window_size个元素对应窗口从1到指定大小的逐步扩大,直接取前缀和即可。
    • 后续元素通过前缀和的差值,快速得到滑动窗口内的总和。

RollingMax函数

  1. 格式适配:调整张量形状为max_pool1d要求的(batch, channels, seq_len)格式。
  2. 填充处理:前部填充负无穷,确保前window_size-1个位置的池化窗口只包含前面的有效元素,实现逐步扩大窗口的效果。
  3. 池化运算:max_pool1d是CUDA加速的内置操作,比手动遍历窗口求最大值效率高得多。

输出验证

运行代码后将得到与示例完全一致的结果:

  • 滚动求和输出:
tensor([[[ 1.]],
        [[ 3.]],
        [[ 6.]],
        [[ 9.]],
        [[12.]],
        [[13.]],
        [[12.]],
        [[ 9.]],
        [[ 6.]]])
  • 滚动最大值输出:
tensor([[[1.]],
        [[2.]],
        [[3.]],
        [[4.]],
        [[5.]],
        [[5.]],
        [[5.]],
        [[4.]],
        [[3.]]])

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.07 10:25:35