使用Keras多分类神经网络时fit_generator参数错误求助
解决
fit_generator()参数冲突的错误 这个错误的根源是你混淆了fit()和fit_generator()的参数用法,导致参数传递冲突了。
错误原因分析
fit_generator()的第一个参数是数据生成器(它会自动返回包含特征和标签的批次数据),而你却把train_set和train_labels作为两个位置参数传了进去——这就导致第二个位置参数train_labels被Keras当成了steps_per_epoch参数,而你后面又显式指定了steps_per_epoch=(train_samples/ batch_size),自然就出现了got multiple values for argument 'steps_per_epoch'的冲突错误。
两种解决方案
方案1:改用model.fit()(推荐,如果你用的是Keras 2.1.0+)
如果你的train_set和train_labels是numpy数组、张量这类直接可用于训练的数据格式,完全不需要用fit_generator(),直接用fit()方法即可:
NN.fit( train_set, train_labels, steps_per_epoch=(train_samples // batch_size), # 注意用整数除法保证参数为整数 epochs=epochs, validation_data=(validation_set, validation_labels), validation_steps=(validation_samples // batch_size) )
小提示:把普通除法改成整数除法
//,因为steps_per_epoch要求传入整数类型的参数。
方案2:正确使用fit_generator()(如果你确实需要用生成器)
如果你必须使用生成器(比如需要实时数据增强),你需要把训练数据和标签包装成一个生成器,比如用ImageDataGenerator的flow()方法:
# 假设你已经定义了ImageDataGenerator实例用于数据增强 train_generator = datagen.flow(train_set, train_labels, batch_size=batch_size) val_generator = datagen.flow(validation_set, validation_labels, batch_size=batch_size) # 调用fit_generator NN.fit_generator( train_generator, steps_per_epoch=(train_samples // batch_size), epochs=epochs, validation_data=val_generator, validation_steps=(validation_samples // batch_size) )
这里生成器会自动每次返回(batch_x, batch_y)格式的批次数据,所以不需要单独传标签参数。
额外提示
从Keras 2.0版本开始,fit()方法已经支持生成器输入了,所以除非你有特殊的定制需求,优先使用fit()会让代码更简洁易读。
内容的提问来源于stack exchange,提问作者Cameron Blake
相关产品推荐
相关产品推荐

