Sequential对象无fit_generator属性报错及解决方案咨询
问题与解决方案
问题说明
运行CNN图像分类代码时触发错误:
AttributeError: 'Sequential' object has no attribute 'fit_generator'
这是因为高版本的Keras(或TensorFlow集成的Keras)已移除fit_generator方法,其功能已合并到fit()中。
解决方案
完全可以用fit()替代fit_generator,仅需修改训练部分代码,其余逻辑无需调整。
修改训练代码
将原代码中的fit_generator调用:
classifier.fit_generator( train_set, steps_per_epoch=8000 // batch_size, epochs=20, validation_data=test_set, validation_steps=2000 // batch_size )
替换为:
classifier.fit( train_set, steps_per_epoch=8000 // batch_size, epochs=20, validation_data=test_set, validation_steps=2000 // batch_size )
额外兼容优化
- 统一导入模块:当前代码混合了
keras.xxx和tensorflow.keras.xxx的导入方式,建议统一使用tensorflow.keras模块,避免版本冲突,开头导入可修改为:
from tensorflow.keras.models import Sequential from tensorflow.keras.layers import Conv2D, MaxPooling2D, Dense, Flatten
(注:Convolution2D是旧命名,Conv2D是官方推荐的新名称,功能完全一致)
fit()方法完全支持原fit_generator的所有参数,包括steps_per_epoch、validation_steps等,替换后可直接正常运行。
内容的提问来源于stack exchange,提问作者Prajwal A.K
相关产品推荐
相关产品推荐

