如何在KerasClassifier中使用数据生成器?解决缺失y属性报错
在KerasClassifier中使用数据生成器的解决方法
KerasClassifier(来自keras.wrappers.scikit_learn)的默认fit方法会强制要求传入x和y参数,这和Keras原生Model.fit支持数据生成器的逻辑冲突,导致你遇到缺失y属性的报错。解决核心是让KerasClassifier适配生成器的输入形式,下面是两种可行方案:
方案1:自定义KerasClassifier子类,重写fit方法
直接重写fit方法,跳过原逻辑对x、y的校验,直接调用底层Keras模型的fit方法传入生成器:
from tensorflow.keras.wrappers.scikit_learn import KerasClassifier from tensorflow.keras.models import Sequential from tensorflow.keras.layers import Dense class GeneratorCompatibleKerasClassifier(KerasClassifier): def fit(self, generator, **kwargs): # 直接调用Model的fit方法,支持生成器输入 return self.model.fit(generator, **kwargs) # 定义模型构建函数 def build_keras_model(): model = Sequential([ Dense(64, activation='relu', input_shape=(10,)), Dense(1, activation='sigmoid') ]) model.compile(loss='binary_crossentropy', optimizer='adam', metrics=['accuracy']) return model # 初始化自定义分类器 classifier = GeneratorCompatibleKerasClassifier( build_fn=build_keras_model, epochs=10, batch_size=32 ) # 直接传入数据生成器训练 classifier.fit(train_generator)
方案2:包装生成器为scikit-learn兼容格式
如果不想修改KerasClassifier,可以把生成器包装成一个返回(x, y)迭代器的对象,同时实现__len__方法(部分scikit-learn逻辑需要):
class GeneratorWrapper: def __init__(self, generator): self.generator = generator self.length = len(generator) # 假设生成器实现了__len__ def __len__(self): return self.length def __iter__(self): return self.generator.__iter__() # 包装你的生成器 wrapped_gen = GeneratorWrapper(train_generator) # 传入KerasClassifier时,x传包装后的生成器,y传占位符 classifier = KerasClassifier(build_fn=build_keras_model, epochs=10, batch_size=32) classifier.fit(x=wrapped_gen, y=None)
注意:第二种方案需要确保你的数据生成器本身实现了__len__方法(比如Keras的ImageDataGenerator.flow_from_directory生成器默认支持),否则需要手动计算批次数量并赋值给self.length。
内容的提问来源于stack exchange,提问作者zyleen
相关产品推荐
相关产品推荐

