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

如何在TensorFlow中延续训练并续写至同一历史文件

TensorFlow模型断点续训、历史记录续写及全局最优模型保存方案

一、训练历史文件可以续写,操作方法如下

训练时model.fit()返回的History对象会记录每轮的指标数据,我们可以把这些数据序列化保存,续训时加载旧历史并合并新数据,实现完整追踪训练演化过程。

1. 首次训练保存历史

训练完成后将历史数据存为JSON或CSV文件,示例用JSON:

import json
import tensorflow as tf

# 示例模型构建
model = tf.keras.Sequential([
    tf.keras.layers.Dense(64, activation='relu', input_shape=(10,)),
    tf.keras.layers.Dense(1)
])
model.compile(optimizer='adam', loss='mse')

# 首次训练
history = model.fit(X_train, y_train, epochs=10)

# 保存历史到文件
with open('training_history.json', 'w') as f:
    json.dump(history.history, f)
# 同时保存当前模型,方便后续续训
model.save('current_model.h5')

2. 续训时加载并续写历史

加载之前的模型和历史文件,续训后合并新旧数据再保存:

import json
import tensorflow as tf

# 加载之前的模型
model = tf.keras.models.load_model('current_model.h5')

# 加载旧历史数据
with open('training_history.json', 'r') as f:
    old_history = json.load(f)

# 用新数据续训,记录新的训练历史
new_history = model.fit(X_new_train, y_new_train, epochs=5)

# 合并历史:把每个指标的新数据追加到旧列表末尾
for metric in old_history.keys():
    old_history[metric].extend(new_history.history[metric])

# 覆盖保存合并后的历史文件
with open('training_history.json', 'w') as f:
    json.dump(old_history, f)
# 同时更新当前模型的保存
model.save('current_model.h5')

如果习惯用CSV,逻辑类似:首次训练后用pandas把history.history转为DataFrame存成CSV,续训后读取CSV,将新的指标数据追加进去再保存。

二、保存全局最优模型的实现方法

默认的ModelCheckpoint只会保存单次训练中的最优模型,要实现全局最优(跨多次续训的最优),可以自定义回调函数,追踪历史上的最佳性能,只有当当前模型超过历史最优时才保存。

1. 自定义全局最优保存回调

import numpy as np
import tensorflow as tf
from tensorflow.keras.callbacks import Callback

class GlobalBestCheckpoint(Callback):
    def __init__(self, save_path, monitor='val_loss', mode='min'):
        super().__init__()
        self.save_path = save_path
        self.monitor = monitor  # 监控的指标,比如val_loss、val_accuracy
        self.mode = mode  # min表示指标越小越好,max表示越大越好
        # 初始化全局最优值:min模式设为无穷大,max模式设为负无穷
        self.best_score = np.inf if mode == 'min' else -np.inf
        # 从文件加载之前的最优值,避免续训时重置
        try:
            with open('best_score.txt', 'r') as f:
                self.best_score = float(f.read())
        except FileNotFoundError:
            pass

    def on_epoch_end(self, epoch, logs=None):
        current_score = logs.get(self.monitor)
        if not current_score:
            return
        # 判断当前性能是否优于全局最优
        if (self.mode == 'min' and current_score < self.best_score) or \
           (self.mode == 'max' and current_score > self.best_score):
            self.best_score = current_score
            # 保存最优模型
            self.model.save(self.save_path)
            # 同时保存当前最优值,方便后续续训加载
            with open('best_score.txt', 'w') as f:
                f.write(str(self.best_score))
            print(f"\n全局最优模型已更新,保存至 {self.save_path}")

2. 训练(含续训)时使用该回调

不管是首次训练还是续训,都加载这个回调,它会自动跨训练周期追踪全局最优:

# 首次训练示例
model = tf.keras.Sequential([
    tf.keras.layers.Dense(64, activation='relu', input_shape=(10,)),
    tf.keras.layers.Dense(1)
])
model.compile(optimizer='adam', loss='mse', metrics=['mae'])

# 初始化全局最优回调,监控验证损失,越小越好
global_checkpoint = GlobalBestCheckpoint('global_best_model.h5', monitor='val_loss', mode='min')

# 首次训练,带验证集
history = model.fit(
    X_train, y_train,
    epochs=10,
    validation_data=(X_val, y_val),
    callbacks=[global_checkpoint]
)
model.save('current_model.h5')

# 续训示例
model = tf.keras.models.load_model('current_model.h5')
# 同样加载全局最优回调,它会自动读取之前保存的best_score.txt
global_checkpoint = GlobalBestCheckpoint('global_best_model.h5', monitor='val_loss', mode='min')

# 用新数据续训
new_history = model.fit(
    X_new_train, y_new_train,
    epochs=5,
    validation_data=(X_new_val, y_new_val),
    callbacks=[global_checkpoint]
)
model.save('current_model.h5')

这样即使某次续训的数据集质量差,只要模型性能没超过历史最优,就不会覆盖之前的全局最优模型,完全满足你的需求。

内容的提问来源于stack exchange,提问作者Jonathan Roy

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.05 23:10:27