Keras模型从检查点重启训练时无法追踪历史val_loss
问题解决:从检查点重启训练后指标记录异常
问题现象
从检查点重启训练时,输出如下:
val_root_mean_squared_error improved from inf to 0.38011
但预期输出应该包含Epoch信息和之前的有效指标值,例如:
Epoch 1: val_root_mean_squared_error improved from 0.583 to 0.38011
或
Epoch 1: val_root_mean_squared_error did not improve from 0.326
核心原因
- 加载完整模型时,
ModelCheckpoint回调的内部状态(如历史最佳指标值)不会被保存,重启后回调会将初始最佳值设为inf(当mode='min'时),导致输出异常。 - 未追踪函数的警告来自自定义损失函数中的
tf.map_fn和LSTM层的动态函数,虽不影响训练,但可通过调整损失函数实现方式消除。
修复步骤
步骤1:手动恢复ModelCheckpoint的最佳指标值
加载模型后,先在验证集上评估当前指标,将结果赋值给回调的best属性,让回调知晓之前的最佳值:
checkpoint_filepath = f'{log_location}/mdl.ckpt' checkpoint = tf.keras.callbacks.ModelCheckpoint( filepath=checkpoint_filepath, save_weights_only=False, monitor=monitor, mode=mode, verbose=1, # 改为1,输出自动包含Epoch信息 save_best_only=True) if os.path.isdir(checkpoint_filepath): print(f'loading model from {checkpoint_filepath}') model = tf.keras.models.load_model(checkpoint_filepath, custom_objects={'loss_fcn': loss_fcn}) # 在验证集上评估,获取当前指标值 val_results = model.evaluate(X_val, y_val, verbose=0) # 找到目标指标在结果中的索引 metric_idx = model.metrics_names.index(monitor) # 手动设置回调的历史最佳值 checkpoint.best = val_results[metric_idx] else: # 无检查点时初始化模型 model = build_bilstm_model(params) model.fit(X_train, y_train, batch_size=params['batch_size'], epochs=epochs, shuffle=True, validation_data=(X_val, y_val), callbacks=[checkpoint, lr_scheduler, tensorboard], verbose=1)
步骤2:优化自定义损失函数,消除未追踪函数警告
将损失函数中的tf.map_fn改为向量化操作,避免使用lambda,让TensorFlow能正确追踪所有函数:
def loss_fcn(y_t, y_p): y_pred = tf.convert_to_tensor(y_p) y_true = tf.cast(y_t, y_pred.dtype) diff = y_pred - y_true # 用向量化操作替代tf.map_fn,避免未追踪函数 abs_diff = tf.abs(diff) neg_mask = diff < 0 pos_mask = diff >= 0 neg_res = tf.math.exp(abs_diff[neg_mask]/3) - 1 pos_res = tf.math.exp(abs_diff[pos_mask]/9) - 1 res = tf.concat([neg_res, pos_res], axis=0) s_score = tf.math.reduce_sum(res) mse = tf.math.reduce_sum(tf.keras.backend.mean(tf.math.squared_difference(y_pred, y_true), axis=-1)) return (s_score + mse) / 2.0
额外说明
- 将
ModelCheckpoint的verbose设为1后,输出会自动包含Epoch序号,匹配预期格式。 - 未追踪函数警告大多不影响训练,但优化损失函数后可彻底消除该警告,避免潜在的加载后函数不可用问题。
内容的提问来源于stack exchange,提问作者darrahts
相关产品推荐
相关产品推荐

