如何在Keras中加载ModelCheckpoint回调保存的最优权重?
正确使用ModelCheckpoint保存并加载最优权重
问题根源
你两次报错的核心原因都是路径指向了目录而非具体文件:
- 第一次保存时,
tmp/checkpoint是目录,但ModelCheckpoint需要明确的文件路径来写入权重数据; - 第二次加载时,你依然指向目录
tmp/checkpoint,但load_weights方法必须接收具体的权重文件路径。
两种正确实现方式
方式一:固定文件名保存最优权重(推荐)
这种方式会在每次发现更优模型时自动覆盖旧文件,最终只保留验证集指标最优的那一份权重,加载时直接用固定路径即可,操作更简洁。
- 配置回调并设置固定文件路径:
checkpoint_filepath = 'tmp/best_weights.h5' model_checkpoint_callback = keras.callbacks.ModelCheckpoint( filepath=checkpoint_filepath, save_weights_only=True, monitor='val_loss', mode='min', save_best_only=True )
- 训练时传入回调:
model.fit(x_train, y_train, validation_data=(x_val, y_val), callbacks=[model_checkpoint_callback])
- 加载最优权重:
net.load_weights('tmp/best_weights.h5')
方式二:按epoch和指标命名保存多份权重文件
如果需要保留不同训练阶段的权重文件,可以用带变量的文件名,但加载时必须指定具体的最优文件路径。
- 先确保保存目录存在,再配置回调:
import os # 自动创建目录,避免训练时因目录不存在报错 os.makedirs('tmp/checkpoint', exist_ok=True) checkpoint_filepath = 'tmp/checkpoint/weights.{epoch:02d}-{val_loss:.2f}.hdf5' model_checkpoint_callback = keras.callbacks.ModelCheckpoint( filepath=checkpoint_filepath, save_weights_only=True, monitor='val_loss', mode='min', save_best_only=True )
- 训练完成后,
tmp/checkpoint目录下会生成类似weights.05-0.08.hdf5的文件,文件名中的val_loss数值越小,代表对应模型的验证集性能越好。 - 加载时指定具体的最优文件路径:
# 替换为你实际生成的最优权重文件名 net.load_weights('tmp/checkpoint/weights.05-0.08.hdf5')
注意事项
- 若使用HDF5格式(
.h5/.hdf5),需确保已安装h5py库,否则可能出现加载失败; save_best_only=True参数会确保只保存验证集指标最优的模型,避免生成大量冗余文件;- 若使用TensorFlow 2.x,也可选择保存为SavedModel格式(去掉
save_weights_only=True或设置save_format='tf'),但加载时需改用model.load_model()方法。
内容的提问来源于stack exchange,提问作者ailauli69
相关产品推荐
相关产品推荐

