添加滑动窗口维度至PyTorch数据引发维度不匹配错误
问题背景
原本基于PyTorch的流水线中,DataLoader返回形状为(4,1,192,320)的4D张量,经Unfold、FFN处理后得到(4,15,256)的输出。添加滑动窗口(窗口大小5)后,DataLoader返回(4,5,1,192,320)的5D张量,导致Unfold因维度不匹配报错,现有模型流水线无法兼容。
问题1:不改造现有模型的解决方案
不需要修改模型,只需在数据输入模型前合并批次与窗口维度,处理完成后再还原维度即可适配现有流水线。
具体步骤:
- 将5D张量
(batch_size, window_size, channels, H, W)reshape为(batch_size*window_size, channels, H, W),转换成Unfold支持的4D格式。 - 用原有模型流水线处理该4D张量。
- 将处理后的结果从
(batch_size*window_size, ...)reshape回(batch_size, window_size, ...),恢复窗口维度。
代码示例:
# 获取5D张量 frames = next(iter(dataloader)) print('Raw (5D): ', tuple(frames.shape)) # (4,5,1,192,320) # 合并批次与窗口维度 batch_size, window_size = frames.shape[:2] frames_4d = frames.reshape(batch_size * window_size, *frames.shape[2:]) print('Reshaped to 4D: ', tuple(frames_4d.shape)) # (20,1,192,320) # 原有流水线处理 unfold = torch.nn.Unfold(kernel_size=64, stride=64) unfolded_ = unfold(frames_4d) unfolded = unfolded_.view(unfolded_.size(0), -1, 64, 64) print('Unfolded: ', tuple(unfolded.shape)) # (20,15,64,64) unfolded_reshaped = unfolded.reshape(unfolded.size(0), -1, 64*64) ffn = FFN(64*64, 256, 0.1) ffn_out = ffn(unfolded_reshaped) print('FFN Output (4D): ', tuple(ffn_out.shape)) # (20,15,256) # 还原批次与窗口维度 ffn_out_5d = ffn_out.reshape(batch_size, window_size, *ffn_out.shape[1:]) print('FFN Output (5D): ', tuple(ffn_out_5d.shape)) # (4,5,15,256)
这种方式完全复用原有模型代码,仅在数据输入输出阶段做维度变换,是最轻量化的适配方案。
问题2:最小改动改造模型的方案
如果希望模型直接处理5D张量,可在模型的关键处理步骤(如Unfold、FFN)中封装维度变换逻辑,改动量极小。
改造思路
在模型的forward方法开头,将5D张量的窗口维度与批次维度合并,用原有逻辑处理后再拆分维度,对外保持5D张量的输入输出接口。
改造后的Unfold封装
创建支持5D输入的Unfold wrapper:
class Unfold5D(nn.Module): def __init__(self, kernel_size, stride): super().__init__() self.unfold = nn.Unfold(kernel_size=kernel_size, stride=stride) def forward(self, x): # x shape: (batch, window, channels, H, W) batch, window = x.shape[:2] x_4d = x.reshape(batch*window, *x.shape[2:]) unfolded_4d = self.unfold(x_4d) # 还原维度:(batch*window, channels*kernel_size^2, num_patches) -> (batch, window, channels*kernel_size^2, num_patches) unfolded_5d = unfolded_4d.reshape(batch, window, *unfolded_4d.shape[1:]) return unfolded_5d
改造后的FFN(支持5D输入)
复用原有FFN,仅在forward中增加维度适配:
class FFN5D(nn.Module): def __init__(self, in_dim, out_dim, dropout=0.1): super().__init__() self.ffn = FFN(in_dim, out_dim, dropout) # 复用原有FFN def forward(self, x): # x shape: (batch, window, num_patches, in_dim) batch, window, num_patches = x.shape[:3] x_3d = x.reshape(batch*window, num_patches, x.shape[-1]) ffn_out_3d = self.ffn(x_3d) # 还原维度 ffn_out_5d = ffn_out_3d.reshape(batch, window, num_patches, ffn_out_3d.shape[-1]) return ffn_out_5d
改造后的流水线示例
frames = next(iter(dataloader)) print('Raw (5D): ', tuple(frames.shape)) # (4,5,1,192,320) unfold_5d = Unfold5D(kernel_size=64, stride=64) unfolded_ = unfold_5d(frames) # 调整维度到(4,5,15,64,64) unfolded = unfolded_.permute(0,1,3,2).view(frames.shape[0], frames.shape[1], -1,64,64) print('Unfolded (5D): ', tuple(unfolded.shape)) # (4,5,15,64,64) unfolded_reshaped = unfolded.reshape(unfolded.size(0), unfolded.size(1), -1, 64*64) ffn_5d = FFN5D(64*64, 256, 0.1) ffn_out = ffn_5d(unfolded_reshaped) print('FFN Output (5D): ', tuple(ffn_out.shape)) # (4,5,15,256)
这种方式仅对原有模型做轻量封装,核心逻辑完全复用,同时支持5D张量的直接输入输出。
内容的提问来源于stack exchange,提问作者Mahesha999
相关产品推荐
相关产品推荐

