运行GitHub停车位检测代码报错fit_generator存在未预期参数samples_per_epoch
问题解决方法
报错原因
这个报错是Keras/TensorFlow版本不兼容导致的。你参考的旧代码基于Keras 2.1.0之前的版本编写,该版本的fit_generator()使用samples_per_epoch、nb_val_samples参数,而新版本中这两个参数名已被移除替换。
解决方案
方案1:适配fit_generator()参数
直接替换旧参数名即可:
- 将
samples_per_epoch替换为steps_per_epoch,参数值为训练样本数 // 批次大小(即nb_train_samples // batch_size,batch_size为你训练生成器设置的单批次样本数) - 将
nb_val_samples替换为validation_steps,参数值为验证样本数 // 批次大小(即nb_validation_samples // batch_size)
修改后代码示例:
### Start training! history_object = model_final.fit_generator( train_generator, steps_per_epoch = nb_train_samples // batch_size, epochs = epochs, validation_data = validation_generator, validation_steps = nb_validation_samples // batch_size, callbacks = [checkpoint, early] )
方案2:使用新版本推荐的fit()方法
TensorFlow 2.x版本的fit()方法原生支持生成器输入,无需再调用fit_generator(),写法更简洁:
### Start training! history_object = model_final.fit( train_generator, steps_per_epoch = nb_train_samples // batch_size, epochs = epochs, validation_data = validation_generator, validation_steps = nb_validation_samples // batch_size, callbacks = [checkpoint, early] )
注意事项
- 如果你的生成器设置为无限循环生成样本,必须指定
steps_per_epoch和validation_steps参数,否则训练过程会无限执行不会自动进入下一轮。 - 如果样本数不能被批次大小整除,可将参数值调整为
nb_train_samples // batch_size + 1避免遗漏最后一批样本。
内容的提问来源于stack exchange,提问作者Nusry KR
相关产品推荐
相关产品推荐

