You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.26 11:09:13