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

如何调整TensorFlow API中检查点的保存间隔(1000步改500步)

TensorFlow调整检查点保存间隔为500步的方法

1. Keras API(tf.keras)场景

在ModelCheckpoint回调中设置save_freq参数即可指定每500步保存一次:

  • 示例代码:
from tensorflow.keras.callbacks import ModelCheckpoint

# 定义检查点回调
checkpoint_callback = ModelCheckpoint(
    filepath='./checkpoints/model_{epoch}_{step}',
    save_freq=500,  # 核心配置:每500步保存一次
    save_weights_only=False,
    verbose=1
)

# 训练时传入该回调
model.fit(
    x_train, y_train,
    epochs=10,
    callbacks=[checkpoint_callback]
)
  • 注意:save_freq设为整数时按训练步数计数,设为'epoch'则按轮次保存,默认是轮次。

2. Estimator API场景

通过RunConfig的save_checkpoints_steps参数配置:

  • 示例代码:
import tensorflow as tf

# 配置训练参数
run_config = tf.estimator.RunConfig(
    save_checkpoints_steps=500,  # 每500步保存检查点
    model_dir='./checkpoints'
)

# 初始化Estimator
estimator = tf.estimator.DNNClassifier(
    feature_columns=feature_cols,
    hidden_units=[128, 64],
    n_classes=10,
    config=run_config
)

# 启动训练
estimator.train(input_fn=train_input_fn, steps=10000)

3. 自定义训练循环(tf.GradientTape)场景

需要手动在步数整除500时触发保存操作:

  • 示例代码:
checkpoint = tf.train.Checkpoint(model=model, optimizer=optimizer)
manager = tf.train.CheckpointManager(checkpoint, './checkpoints', max_to_keep=5)

step = 0
for epoch in range(epochs):
    for x, y in train_dataset:
        with tf.GradientTape() as tape:
            predictions = model(x)
            loss = loss_fn(y, predictions)
        gradients = tape.gradient(loss, model.trainable_variables)
        optimizer.apply_gradients(zip(gradients, model.trainable_variables))
        
        step += 1
        if step % 500 == 0:
            manager.save()
            print(f"检查点已保存,当前步数:{step}")

内容的提问来源于stack exchange,提问作者Shaw AN

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.12 23:15:03