PyTorch中如何用Unfold/Fold实现特定滑动块提取与矩阵重复?
解决方案
针对你的需求,这里提供两种实现方式,都能准确提取你需要的滑动块并整理成目标形状:
方法一:使用带膨胀参数的unfold(简洁高效)
利用torch.nn.functional.unfold的dilation参数,直接在高度维度设置膨胀间隔,一次性提取所有目标块:
import torch import torch.nn.functional as F # 初始化原始张量 i1 = torch.randn(1, 32, 320, 640) # 配置参数 kernel_h, kernel_w = 4, 16 # 高度取4个间隔80的点,宽度窗口16 stride_h, stride_w = 1, 1 # 高度步长1(遍历i=0~79),宽度步长1(滑动窗口) dilation_h, dilation_w = 80, 1 # 高度膨胀80,实现间隔采样 # 执行unfold操作 unfolded = F.unfold( i1, kernel_size=(kernel_h, kernel_w), stride=(stride_h, stride_w), dilation=(dilation_h, dilation_w) ) # 调整形状到目标格式:[1, 总块数, 32, 4, 16] # unfolded形状是[1, 32*4*16, 80*625],转置后reshape result = unfolded.transpose(1, 2).reshape(1, -1, 32, 4, 16) print(result.shape) # 输出: torch.Size([1, 50000, 32, 4, 16])
说明:
dilation_h=80让高度维度的4个kernel点间隔80采样,正好覆盖i, i+80, i+160, i+240- 总块数为
80(i的取值数) × 625(宽度滑动窗口数)=50000,完全匹配你的需求 transpose(1,2)将窗口数维度移到第二位,再reshape拆分出通道、高度采样数、宽度窗口尺寸
方法二:手动切片+滑动窗口(直观易理解)
先手动提取高度维度的间隔采样点,再对宽度维度做滑动窗口提取,适合需要更精细控制的场景:
import torch import torch.nn.functional as F # 初始化原始张量 i1 = torch.randn(1, 32, 320, 640) # 1. 生成所有高度采样索引:i从0~79,每个i对应4个间隔80的点 base_h_indices = torch.arange(0, 320, 80) # [0,80,160,240] all_h_indices = torch.arange(80).unsqueeze(1) + base_h_indices.unsqueeze(0) # shape [80,4] # 2. 提取对应高度位置的张量:shape [1,32,80,4,640] sampled_tensor = i1[:, :, all_h_indices, :] # 3. 对宽度维度做滑动窗口提取,生成[1,32,80,4,625,16]的张量 window_size = 16 stride = 1 num_windows = sampled_tensor.shape[-1] - window_size + 1 # 使用torch.as_strided创建滑动窗口(避免复制数据,高效) slided_tensor = torch.as_strided( sampled_tensor, shape=sampled_tensor.shape[:-1] + (num_windows, window_size), strides=sampled_tensor.stride() + (sampled_tensor.stride()[-1],) ) # 4. 调整维度到目标格式:[1, 50000, 32,4,16] result = slided_tensor.permute(0, 2, 4, 1, 3, 5).reshape(1, -1, 32, 4, 16) print(result.shape) # 输出: torch.Size([1, 50000, 32, 4, 16])
说明:
- 先通过索引生成器精准获取每个
i对应的4个高度位置 - 用
torch.as_strided在宽度维度创建滑动窗口,不会额外占用内存 - 最后通过维度重排和合并,得到和方法一完全一致的结果
内容的提问来源于stack exchange,提问作者LookerHan9327
相关产品推荐
相关产品推荐

