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

Mean-Teacher半监督U-Net中Consistency Loss失效且总损失上升问题

半监督Mean-Teacher(双U-Net)训练异常:一致性损失未反向传播、总损失上升的修复方案

问题概述

在半监督训练场景下采用双U-Net构建Mean-Teacher架构,训练时出现以下异常:

  • 一致性损失(Consistency Loss)未正常反向传播
  • 总损失持续上升
    已确认所有损失的required_grad=True,现结合提供的训练函数与模型创建代码,给出针对性修复方案。

核心问题诊断

  1. 模型引用错误:训练函数中使用全局model而非类实例的self.model,导致优化器更新的参数与实际训练的模型不匹配
  2. EMA模型模式错误:将EMA模型设置为train()模式,Mean-Teacher架构中EMA模型应保持eval()模式(仅通过EMA更新参数,不进行反向传播训练)
  3. 不确定性计算逻辑错误:未标注数据的不确定性计算错误地使用了标签(label_batch[self.labeled_bs:]),半监督场景下未标注数据无有效标签,此操作会导致权重计算完全失效
  4. 一致性损失的权重与梯度传递隐患:部分计算步骤未明确梯度传递路径,且一致性权重的调度逻辑可能不合理

修复后的代码示例

1. 修正模型创建与引用

def create_model(ema=False):
    # Network definition
    model = UNet(1, 1)
    model = model.cuda()
    if ema:
        for param in model.parameters():
            param.detach_()
        model.eval()  # EMA模型默认设为评估模式
    return model

# 类内部初始化时将模型赋值给实例属性
self.model = create_model()
self.ema_model = create_model(ema=True)
self.criterion = MultiTaskLoss().to(device)
self.optimizer = optim.Adam(self.model.parameters(), lr=1e-4)

2. 修正训练函数逻辑

def _train(self, writer):
    self.model.train()  # 主模型设为训练模式
    self.ema_model.eval()  # EMA模型固定为评估模式,禁止训练模式下的batch norm等行为
    cons = torch.as_tensor(0, dtype=torch.float32, device=device)
    cons_weight = torch.as_tensor(0, dtype=torch.float32, device=device)

    temp = []
    n_batch = 0

    batch_iter = tqdm(enumerate(self.training_DataLoader), 'Training', total=len(self.training_DataLoader), leave=False)
    for i_batch, (x, y) in batch_iter:
        n_batch += 1

        volume_batch, label_batch = x, y[0]
        volume_batch, label_batch = volume_batch.cuda(), label_batch.cuda()
        unlabeled_volume_batch = volume_batch[self.labeled_bs:]

        # 主模型对全批次数据(标注+未标注)的预测
        outputs = self.model(volume_batch)

        # EMA模型对加噪未标注数据的预测(禁用梯度计算)
        noise = torch.clamp(torch.randn_like(unlabeled_volume_batch) * 0.1, -0.2, 0.2)
        ema_inputs = unlabeled_volume_batch + noise
        with torch.no_grad():
            ema_output = self.ema_model(ema_inputs)

        # 基于EMA模型多次预测计算不确定性(修正:使用预测本身的熵,而非未标注数据的标签)
        T = 8
        volume_batch_r = unlabeled_volume_batch.repeat(2, 1, 1, 1)
        stride = volume_batch_r.shape[0] // 2
        preds = torch.zeros([stride * T, 1, 384, 384]).cuda()
        for i in range(T//2):
            ema_inputs_r = volume_batch_r + torch.clamp(torch.randn_like(volume_batch_r) * 0.1, -0.2, 0.2)
            with torch.no_grad():
                preds[2 * stride * i:2 * stride * (i + 1)] = self.ema_model(ema_inputs_r)
        preds = F.softmax(preds, dim=1)
        preds = preds.reshape(T, stride, 1, 384, 384)
        preds = torch.mean(preds, dim=0)
        # 修正:使用预测的熵作为不确定性(半监督场景下未标注数据无标签)
        uncertainty = -1.0 * torch.sum(preds * torch.log(preds + 1e-6), dim=1, keepdim=True)

        weights = F.softmax(1 - uncertainty, dim=0)
        ema_probs = torch.sum(preds * weights, dim=0)
        ema_seg_uncertainty = -1.0 * torch.sum(ema_probs * torch.log2(ema_probs + 1e-6), dim=1, keepdim=True)

        # 计算损失
        supervised_loss = self.criterion(outputs[:self.labeled_bs], label_batch[:self.labeled_bs])

        # 一致性权重调度(建议使用基于迭代次数的调度,而非epoch)
        consistency_weight = get_current_consistency_weight(self.iter_num // 150)
        # 计算一致性损失:主模型未标注部分输出与EMA模型输出的MSE,乘以不确定性权重
        consistency_dist = torch.pow(outputs[self.labeled_bs:] - ema_output, 2)
        consistency_dist = consistency_dist * (1 - ema_seg_uncertainty)
        consistency_dist = torch.mean(consistency_dist)
        consistency_loss = consistency_weight * consistency_dist

        cons += consistency_loss
        cons_weight += consistency_weight

        # 累计损失统计
        if n_batch == 1:
            total_loss_sum = supervised_loss.item() + consistency_loss.item()
            sup_loss_sum = supervised_loss.item()
        else:
            total_loss_sum += supervised_loss.item() + consistency_loss.item()
            sup_loss_sum += supervised_loss.item()

        # 反向传播与优化
        total_loss = supervised_loss + consistency_loss
        self.optimizer.zero_grad()
        total_loss.backward()
        self.optimizer.step()
        # 更新EMA模型参数
        update_ema_variables(self.model, self.ema_model, args.ema_decay, self.iter_num)

        self.iter_num += 1

        logging.info('Epoch %d | iteration %d : total_loss : %f, sup_loss: %f, consistency_loss: %f, cons_dist: %f, cons_weight: %f' %
                     (self.epoch, self.iter_num, total_loss.item(), supervised_loss.item(), consistency_loss.item(), consistency_dist.item(), consistency_weight))

    # 训练后统计
    avg_cons_loss = round(cons.item() / n_batch, 5)
    avg_cons_weight = cons_weight.item() / n_batch
    avg_total_loss = total_loss_sum / n_batch
    avg_sup_loss = sup_loss_sum / n_batch

    print(f" CONSISTENCY LOSS: {avg_cons_loss}, CONSISTENCY WEIGHT: {avg_cons_weight}")
    print(f" TRAIN LOSS: {avg_total_loss}")
    self.task_losses.append([avg_sup_loss, avg_cons_loss, avg_cons_weight])
    self.training_loss.append(avg_total_loss)
    self.consistency_loss.append(round(cons.item() / n_batch, 6))
    self.learning_rate.append(self.optimizer.param_groups[0]['lr'])
    torch.save(self.model.state_dict(), self.model_name)
    print("Model Saved")

关键修复点说明

  • 模型引用修正:将全局model/ema_model替换为类实例的self.model/self.ema_model,确保优化器更新的是当前训练的模型参数
  • EMA模型模式修正:EMA模型始终保持eval()模式,避免训练模式下的batch normalization等组件引入额外噪声,保证输出稳定
  • 不确定性计算修正:恢复使用预测结果的熵作为不确定性指标,半监督场景下未标注数据无有效标签,原逻辑会导致权重计算完全偏离预期
  • 一致性权重调度优化:建议使用基于迭代次数的权重调度(而非epoch),更贴合Mean-Teacher的训练节奏
  • 代码可读性优化:统一损失统计变量命名,避免歧义

内容的提问来源于stack exchange,提问作者MARCO CALABRESE

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.24 07:45:56