如何调整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
相关产品推荐
相关产品推荐

