You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

结合混淆矩阵的图像损伤等级回归分割损失函数优化咨询

回归损失函数优化方案

原有损失的核心问题

当前使用的普通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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.30 21:54:03