PyTorch中基于Bx3xHxW张量高效构建右移图像栈的方法
实现方案
核心思路完全基于向量化操作,无任何for循环,GPU/CPU适配性强、内存开销极低:先在原始图像宽度方向左侧补F-1个0,再通过滑动窗口直接提取所有偏移量对应的结果,最后调整维度得到目标张量。
PyTorch 实现
import torch def generate_shift_stack(img: torch.Tensor, F: int = 64) -> torch.Tensor: B, C, H, W = img.shape # 宽度维度左侧补F-1个0,pad参数顺序:(W左补, W右补, H上补, H下补) padded = torch.nn.functional.pad(img, (F-1, 0, 0, 0), mode='constant', value=0) # 宽度维度滑动取窗口,窗口大小等于原宽度,步长为1,得到形状B×C×H×F×W shifted = padded.unfold(dimension=-1, size=W, step=1) # 调整维度顺序,合并通道和偏移维度得到目标形状 shifted = shifted.permute(0, 1, 3, 2, 4).reshape(B, C*F, H, W) return shifted
正确性验证示例
输入测试张量shape = (1, 1, 1, 3),值为[[[[1,2,3]]]],F=2:
- 补零后张量为
[[[[0,1,2,3]]]] - 滑动窗口提取结果为偏移0次的
[1,2,3]、偏移1次的[0,1,2] - 输出形状为
(1, 2, 1, 3),完全符合规则要求
NumPy 实现
逻辑和PyTorch完全一致,可直接复用:
import numpy as np def generate_shift_stack_np(img: np.ndarray, F: int = 64) -> np.ndarray: B, C, H, W = img.shape # 宽度维度左侧补F-1个0 padded = np.pad(img, pad_width=((0,0), (0,0), (0,0), (F-1, 0)), mode='constant', constant_values=0) # 滑动窗口提取所有偏移结果 shifted = np.lib.stride_tricks.sliding_window_view(padded, window_shape=W, axis=-1) # 调整维度合并通道和偏移维度 shifted = shifted.transpose(0, 1, 3, 2, 4).reshape(B, C*F, H, W) return shifted
方案优势
- 全向量化运算,无任何Python层循环,性能完全拉满
- 滑动窗口操作属于视图操作,不会额外复制冗余数据,内存效率极高
- 兼容任意偏移量
F、任意输入尺寸,无需修改逻辑
内容的提问来源于stack exchange,提问作者Mohit Lamba
相关产品推荐
相关产品推荐

