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

使用MirroredStrategy训练的UNet模型加载后性能骤降求助

问题

我使用tf.distribute.MirroredStrategy在双NVIDIA 4090的多GPU环境下训练了一个UNet二值分割模型。训练过程中模型表现正常,dice loss从初始约0.8降至0.1。我通过ModelCheckpoint回调保存训练过程中的最优模型,但从.h5文件加载模型后,预测效果极差(仅分割出少量随机像素),即使使用训练时曾成功预测的验证集数据也是如此。在切换到多GPU/MirroredStrategy配置前未出现该问题,我尝试过保存完整模型和仅保存权重两种方式,均存在此问题,请问可能是什么原因导致的?

以下是我的训练函数:

strategy = tf.distribute.MirroredStrategy()
with strategy.scope():
  model = UNet(
               (image_width, image_height, 3), 
               batchnorm=True, 
               start_ch=start_channel_count, 
               depth=layer_count,     
               residual=use_residual)
  model.compile(optimizer=Adam(learning_rate=learning_rate), 
                loss=dice_coef_loss, metrics= 
                  [tf.keras.metrics.BinaryAccuracy(), 
                  tf.keras.metrics.MeanIoU(num_classes=2)])

def scheduler(epoch, lr):
  if epoch < 10:
      return lr
  else:
      return lr * tf.math.exp(-0.1)
mc = ModelCheckpoint(os.path.join(run_dir_path, "model.h5"), 
                     monitor='val_loss', verbose=1, 
                     save_best_only=True, 
                     save_weights_only = True)

history = model.fit(train_data_generator,
                    validation_data=validation_data_generator,
                    callbacks=[mc], 
                    validation_steps=
                      math.ceil(validation_count/batch_size),
                    steps_per_epoch=
                      math.ceil(train_count / batch_size),
                    epochs=100)

训练后直接评估模型时表现正常:

model.evaluate(validation_data, 
               steps=math.ceil(validation_count / batch_size))

但加载.h5文件中的权重后再评估,性能就变得很差:

model.load_weights(os.path.join(run_dir_path, "model.h5"))
model.evaluate(validation_data, 
               steps=math.ceil(validation_count / batch_size))
原因分析与解决方案
  • 模型加载未在策略作用域内重建
    用MirroredStrategy训练时,模型是在strategy.scope()下创建的,权重会适配多GPU的分布式结构。如果加载权重前没在相同的策略作用域内重建完全一致的模型,直接加载会导致权重不匹配,引发预测异常。
    解决:加载权重前,必须先在strategy.scope()内重新定义、编译和训练时完全相同的模型,再执行load_weights。

  • BatchNormalization层运行模式错误
    训练时BN层处于训练模式,会更新均值和方差;加载模型后预测时,如果没手动设置training=False,或者多GPU环境下BN层的状态未正确恢复,会导致BN层使用错误的统计量,输出乱序结果。
    解决:预测时明确调用model.predict(x, training=False);若不需要继续训练,加载模型后可将所有BN层的trainable设为False。同时确保保存模型时BN层的移动均值、方差被正确存储。

  • 多GPU权重的保存/加载适配问题
    MirroredStrategy训练的模型权重是分布式的(每个GPU一份副本),直接保存的权重可能包含多份变量,单GPU环境加载时会出现维度不匹配。
    解决:训练结束后,先通过model = strategy.unwrap_model(model)提取单GPU基础模型,再保存权重;加载时直接用单GPU模型加载,无需再使用策略作用域。

  • 数据预处理不一致
    即使是同一验证集,若加载模型后预测时的预处理流程和训练时不一致(比如意外开启了数据增强、归一化参数错误),会导致输入数据分布偏离训练时的状态,模型输出异常。
    解决:检查预测阶段的数据生成器或预处理函数,确保和训练时完全一致——比如关闭预测时的数据增强,归一化用的均值、方差和训练阶段保持相同。

内容的提问来源于stack exchange,提问作者Michal Cicatka

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.06 07:45:35