结合混淆矩阵的图像损伤等级回归分割损失函数优化咨询
回归损失函数优化方案
原有损失的核心问题
当前使用的普通MSE损失默认所有像素的误差权重一致,由于数据集存在100:1的类别不平衡,占比99%的无损伤像素的误差会主导损失更新,模型会优先拟合无损伤像素的特征,完全忽略占比极低的1-5级损伤像素的误差,最终导致损伤等级的预测和标签无相关性。
方案1:加权均方误差改造
不修改原始数据集,仅通过损失加权提高1-5级损伤像素的误差贡献,完整保留无损伤区域的信息,满足损伤旁无损伤区域的参考需求。
首先统计训练集中每个损伤等级(0-5)的像素总数,按如下规则计算每个等级的权重:
某等级权重 = 总像素数 / (6 * 该等级像素数)
该计算方式可以保证每个等级的像素对总损失的贡献权重一致,抵消类别不平衡的影响。如果需要重点提升严重损伤的预测精度,可以额外调高4、5级的权重。
修改后的损失代码如下:
def weighted_mse_loss(pred, mask, class_weights=[0.17, 20, 20, 20, 20, 20]): # pred: 模型回归头原始输出 # mask: 真实损伤等级,已归一化到[0,1]区间(0对应0级,1对应5级) pred = torch.sigmoid(pred) # 将归一化的mask转回0-5的整数等级,用于匹配权重 mask_level = torch.round(mask * 5).long() # 生成逐像素权重 class_weight_tensor = torch.tensor(class_weights, device=pred.device) pixel_weight = torch.take(class_weight_tensor, mask_level).unsqueeze(1).unsqueeze(2) # 计算加权MSE mse = F.mse_loss(pred, mask, reduction='none') weighted_mse = (pixel_weight * mse).mean() return weighted_mse
代码中的class_weights为示例值,可根据你自己的数据集像素统计结果修改。
方案2:融合结构权重的回归损失
可以复用你二分类任务中已经验证有效的边缘加权逻辑,进一步提升损伤区域边缘像素的损失权重,这类边缘像素对检查人员的参考价值更高:
def structured_weighted_mse_loss(pred, mask): pred_sig = torch.sigmoid(pred) # 生成二分类受损掩码,计算结构权重 binary_mask = (mask > 0).float() edge_weight = 1 + 5 * torch.abs(F.avg_pool2d(binary_mask, kernel_size=31, stride=1, padding=15) - binary_mask) # 生成类别平衡权重 class_weights = [0.17, 20, 20, 20, 20, 20] mask_level = torch.round(mask *5).long() class_weight_tensor = torch.tensor(class_weights, device=pred.device) class_weight = torch.take(class_weight_tensor, mask_level).unsqueeze(1).unsqueeze(2) # 合并两类权重 total_weight = class_weight * edge_weight # 计算加权MSE mse = F.mse_loss(pred_sig, mask, reduction='none') weighted_mse = (total_weight * mse).sum(dim=(2,3)) / total_weight.sum(dim=(2,3)) return weighted_mse.mean()
方案3:多任务联合训练(推荐)
你当前二分类任务的IoU已经超过0.6,效果较为稳定,可以采用多任务训练的方式,让模型共享 backbone 特征,同时输出二分类头(判断是否受损)和回归头(预测损伤等级),损失为两个任务损失的加权和:
def multi_task_total_loss(pred_cls, pred_reg, mask): # pred_cls: 二分类头原始输出 # pred_reg: 回归头原始输出 binary_mask = (mask > 0).float() # 二分类损失复用你已有的structure_loss loss_cls = structure_loss(pred_cls, binary_mask) # 回归损失使用方案2的结构化加权MSE loss_reg = structured_weighted_mse_loss(pred_reg, mask) # 可根据效果调整回归损失的权重 return loss_cls + 2 * loss_reg
训练监控注意事项
训练过程中不要仅参考全局MSE指标,该指标会被占比极高的无损伤像素主导,没有参考价值。需要单独统计损伤区域(mask>0的像素)的MSE、MAE以及等级混淆矩阵,用这类指标判断模型对损伤等级的预测效果。
内容的提问来源于stack exchange,提问作者sixtytrees
相关产品推荐
相关产品推荐

