PyTorch是否支持类似NumPy stride_tricks的功能?如何实现明确行为的重叠索引?
在PyTorch中实现类似NumPy
stride_tricks.as_strided的安全替代方案 你提到在NumPy中可以通过as_strided修改步长实现滑动窗口这类实用功能,比如以下示例:
import numpy as np from numpy.lib.stride_tricks import as_strided a = np.arange(15).reshape(3,5) print(a) # [[ 0 1 2 3 4] # [ 5 6 7 8 9] # [10 11 12 13 14]] b = as_strided(a, shape=(3,3,3), strides=(a.strides[-1],)+a.strides) print(b) # [[[ 0 1 2] # [ 5 6 7] # [10 11 12]] # [[ 1 2 3] # [ 6 7 8] # [11 12 13]] # [[ 2 3 4] # [ 7 8 9] # [12 13 14]]] # 计算3x3窗口的和 print(b.sum(axis=(1,2))) # [54 63 72]
而PyTorch的torch.as_strided因内存重叠索引行为未定义,不适合这类场景。以下是两种官方支持、文档明确的替代方案:
方法1:使用torch.nn.Unfold(适合滑动窗口提取)
Unfold是PyTorch专为提取张量滑动窗口设计的API,所有行为都有明确规范,无未定义风险。示例代码如下:
import torch a = torch.arange(15).reshape(3, 5) # 在列维度(dim=1)提取窗口:窗口大小3,步长1 unfolded = a.unfold(dimension=1, size=3, step=1) # unfolded形状为(3, 3, 3),与NumPy示例中的b完全对应 print(unfolded) # 计算窗口求和 window_sums = unfolded.sum(dim=(1, 2)) print(window_sums) # tensor([54, 63, 72])
方法2:使用torch.sliding_window_view(PyTorch 1.10+,接近NumPy视图用法)
sliding_window_view是官方提供的安全滑动窗口视图API,它会自动验证参数合法性,保证行为可预期,且无需额外内存副本(条件允许时)。示例代码如下:
import torch a = torch.arange(15).reshape(3, 5) # 创建滑动窗口视图:窗口大小(3,3),滑动步长(1,1) window_view = torch.sliding_window_view(a, window_size=(3, 3), step=(1, 1)) # 移除行维度滑动产生的冗余维度(因原张量仅3行,行方向仅能滑动1次) window_view = window_view.squeeze(0) print(window_view) # 计算窗口求和 window_sums = window_view.sum(dim=(1, 2)) print(window_sums) # tensor([54, 63, 72])
方案选择
- 如果是为后续神经网络层准备输入特征,优先用
Unfold,它是CV领域提取滑动窗口的标准工具; - 如果仅需要创建窗口视图进行计算,
sliding_window_view的用法更贴近NumPy的as_strided,且安全可靠。
内容的提问来源于stack exchange,提问作者Fırat Kıyak
相关产品推荐
相关产品推荐

