【问题标题】:How to handle odd resolutions in Unet architecture PyTorch如何在 Unet 架构 PyTorch 中处理奇数分辨率
【发布时间】:2021-05-07 18:09:57
【问题描述】:

我正在 PyTorch 中实现基于 U-Net 的架构。在火车时间,我有大小为256x256 的补丁,这不会造成任何问题。但是在测试时,我有全高清图像 (1920x1080)。这会导致跳过连接期间出现问题。

下采样1920x1080 3 次得到240x135。如果我再下采样一次,则分辨率变为120x68,上采样时会得到240x136。现在,我无法连接这两个特征图。我该如何解决这个问题?

PS:我认为这是一个相当普遍的问题,但我没有得到任何解决方案,甚至在网络上的任何地方都没有提到这个问题。我错过了什么吗?

【问题讨论】:

  • 您是否尝试过使用torch.nn.MaxPool2d 沿每个维度使用不同的因子进行下采样?您可以使用修复 kernel size = (8, 5) 这将给您 240 x 216 然后您可以填充数组以满足所需的大小 256 x 256 而不会过多地扭曲图像。
  • PS:我建议Maxpooling,但也可以是AveragePooling
  • 没有。我不能用那个。实际上,我正在为我的研究对 Partial ConvNet 论文进行基准测试。不知道能不能这样修改架构。
  • 请问您为什么认为这行不通?或者仅仅是这样一个事实,它意味着对架构进行调整。谢谢
  • 如果我的操作正确,那将是跳过连接。您不会对架构本身进行调整,您将在进入网络之前对输入数据添加预处理步骤。现在,如果您的目标是图像分割,而不是下采样,您还可以尝试将高清图片分割成256 x 256 块,这样您就不会影响分辨率。

标签: python image-processing deep-learning pytorch hourglass


【解决方案1】:

在解码过程中经常涉及跳过连接的分段网络中,这是一个非常常见的问题。网络通常(取决于实际架构)要求输入大小的边长为最大步幅的整数倍(8、16、32 等)。

主要有两种方式:

  1. 将输入调整为最接近的可行大小。
  2. 将输入填充到下一个更大的可行大小。

我更喜欢 (2),因为 (1) 会导致所有像素的像素级别发生微小变化,从而导致不必要的模糊。请注意,我们通常需要在这两种方法中恢复原始形状。

我最喜欢的代码 sn-p 用于此任务(高度/宽度的对称填充):

import torch
import torch.nn.functional as F

def pad_to(x, stride):
    h, w = x.shape[-2:]

    if h % stride > 0:
        new_h = h + stride - h % stride
    else:
        new_h = h
    if w % stride > 0:
        new_w = w + stride - w % stride
    else:
        new_w = w
    lh, uh = int((new_h-h) / 2), int(new_h-h) - int((new_h-h) / 2)
    lw, uw = int((new_w-w) / 2), int(new_w-w) - int((new_w-w) / 2)
    pads = (lw, uw, lh, uh)

    # zero-padding by default.
    # See others at https://pytorch.org/docs/stable/nn.functional.html#torch.nn.functional.pad
    out = F.pad(x, pads, "constant", 0)

    return out, pads

def unpad(x, pad):
    if pad[2]+pad[3] > 0:
        x = x[:,:,pad[2]:-pad[3],:]
    if pad[0]+pad[1] > 0:
        x = x[:,:,:,pad[0]:-pad[1]]
    return x

一个测试sn-p:

x = torch.zeros(4, 3, 1080, 1920) # Raw data
x_pad, pads = pad_to(x, 16) # Padded data, feed this to your network 
x_unpad = unpad(x_pad, pads) # Un-pad the network output to recover the original shape

print('Original: ', x.shape)
print('Padded: ', x_pad.shape)
print('Recovered: ', x_unpad.shape)

输出:

Original:  torch.Size([4, 3, 1080, 1920])
Padded:  torch.Size([4, 3, 1088, 1920])
Recovered:  torch.Size([4, 3, 1080, 1920])

参考:https://github.com/seoungwugoh/STM/blob/905f11492a6692dd0d0fa395881a8ec09b211a36/helpers.py#L33

【讨论】:

  • 谢谢!填充似乎是更好的选择。
猜你喜欢
  • 1970-01-01
  • 1970-01-01
  • 1970-01-01
  • 1970-01-01
  • 2021-10-17
  • 2014-11-05
  • 1970-01-01
  • 2012-01-29
  • 2011-11-11
相关资源
最近更新 更多