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

Keras训练中如何在1000个epoch后存模型并确认所用GPU

嘿,我帮你搞定这两个需求,直接上修改方案和验证方法:

一、修改模型保存逻辑,仅在第1000个epoch后保存

你原来用的ModelCheckpoint默认会在每个epoch后保存模型,要实现只在第1000个epoch保存,我们可以自定义一个Keras回调函数,只在指定epoch触发保存动作。

替换掉原来的ModelCheckpoint部分,改用下面的自定义回调类:

from keras.callbacks import Callback

class CustomModelCheckpoint(Callback):
    def __init__(self, save_path, target_epoch=1000):
        super(CustomModelCheckpoint, self).__init__()
        self.save_path = save_path
        self.target_epoch = target_epoch

    def on_epoch_end(self, epoch, logs=None):
        # Keras的epoch从0开始计数,第1000个epoch对应epoch值为999,所以用epoch+1判断更直观
        if epoch + 1 == self.target_epoch:
            self.model.save(self.save_path)
            print(f"模型已在第{self.target_epoch}个epoch后保存到{self.save_path}")

然后在训练时调用这个回调:

# 替换原来的ModelCheckpoint实例化
custom_checkpoint = CustomModelCheckpoint(save_path='your_best_model.h5', target_epoch=1000)

# 训练时传入回调列表(保留你的CSVLogger)
model.fit(x_train, y_train, epochs=10000, callbacks=[CSVLogger('training.log'), custom_checkpoint])
二、确认当前训练使用的GPU

你代码里已经设置了os.environ['CUDA_VISIBLE_DEVICES'] = '1',这里的编号是从0开始计数的,理论上指定的是第2块GPU(编号0为第一块)。不过可以用以下几种方法验证:

方法1:在代码中打印GPU信息

在代码开头添加以下内容,直接查看TensorFlow识别到的GPU:

# 打印系统中所有物理GPU
physical_devices = tf.config.list_physical_devices('GPU')
print(f"系统物理GPU列表:{physical_devices}")

# 打印当前程序可见的GPU(由CUDA_VISIBLE_DEVICES控制)
visible_devices = tf.config.get_visible_devices('GPU')
print(f"当前程序可见GPU:{visible_devices}")

输出结果里的编号会和你设置的CUDA_VISIBLE_DEVICES对应,比如设置了1,可见GPU就会是编号1的那块。

方法2:用系统命令查看显存占用

打开终端运行nvidia-smi命令,查看两块GPU的显存使用情况:

  • 训练开始后,哪块GPU的显存被大量占用(通常会占满大部分显存),就是当前训练使用的GPU。
  • 比如GPU 1列下的Memory-Usage有明显增长,就对应你代码里设置的CUDA_VISIBLE_DEVICES='1'。

方法3:打印设备分配日志

在代码开头添加以下配置,让TensorFlow输出每个操作的设备分配信息:

tf.debugging.set_log_device_placement(True)

运行代码后,控制台会输出类似MatMul: (MatMul): /job:localhost/replica:0/task:0/device:GPU:1的日志,这里的GPU:1就是当前使用的GPU。

把这些修改整合到你的代码里,就能满足你的需求啦!

内容的提问来源于stack exchange,提问作者Jaffer Wilson

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.19 03:15:54