如何仅对5D张量的最后两个维度进行上采样?
问题分析与解决方案
问题根源
你的代码出现该问题的核心原因有两点:
- 使用了
trilinear插值模式:该模式专门针对5D张量(形状为(N, C, D, H, W)),会同时对**最后三个维度(D, H, W)**执行缩放操作; - 传入的
scale_factor=2是标量:会让所有被插值的维度统一放大2倍,导致原本(2,4,3,10,20)的张量中,D维度(3→6)、H维度(10→20)、W维度(20→40)全部被缩放,不符合你仅保留D维度不变的需求。
解决方法
方法一:指定维度专属缩放倍数(推荐)
直接将scale_factor改为元组,明确指定每个维度的缩放倍数——D维度设为1(保持不变),H、W维度设为2(放大2倍),同时保留trilinear模式即可。
修改后的完整代码:
import torch import torch.nn as nn from torch.nn.functional import interpolate class Upsample(nn.Module): def __init__(self, scale_factor, mode, align_corners=False): super().__init__() # 补全父类初始化,避免潜在的Module注册问题 self.interp = interpolate self.scale_factor = scale_factor self.mode = mode self.align_corners = align_corners def forward(self, x): x = self.interp(x, scale_factor=self.scale_factor, mode=self.mode, align_corners=self.align_corners) return x class Main(nn.Module): def __init__(self): super(Main, self).__init__() # 用元组指定D→1倍(不变),H→2倍,W→2倍 self.upsample = Upsample(scale_factor=(1,2,2), mode='trilinear') def forward(self, x): x = self.upsample(x) return x # 测试验证 x = torch.randn(2,4,3,10,20) model = Main() output = model(x) print(output.shape) # 输出 torch.Size([2, 4, 3, 20, 40])
方法二:转换为4D张量使用bilinear插值
如果更倾向于使用专门针对空间维度(H,W)的bilinear插值,可以先将5D张量的通道维度(C)和深度维度(D)合并,转为4D张量后执行插值,最后再拆分回原维度结构。
修改后的完整代码:
import torch import torch.nn as nn from torch.nn.functional import interpolate class Upsample(nn.Module): def __init__(self, scale_factor, mode, align_corners=False): super().__init__() # 补全父类初始化 self.interp = interpolate self.scale_factor = scale_factor self.mode = mode self.align_corners = align_corners def forward(self, x): if x.dim() == 5: # 5D张量转4D:合并C和D维度为新的通道维度 N, C, D, H, W = x.shape x = x.view(N, C*D, H, W) # bilinear插值仅缩放H、W维度 x = self.interp(x, scale_factor=self.scale_factor, mode=self.mode, align_corners=self.align_corners) # 4D张量转5D:拆分回原维度结构 new_H, new_W = x.shape[-2:] x = x.view(N, C, D, new_H, new_W) else: x = self.interp(x, scale_factor=self.scale_factor, mode=self.mode, align_corners=self.align_corners) return x class Main(nn.Module): def __init__(self): super(Main, self).__init__() self.upsample = Upsample(scale_factor=2, mode='bilinear') def forward(self, x): x = self.upsample(x) return x # 测试验证 x = torch.randn(2,4,3,10,20) model = Main() output = model(x) print(output.shape) # 输出 torch.Size([2, 4, 3, 20, 40])
内容的提问来源于stack exchange,提问作者dtr43
相关产品推荐
相关产品推荐

