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

将生成器传入TensorFlow的model.fit时需注意什么?训练精度不提升

问题:用生成器替换fit的x/y参数后模型精度无法正常提升

尝试用生成器替代tf.keras.Model.fit()中的x、y训练数据参数,按照文档要求,生成器返回元组(x_vals, y_vals)(其中x_vals和y_vals是按batch_size拼接的样本与标签),同时指定了steps_per_epoch。但替换后,原本能正常训练的模型精度先小幅上升后回落至随机水平,而直接传入原始数据或让生成器生成全量样本再传入fit则能正常训练。

错误原因

问题出在自定义生成器的实现上:

  • 生成器没有对训练数据进行打乱,而原始fit调用中使用了shuffle=True,固定的样本顺序会让模型容易陷入局部最优,甚至因重复学习相同序列的样本导致泛化能力急剧下降。
  • 生成器的索引持续递增,每个epoch不会重置也不会重新打乱数据,导致后续epoch的样本序列完全重复,模型无法学习到有效特征。

修复后的生成器实现

调整生成器,在每个epoch开始时打乱数据索引,并按批次遍历数据:

import numpy as np
import tensorflow as tf

BATCH_SIZE = 32

def load_cifar():
    (x_train, y_train), (x_test, y_test) = tf.keras.datasets.cifar10.load_data()
    assert x_train.shape == (50000, 32, 32, 3)
    assert x_test.shape == (10000, 32, 32, 3)
    assert y_train.shape == (50000, 1)
    assert y_test.shape == (10000, 1)

    x_train = np.true_divide(x_train,255,dtype=np.single)
    x_test =  np.true_divide(x_test,255,dtype=np.single)
    y_train = y_train.astype(np.single)
    y_test =  y_test.astype(np.single)

    return (x_train,y_train), (x_test,y_test)

(train_x, train_y) , (validation_x, validation_y) = load_cifar()

# 修复后的生成器
def data_generator_fixed(input_data_x:np.ndarray,
                          input_data_y:np.ndarray,
                          batch_size=BATCH_SIZE,
                          ):
    num_samples = input_data_x.shape[0]
    while True:
        # 每个epoch开始时打乱数据索引
        indices = np.random.permutation(num_samples)
        for start_idx in range(0, num_samples, batch_size):
            end_idx = min(start_idx + batch_size, num_samples)
            batch_indices = indices[start_idx:end_idx]
            batch_x = input_data_x[batch_indices]
            batch_y = input_data_y[batch_indices]
            
            # 可选:补全最后一个不足batch_size的批次
            if len(batch_x) < batch_size:
                pad_size = batch_size - len(batch_x)
                batch_x = np.concatenate([batch_x, input_data_x[indices[:pad_size]]])
                batch_y = np.concatenate([batch_y, input_data_y[indices[:pad_size]]])
                
            yield batch_x, batch_y

def make_model():
    model = tf.keras.models.Sequential()
    model.add(tf.keras.layers.Flatten())
    model.add(tf.keras.layers.Dense(10,tf.nn.softmax))
    model.build([None] +list(train_x[0,:,:,:].shape))
    return model

# 训练代码
generator = data_generator_fixed(train_x,train_y,batch_size=BATCH_SIZE)
model = make_model()
model.summary()
optimizer_adam=tf.keras.optimizers.Adam(learning_rate=0.0005/32,beta_1=0.9,beta_2=0.999,epsilon=1e-07)
model.compile(optimizer_adam,loss="sparse_categorical_crossentropy",metrics="accuracy")
model.fit(generator, validation_data=(validation_x,validation_y),epochs=10,
          steps_per_epoch=train_x.shape[0]//BATCH_SIZE,
          )

额外说明

  • 修复后的生成器在每个epoch都会重新打乱数据索引,和原始fit的shuffle=True行为一致,确保模型每次迭代看到的样本顺序不同。
  • 处理了最后一个batch可能不足batch_size的情况,避免因维度不一致导致训练报错(也可以选择不补全,此时steps_per_epoch需设为num_samples // batch_size + 1)。
  • 当生成器生成全量样本再传入fit时,fit内部会自动执行shuffle,所以模型能正常训练,这也验证了数据打乱是解决问题的关键。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.30 05:37:53