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

使用Keras Functional API调用fit_generator时出现输入形状错误

解决Keras自定义生成器训练时的输入形状错误问题

这是使用自定义生成器时非常常见的坑——你的生成器每次返回单个样本,但Keras期望生成器返回批量数据,这就是导致形状不匹配的核心原因。我来帮你梳理解决方案:

问题根源

你当前的生成器每次yield的是单条序列和单个标签,比如输入的形状是(seq_length,),标签是单个数值。但Keras模型训练时,要求输入是带批量维度的张量(例如形状为(batch_size, seq_length)的数组),单个样本的形状和模型输入层定义的形状不匹配,自然会抛出错误。

修改后的生成器代码

调整生成器,让它每次返回一个批次的数据,同时加入可选的数据打乱操作提升训练效果:

import numpy as np
from tensorflow.keras.preprocessing.sequence import pad_sequences

def data_generator(batch_size):
    while True:
        # 打乱数据索引,避免模型记住数据顺序(可选但推荐)
        indices = np.random.permutation(len(sequences1))
        # 按批次切割数据
        for start_idx in range(0, len(indices), batch_size):
            batch_indices = indices[start_idx:start_idx + batch_size]
            # 收集当前批次的输入和标签
            batch_seq1 = [sequences1[idx] for idx in batch_indices]
            batch_seq2 = [sequences2[idx] for idx in batch_indices]
            batch_labels = [trainLabels[idx] for idx in batch_indices]
            
            # 关键操作:如果序列是变长的,必须统一长度才能组成数组
            # 根据你的模型需求选择padding方式(pre/post)
            batch_seq1 = pad_sequences(batch_seq1, padding='post')
            batch_seq2 = pad_sequences(batch_seq2, padding='post')
            
            # 返回批量格式的输入和标签
            yield ([np.array(batch_seq1), np.array(batch_seq2)], np.array(batch_labels))

训练调用方式

在新版本Keras(2.1.0+)中,fit_generator已被整合到fit方法里,推荐直接使用:

batch_size = 32  # 根据你的显存和数据规模调整
steps_per_epoch = len(sequences1) // batch_size  # 每个epoch需要执行的步数

model.fit(
    data_generator(batch_size),
    steps_per_epoch=steps_per_epoch,
    epochs=10  # 根据你的训练需求调整
)

额外注意事项

  • 如果你的序列是定长的,可以省略pad_sequences步骤,直接转成numpy数组即可;如果是变长序列,必须做padding,或者在模型输入层设置input_shape=(None,)接受变长输入,但批量时仍需统一长度才能生成数组。
  • 确保模型输入层的形状和批量数据的形状匹配:比如输入层定义为Input(shape=(seq_length,)),那么批量后的batch_seq1形状应该是(batch_size, seq_length)。

内容的提问来源于stack exchange,提问作者Amir Panahandeh

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.27 04:21:12