【问题标题】:Dimension mismatch between output mask and original mask in dice loss in Semantic segmentation语义分割中骰子损失中输出掩码和原始掩码之间的尺寸不匹配
【发布时间】:2021-09-23 09:12:32
【问题描述】:

我正在做多类语义分割(4类+背景)。我的掩码维度是 (256, 256, 3),输出掩码维度是 (256, 256, 5)。我拿了5,因为这是班级的数量。

骰子损失函数

inputs = inputs.view(-1)
targets = targets.view(-1)
        
intersection = (inputs * targets).sum() ---> error                       
dice = (2.*intersection + smooth)/(inputs.sum() + targets.sum() + smooth)  
        
return 1 - dice

我应该怎么做才能使两个维度相同?掩码是从 TIF 文件中提取的。

我在下面附上了我的面具图片。

【问题讨论】:

  • 你能显示错误信息吗? inputstargets 的形状是什么。另外,“我的掩码维度是 (256, 256, 3),输出掩码维度是 (256, 256, 5)” 这两个张量相对于 inputs 和 @ 987654327@?
  • @Ivan 我正在做语义分割。输入图像大小为 (256, 256, 3),模型输出为 (256, 256, 5),因为图像中有 5 个标签。目标是大小为 (256, 256, 3) 的图像的掩码,即问题中的图像。
  • 我已经提供了答案,见下文。

标签: deep-learning computer-vision pytorch image-segmentation


【解决方案1】:

我相信您必须首先对目标掩码进行一次热编码。 我建议你先阅读这篇好文章,以便更好地掌握语义分割的所有细节https://www.jeremyjordan.me/semantic-segmentation/

确保预测和目标形状匹配,无需使用view(-1) 展平张量。

另外作为个人建议,Pytorch 张量优先使用通道。

【讨论】:

  • 是的,这是我如何匹配这 2 个形状的问题,因为目标蒙版有 3 个维度,而预测的蒙版由于 5 个标签而有 5 个维度。 Yess 频道排名第一
  • 如果我理解正确,您必须先对目标掩码进行一次热编码。如果您有一个当前像素值等于类的目标图像(例如,在 0-4 范围内),您必须将其转换为形状为 5xWxH 的目标图像,每个像素值都是 0-1(按类对其进行一次热编码) .你现在在目标蒙版的通道维度中拥有的 3 我敢打赌这只是保存到 .png 的警告。
【解决方案2】:

我假设您显示的目标分割是一个 RGB 编码的地图。您希望将此 3 通道图像转换为 1 通道标签图。

假设seg 是您的真实分割图,形状为(b, 3, h, w)。标签到颜色的映射可以任意设置为:

colors = torch.FloatTensor([[0, 0, 0],
                            [1, 1, 0],
                            [1, 0, 0],
                            [0, 1, 0],
                            [0, 0, 1]])

为每种颜色构造一个匹配像素的掩码,并在这些像素位置的新张量中分配相应的标签:

b, _, h, w = seg.shape
gt = torch.zeros(b,1,h,w)
seg_perm = seg.permute(0,2,3,1)

for label, color in enumerate(colors):
    mask = torch.all(seg_perm == color, dim=-1).unsqueeze(1)
    gt[mask] = label

以下面的分割图为例:

>>> seg = tensor([[[[1., 1., 0., 0.],
                    [1., 0., 0., 0.]],

                   [[0., 1., 0., 0.],
                    [0., 1., 0., 1.]],

                   [[0., 0., 0., 0.],
                    [0., 0., 1., 0.]]]])

出于可视化目的:

>>> T.ToPILImage()(seg[0].repeat_interleave(100,2).repeat_interleave(100,1))

生成的标签映射将:

>>> gt
tensor([[[[2., 1., 0., 0.],
          [2., 3., 4., 3.]]]])

【讨论】:

    猜你喜欢
    • 1970-01-01
    • 1970-01-01
    • 2021-05-12
    • 1970-01-01
    • 2018-10-30
    • 2018-10-21
    • 1970-01-01
    • 2020-06-14
    • 1970-01-01
    相关资源
    最近更新 更多