Colab中中断TensorFlow训练时保存模型报错:NameError求解
解决Keras训练中断时保存模型的NameError问题
嘿,我来帮你搞定这个问题!首先,你遇到的NameError: name 'models' is not defined是个典型的实例属性访问错误,咱们一步步拆解修复,顺便优化你的中断保存逻辑:
1. 直接错误的根源
你在on_epoch_end方法里写了models.append(model),但你在on_train_begin里初始化的是实例属性self.models,Python找不到全局的models变量,自然就报错了。另外,你要添加的模型应该是回调内置的self.model(当前训练的模型实例),而不是外部的model变量。
2. 修正回调函数的核心问题
先把回调里的错误改掉,同时优化逻辑:
import tensorflow as tf class myCallback(tf.keras.callbacks.Callback): def on_train_begin(self, logs={}): # 初始化实例属性,必须用self.前缀 self.models = [] # 可以在这里设置目标准确率,更灵活 self.target_acc = 0.2 def on_epoch_end(self, epoch, logs={}): # 把模型实例存入列表(不过更推荐直接存到文件,省内存) self.models.append(self.model) # 注意:TensorFlow 2.x的Keras里,logs的准确率键是'accuracy',不是'acc' current_acc = logs.get('accuracy', 0) if current_acc > self.target_acc: print(f'\nReached {self.target_acc*100}% accuracy, stopping training!') self.model.stop_training = True
3. 正确处理训练中断(KeyboardInterrupt)
你原来的try-except放在on_epoch_end里是抓不到训练过程中用户按Ctrl+C的中断的,因为on_epoch_end只在每个epoch结束时触发。正确的做法是在model.fit外面包裹try-except块,捕获中断后立即保存模型:
from keras.models import load_model # 编译模型 model.compile(loss='categorical_crossentropy', optimizer='adam', metrics=['accuracy']) # 初始化回调 callbacks = myCallback() # 用try-except包裹训练逻辑,捕获中断 try: history = model.fit(X_modified, ys, epochs=127, verbose=1, callbacks=[callbacks]) except KeyboardInterrupt: print('\nTraining interrupted by user! Saving current model...') file_Name = "shahnameh_embdding64_bidirectional_LSTM400_softmax_interrupted.h5" model.save(file_Name) print(f'Model saved to {file_Name}') # 也可以保存最后一个epoch的模型(从回调的models列表里取) # last_model = callbacks.models[-1] # last_model.save(file_Name) # 训练正常结束时保存模型 file_Name = "shahnameh_embdding64_bidirectional_LSTM400_softmax.h5" model.save(file_Name)
4. 额外优化建议
- 不要用列表存模型实例:模型实例占用大量内存,训练129轮的话会把内存撑爆,直接每个epoch保存到文件(比如用
ModelCheckpoint回调)更靠谱:# 内置的ModelCheckpoint回调,自动保存每个epoch的模型,或最好的模型 checkpoint_callback = tf.keras.callbacks.ModelCheckpoint( filepath='model_epoch_{epoch:02d}.h5', save_freq='epoch', save_best_only=False ) # 训练时把这个回调也加进去 history = model.fit(..., callbacks=[callbacks, checkpoint_callback]) - 准确率键名兼容:如果你的Keras版本比较旧,
logs里的键是'acc',可以用logs.get('accuracy', logs.get('acc', 0))来兼容新旧版本。 - 中断时保存训练历史:如果需要保存训练的loss和准确率曲线,可以把
history.history保存成json文件,方便后续分析。
内容的提问来源于stack exchange,提问作者chikitin
相关产品推荐
相关产品推荐

