如何在Python中正确使用fit或fit_generator解决交叉验证训练问题?
解决方案
首先,Model.fit() 不存在 generator 参数,这是你报错的核心原因。在新版本Keras/TensorFlow中,fit() 直接支持生成器作为输入,无需额外的generator参数,只需把生成器直接传给fit()的第一个位置参数(即x参数)即可。
另外,你之前设置的steps_per_epoch=batches.n是错误的——batches.n是训练集的总样本数,而steps_per_epoch需要的是每轮训练的步数(即总样本数除以batch_size)。直接用len(batches)就能得到正确步数,因为flow生成器的长度就是总样本数除以batch_size的计算结果。同理validation_steps也应该用len(val_batches)。
修正后的训练代码
history = model.fit( batches, # 直接传入生成器作为第一个参数 steps_per_epoch=len(batches), epochs=3, validation_data=val_batches, validation_steps=len(val_batches) )
额外说明
- 若想显式指定参数名,可写
x=batches,效果和直接传位置参数一致。 - 之前
fit_generator中用batches.n作为steps_per_epoch会导致每轮执行37800步(你的样本总数),但每步已处理64个样本,这会重复训练数据多次;修正后每轮只会执行正确步数(37800//64≈591步),训练效率和结果都会更合理。
内容的提问来源于stack exchange,提问作者Thanh Pham
相关产品推荐
相关产品推荐

