Keras升级后model.fit_generator抛出NotImplementedError问题求助
解决升级TensorFlow/Keras后
fit_generator抛出NotImplementedError的问题 问题根源
你遇到的这个错误本质是TensorFlow 2.x及后续的Keras版本已经彻底废弃了fit_generator方法,同时旧版Keras的Iterator基类在新版本中的内部实现逻辑发生了变化,导致你基于它自定义的生成器不再兼容。旧版里fit_generator是专门为这类迭代器设计的,但新版本统一用model.fit处理所有输入类型(生成器、tf.data.Dataset等),旧Iterator的部分接口无法被新的fit方法正确识别,从而触发了未实现错误。
解决方案1:直接替换fit_generator为model.fit
新版本的model.fit完全支持生成器输入,参数和fit_generator几乎一致,这是最快速的修复方式:
model.fit( generator=generator_train, steps_per_epoch=generator_train.n // batch_size, # 建议用整数除法,避免浮点数类型问题 epochs=20, verbose=1, validation_data=generator_val, validation_steps=math.ceil(generator_val.n / batch_size), callbacks=[tb_callback, saver_callback], use_multiprocessing=False, initial_epoch=0 )
解决方案2:调整自定义生成器的兼容性
你的BoxCarsDataGenerator继承了旧的keras.preprocessing.image.Iterator,这个类在新版本中已被标记为过时,你可以通过两种方式调整:
方式A:给生成器添加__getitem__方法
新版本Keras要求自定义迭代器实现__getitem__方法(仅靠__next__不再足够),因为fit方法会尝试通过索引访问批次数据。在你的类中添加以下代码:
class BoxCarsDataGenerator(Iterator): # 保留原有的__init__和__next__方法... def __getitem__(self, idx): # 直接复用__next__的逻辑即可 return self.__next__()
方式B:改为普通Python生成器(推荐)
放弃继承旧的Iterator类,改成普通的Python生成器函数,兼容性更好,逻辑也更直观:
def boxcars_data_generator(dataset, part, batch_size=8, training_mode=False, generate_y=True, image_size=(224,224)): assert image_size == (224,224), "only images 224x224 are supported by unpack_3DBB for now" assert dataset.X[part] is not None, "load some classification split first" if dataset.atlas is None: dataset.load_atlas() num_samples = dataset.X[part].shape[0] indices = np.arange(num_samples) if training_mode: np.random.shuffle(indices) while True: for start in range(0, num_samples, batch_size): end = min(start + batch_size, num_samples) batch_indices = indices[start:end] x = np.empty([len(batch_indices)] + list(image_size) + [3], dtype=np.float32) for i, ind in enumerate(batch_indices): vehicle_id, instance_id = dataset.X[part][ind] vehicle, instance, bb3d = dataset.get_vehicle_instance_data(vehicle_id, instance_id) image = dataset.get_image(vehicle_id, instance_id) if training_mode: image = alter_HSV(image) image = image_drop(image) bb_noise = np.clip(np.random.randn(2) * 1.5, -5, 5) flip = bool(random.getrandbits(1)) image, bb3d = add_bb_noise_flip(image, bb3d, flip, bb_noise) image = unpack_3DBB(image, bb3d) image = (image.astype(np.float32) - 116)/128. x[i, ...] = image if not generate_y: yield x else: y = dataset.Y[part][batch_indices] yield x, y
之后创建生成器实例的方式改为:
generator_train = boxcars_data_generator(dataset, 'train', batch_size=batch_size, training_mode=True) generator_val = boxcars_data_generator(dataset, 'val', batch_size=batch_size, training_mode=False)
额外注意事项
- 确保你的TensorFlow/Keras是最新稳定版,避免遗留的版本兼容问题;
- 如果启用多进程(
use_multiprocessing=True),普通生成器要注意随机种子的设置,避免进程间随机性冲突; - 检查
alter_HSV、image_drop等自定义数据增强函数的输出数据类型,确保符合TensorFlow的要求。
内容的提问来源于stack exchange,提问作者Roglleg RMA
相关产品推荐
相关产品推荐

