【问题标题】:Multi-channel, 2D mask weights using BCEWithLogitsLoss in Pytorch在 Pytorch 中使用 BCEWithLogitsLoss 的多通道 2D 掩码权重
【发布时间】:2022-07-29 15:55:45
【问题描述】:

我有一组 256x256 的图像,每个图像都标有九个二进制 256x256 掩码。我正在尝试计算 pos_weight 以便使用 Pytorch 对 BCEWithLogitsLoss 进行加权。

我的掩码张量的形状是tensor([1000, 9, 256, 256]),其中 1000 是训练图像的数量,9 是掩码通道的数量(全部编码为 0/1),256 是每个图像边的大小。

为了计算 pos_weight,我将每个掩码中的零相加,然后将该数字除以每个掩码中所有零的总和(遵循建议的 here。):

(masks[:,channel,:,:]==0).sum()/masks[:,channel,:,:].sum()

计算每个遮罩通道的权重会提供一个形状为tensor([9]) 的张量,这对我来说似乎很直观,因为我想要为九个遮罩通道中的每一个通道设置一个 pos_weight 值。但是,当我尝试拟合我的模型时,我收到以下错误消息:

RuntimeError: The size of tensor a (9) must match the size of
tensor b (256) at non-singleton dimension 3

此错误消息令人惊讶,因为它表明权重需要是图像一侧的大小,而不是遮罩通道的数量。 pos_weight 应该是什么形状,我如何指定它应该为遮罩通道而不是图像像素提供权重?

【问题讨论】:

    标签: python deep-learning pytorch loss-function weighted


    【解决方案1】:

    TLDR;这是一个广播问题,令人惊讶的是 PyTorch 的 nn.BCEWithLogitsLossF.binary_cross_entropy_with_logits 没有处理。实际上可能值得发布一个链接到此 SO 线程的 Github 问题,以通知开发人员这种不良行为。

    nn.BCEWithLogitsLoss的文档页面中,声明了提供的正权重张量pos_weight

    必须是长度等于类数的向量。

    这当然是您所期望的(这是正确的),因为正权重是指为每个单独的类赋予正实例的权重。由于您的预测和目标张量是多维的,因此 PyTorch 似乎无法正确处理。


    总之,这是一个最小的示例,展示了如何绕过此错误,还展示了二进制交叉熵的手动计算,作为参考。

    以下是预测张量和目标张量predlabel 的设置:

    >>> c=2;b=5;h=3;w=3
    >>> pred = torch.rand(b,c,h,w)
    >>> label = torch.randint(0,2, (b,c,h,w), dtype=float)
    

    现在对于正权重的定义,请注意前导单例维度:

    >>> pos_weight = torch.rand(c,1,1) 
    

    在您的情况下,使用您现有的长度为c 的一维张量,您只需为高度和宽度维度解压缩两个额外维度。这意味着执行以下操作:pos_weight = pos_weight[:,None,None]

    使用 logits 函数或其 oop 等效函数调用 bce:

    >>> F.binary_cross_entropy_with_logits(pred, label, pos_weight=pos_weight).mean()
    

    在纯代码中相当于:

    >>> z = torch.sigmoid(pred)
    >>> bce = -(pos_weight*label*torch.log(z) + (1-label)*torch.log(1-z))
    

    请注意,如果 class 维度在您的预测和目标张量中是最后一个,则内置函数将具有所需的行为(没有错误消息)。

    >>> pos_weight = torch.rand(c)
    >>> F.binary_cross_entropy_with_logits(
    ...    pred.transpose(1,-1), 
    ...    label.transpose(1,-1), 
    ...    pos_weight=pos_weight)
    

    换句话说,我们正在应用格式为NHWC 的函数,这意味着格式为Cpos_weight 可以正确相乘。所以上面的结果有效地产生了相同的结果:

    >>> F.binary_cross_entropy_with_logits(
    ...    pred, 
    ...    label, 
    ...    pos_weight=pos_weight[:,None,None])
    

    您可以在BCEWithLogitsLoss in another thread here 中阅读有关pos_weight 的更多信息

    【讨论】:

      猜你喜欢
      • 2019-11-23
      • 2019-05-01
      • 2020-04-27
      • 2021-07-06
      • 1970-01-01
      • 2020-08-18
      • 2018-05-06
      • 1970-01-01
      • 1970-01-01
      相关资源
      最近更新 更多