BCEWithLogitsLoss多标签分类mask损失计算与epoch损失累计问题
PyTorch BCEWithLogitsLoss多标签任务Epoch损失累计问题解答
问题1:带mask的批次损失计算与累计逻辑
原代码#1位置的写法不正确,问题点和正确逻辑如下:
- 首先默认初始化的
BCEWithLogitsLoss参数reduction='mean',会直接返回0维的标量平均损失,此时和mask矩阵相乘属于广播操作,得到的矩阵每个位置值为「标量损失*对应位置mask值」,再调用.mean()得到的结果完全不代表真实的masked损失,计算值会比真实值偏低。 - 计算带mask的损失时,需要先将损失函数的
reduction参数设为'none',拿到每个样本每个类别维度的独立损失值,乘mask屏蔽无效位置后,仅对有效位置求平均得到单batch的真实平均损失。 - 不建议直接把每个batch的平均损失乘以
outputs.shape[0](即batch size)后累加,原因有两个:一是如果你的mask是标签级的(即同一个样本内部分类别标签无效),单batch的有效计算单元数不是batch size,而是mask值为1的总个数;二是如果dataloader没有设置drop_last=True,最后一个batch的样本数会小于设定的batch size,直接乘固定batch size会累计错误的总损失。 - 正确的累计方式是同时维护两个统计量:整个epoch的有效位置损失总和、整个epoch的有效位置总计数,反向传播用当前batch的平均损失即可,不需要改动反向传播的逻辑。
- 原代码此处存在两个笔误:一是
model.parameters要加括号写成model.parameters()才是可传入优化器的参数列表;二是反向传播时调用的loss未定义,应该用计算得到的batch_loss。
问题2:Epoch平均损失的上报逻辑
原代码#2位置的写法不正确:
- 直接将累计的batch损失和除以dataloader长度(即batch总数),仅在所有batch样本数完全一致、每个batch内有效mask的数量完全相等时才近似准确,一旦存在末位batch样本不足、不同样本有效标签数不同的情况,计算出的epoch损失是有偏的。
- 正确的上报值应该是整个epoch累计的有效位置损失总和,除以整个epoch累计的有效位置总计数,这个值才是整个训练集上的真实平均损失。
修正后参考代码
# 注意初始化损失函数时指定reduction='none',拿到逐元素损失 criterion = torch.nn.BCEWithLogitsLoss(pos_weight=pos_weight, reduction='none') optimizer = torch.optim.Adam(model.parameters(), lr=1e-3, weight_decay=1e-5) for epoch in range(10): # 累计两个统计量:总损失和、有效单元总数 epoch_loss_sum = 0. epoch_valid_count = 0 for inputs, gt_labels, masks in training_dataloader: optimizer.zero_grad() outputs = model(inputs) # 逐元素计算损失,shape和outputs、gt_labels、masks一致:[batch_size, num_classes] per_element_loss = criterion(outputs, gt_labels.float()) # 屏蔽无效位置的损失 masked_loss = per_element_loss * masks # 当前batch的平均损失:仅对有效位置求平均,用于反向传播 batch_loss = masked_loss.sum() / masks.sum() # 累计到epoch总统计量,注意用.item()取出张量数值避免显存泄漏 epoch_loss_sum += masked_loss.sum().item() epoch_valid_count += masks.sum().item() # 反向传播用batch的平均损失即可 batch_loss.backward() optimizer.step() # 计算整个epoch的真实平均损失 epoch_avg_loss = epoch_loss_sum / epoch_valid_count print(f'EPOCH LOSS: {epoch_avg_loss:.3f}')
如果你的mask是样本级(shape为
[batch_size],标记整个样本是否有效),逻辑完全一致,只需要在计算损失时把mask补全维度到[batch_size, 1]保证广播逻辑正确即可。
内容的提问来源于stack exchange,提问作者Sam
相关产品推荐
相关产品推荐

