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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.06 16:57:49