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

如何使用fit_generator实现多输入?能否使用两个fit_generator?

解决双输入模型的fit_generator使用问题

首先,你遇到的报错核心原因很明确:你的模型定义了2个输入,但传入生成器的输出只包含了1组数据(大概率是把X1和X2混在一起,或者只传了其中一个)。下面分两部分帮你解决问题:

一、正确使用单个fit_generator处理多输入

要让生成器适配双输入模型,你的生成器需要每次返回一个元组,其中第一个元素是包含两个输入数组的列表,第二个元素是对应标签数组。简单来说,生成器的输出格式必须是:([X1_batch, X2_batch], Y_batch)。

自定义生成器示例

假设你已经有加载X1、X2、Y的逻辑,这里写一个基础的自定义生成器:

import numpy as np

def multi_input_generator(X1, X2, Y, batch_size=32):
    num_samples = len(X1)
    while True:
        # 训练时随机打乱样本顺序(可选,但推荐)
        indices = np.random.permutation(num_samples)
        for i in range(0, num_samples, batch_size):
            batch_indices = indices[i:i+batch_size]
            # 提取当前批次的两个输入和对应标签
            X1_batch = X1[batch_indices]
            X2_batch = X2[batch_indices]
            Y_batch = Y[batch_indices]
            # 返回符合模型要求的格式
            yield ([X1_batch, X2_batch], Y_batch)

训练时直接传入这个生成器即可(注意:Keras 2.10+版本中fit_generator已整合到fit方法,直接用model.fit()也能兼容):

# 假设你的双输入模型已定义完成
model.fit(
    multi_input_generator(X1, X2, Y, batch_size=32),
    steps_per_epoch=len(X1)//32,
    epochs=10
)

内置生成器(如ImageDataGenerator)的组合方式

如果你用ImageDataGenerator做数据增强,可以分别为X1、X2创建生成器,再通过zip组合,同时保证标签同步:

from tensorflow.keras.preprocessing.image import ImageDataGenerator

# 为两个输入分别定义数据增强规则
datagen_x1 = ImageDataGenerator(rescale=1./255, rotation_range=10)
datagen_x2 = ImageDataGenerator(rescale=1./255, width_shift_range=0.1)

# 生成输入批次,注意标签共用同一个Y,且设置相同seed保证打乱顺序一致
generator_x1 = datagen_x1.flow(X1, Y, batch_size=32, seed=42)
generator_x2 = datagen_x2.flow(X2, Y, batch_size=32, seed=42)

# 包装成符合要求的生成器
def combined_generator():
    for (x1_batch, y_batch), (x2_batch, _) in zip(generator_x1, generator_x2):
        yield ([x1_batch, x2_batch], y_batch)

# 启动训练
model.fit(
    combined_generator(),
    steps_per_epoch=len(X1)//32,
    epochs=10
)

二、能不能使用两个fit_generator?

答案是绝对不可以。Keras的模型训练过程中,只能接收一个统一的数据源(生成器或数据数组)来同步提供所有输入和对应的标签。如果用两个fit_generator,模型无法将两个生成器的输入与标签正确匹配,会直接导致输入-标签错位,甚至抛出新的格式错误。必须把两个输入整合到同一个生成器的输出里,或者像上面那样用zip组合后包装成单一生成器。

额外注意事项

  • 确保X1和X2的样本数量完全一致,否则生成器迭代时会出现长度不匹配的问题;
  • 如果使用数据打乱,要保证两个输入的打乱顺序完全同步(比如设置相同的seed,或者在自定义生成器中统一打乱索引),避免输入和标签错位。

内容的提问来源于stack exchange,提问作者김태형

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.20 11:37:38