如何用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函数
- 轴适配:将目标轴转置到0轴统一处理,避免为不同轴编写重复逻辑。
- 前缀和计算:
torch.cumsum是PyTorch优化后的向量操作,比循环快几个数量级。 - 窗口求和逻辑:
- 前
window_size个元素对应窗口从1到指定大小的逐步扩大,直接取前缀和即可。 - 后续元素通过前缀和的差值,快速得到滑动窗口内的总和。
- 前
RollingMax函数
- 格式适配:调整张量形状为
max_pool1d要求的(batch, channels, seq_len)格式。 - 填充处理:前部填充负无穷,确保前
window_size-1个位置的池化窗口只包含前面的有效元素,实现逐步扩大窗口的效果。 - 池化运算:
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
相关产品推荐
相关产品推荐

