Mean-Teacher半监督U-Net中Consistency Loss失效且总损失上升问题
半监督Mean-Teacher(双U-Net)训练异常:一致性损失未反向传播、总损失上升的修复方案
问题概述
在半监督训练场景下采用双U-Net构建Mean-Teacher架构,训练时出现以下异常:
- 一致性损失(Consistency Loss)未正常反向传播
- 总损失持续上升
已确认所有损失的required_grad=True,现结合提供的训练函数与模型创建代码,给出针对性修复方案。
核心问题诊断
- 模型引用错误:训练函数中使用全局
model而非类实例的self.model,导致优化器更新的参数与实际训练的模型不匹配 - EMA模型模式错误:将EMA模型设置为
train()模式,Mean-Teacher架构中EMA模型应保持eval()模式(仅通过EMA更新参数,不进行反向传播训练) - 不确定性计算逻辑错误:未标注数据的不确定性计算错误地使用了标签(
label_batch[self.labeled_bs:]),半监督场景下未标注数据无有效标签,此操作会导致权重计算完全失效 - 一致性损失的权重与梯度传递隐患:部分计算步骤未明确梯度传递路径,且一致性权重的调度逻辑可能不合理
修复后的代码示例
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
相关产品推荐
相关产品推荐

