使用EarlyStopping回调时AutoKeras.ImageClassifier返回空History对象
AutoKeras ImageClassifier触发EarlyStopping后history返回None的解决办法
问题场景
在同一数据集上测试含autokeras.ImageClassifier的多个模型,数据集加载代码:
img_size = (100,120,3) train_dataset = get_dataset(x_train, y_train, img_size[:-1], 128) valid_dataset = get_dataset(x_valid, y_valid, img_size[:-1], 128) test_dataset = get_dataset(x_test, y_test, img_size[:-1], 128)
创建模型并添加tf.keras.callbacks.EarlyStopping回调训练:
# 创建网络 model = ak.ImageClassifier(overwrite=True, max_trials=1, metrics=['accuracy']) # 训练网络 early_stop = tf.keras.callbacks.EarlyStopping(monitor='val_loss', mode='min', patience=2) #其余参数默认 history = model.fit(train_dataset, epochs=10, validation_data=valid_dataset, callbacks=[early_stop]) # 评估网络 model.evaluate(test_dataset)
核心问题:训练因EarlyStopping终止时,history返回None;移除回调后,能正常获取history对象。训练终止时输出包含模型保存警告及检查点变量未找到的提示。
原因分析
AutoKeras的fit方法内部封装了模型搜索与训练流程,外部传入的EarlyStopping触发终止时,可能打断了其内部的history返回逻辑,导致无法正常返回训练历史。同时训练日志中的检查点警告也说明回调触发时,模型的保存/恢复流程未完全完成,进一步影响了history的生成。
解决方案
1. 自定义回调手动记录训练历史
绕过AutoKeras的history返回逻辑,自己写一个回调来捕获每个epoch的指标:
from tensorflow.keras.callbacks import Callback class HistoryRecorder(Callback): def on_train_begin(self, logs=None): self.history = {'loss': [], 'val_loss': [], 'accuracy': [], 'val_accuracy': []} def on_epoch_end(self, epoch, logs=None): logs = logs or {} self.history['loss'].append(logs.get('loss')) self.history['val_loss'].append(logs.get('val_loss')) self.history['accuracy'].append(logs.get('accuracy')) self.history['val_accuracy'].append(logs.get('val_accuracy')) # 初始化回调 history_recorder = HistoryRecorder() early_stop = tf.keras.callbacks.EarlyStopping(monitor='val_loss', mode='min', patience=2) # 训练时传入两个回调 model.fit(train_dataset, epochs=10, validation_data=valid_dataset, callbacks=[early_stop, history_recorder]) # 从自定义回调中获取历史数据 history = history_recorder.history
2. 升级AutoKeras版本
旧版AutoKeras可能存在与TensorFlow回调的兼容bug,升级到最新稳定版试试:
pip install --upgrade autokeras
3. 手动记录或借助TensorBoard
如果上面的方法都不行,训练时可以在每个epoch结束时打印指标并保存到本地文件,后续整理成history格式;也可以添加TensorBoard回调,将指标写入日志文件,之后从日志中提取训练历史:
from tensorflow.keras.callbacks import TensorBoard tensorboard_callback = TensorBoard(log_dir='./logs', histogram_freq=1) model.fit(train_dataset, epochs=10, validation_data=valid_dataset, callbacks=[early_stop, tensorboard_callback])
内容的提问来源于stack exchange,提问作者Videgain
相关产品推荐
相关产品推荐

