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

如何用PyTorch高效实现张量的滑动窗口切片?

高效创建滑动窗口张量的PyTorch方法

你可以用PyTorch内置的**unfold**方法直接实现矢量化的滑动窗口切片,完全替代串行的列表推导式,速度提升非常明显。另外也可以用as_strided手动构造视图,但unfold更安全易用,优先推荐。

方法1:使用torch.Tensor.unfold(推荐)

unfold是PyTorch专门为滑动窗口切片设计的API,支持指定维度、窗口大小和步长,底层是矢量化实现,不管CPU还是GPU都能高效运行。

对应你的需求,代码如下:

# 确保n = w.shape[-1] - self.l_c + 1(窗口数量的正确计算方式)
weight = w.unfold(dimension=-1, size=self.l_c, step=1)
  • 输出形状:假设原张量w的形状是(a, b, c, d),那么unfold后的形状为(a, b, c, n, self.l_c),其中n是窗口数量。
  • 如果需要和原列表推导式的输出结构(n个(a,b,c,self.l_c)的张量)一致,可以用torch.unbind拆分:
    weight_list = torch.unbind(weight, dim=3)
    
    但除非你真的需要拆分后的列表,否则直接使用unfold返回的张量做后续运算效率更高,避免不必要的拆分操作。

方法2:使用torch.as_strided(灵活但需谨慎)

as_strided可以手动构造张量的视图,通过指定新形状和跨步来实现滑动窗口。这种方法不复制数据,但需要准确计算跨步和形状,否则会出现越界错误。

代码示例:

original_shape = w.shape
dim_size = original_shape[-1]
# 确保n的计算正确:n = dim_size - self.l_c + 1
target_shape = original_shape[:-1] + (n, self.l_c)
# 构造跨步:原张量的跨步 + 最后一维的跨步(窗口每次移动1步)
strides = w.stride() + (w.stride()[-1],)
weight = torch.as_strided(w, size=target_shape, stride=strides)

这个方法得到的结果和unfold完全一致,但需要自己保证n的取值正确,否则会访问到张量外的内存,引发错误。

为什么比列表推导式快?

列表推导式是Python层面的循环,每次切片虽然可能创建视图,但循环本身是串行的,无法利用PyTorch的矢量化加速。而unfold和as_strided都是底层C++实现的矢量化操作,能一次性完成所有窗口的切片,效率提升几个数量级。

内容的提问来源于stack exchange,提问作者Inyoung Kim 김인영

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.11 05:20:32