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

TensorFlow Keras中model.fit结合生成器出现额外批次问题排查

问题:TensorFlow自定义生成器出现额外批次,epoch批次计数异常

问题场景

用自定义生成器配合model.fit训练模型时,出现了不符合预期的行为:

  • 训练生成器做完len(train_indices)(607个)批次后,会多生成1个Epoch:2,Batch:1的批次
  • 验证生成器做完len(val_indices)(67个)批次后,同样多生成1个Epoch:2,Batch:1的批次
  • 后续epoch的批次计数直接从2开始,没法回到1重新计数

相关代码

model.fit调用

transformer.fit(
   x=data_generation.generate_dataset(batch_size, dontchange, train_indices, filenames),
   epochs=epochs,
   steps_per_epoch=len(train_indices),
   validation_data=data_generation.generate_dataset(batch_size, dontchange, val_indices, filenames),
   validation_steps=len(val_indices)
)

自定义generate_dataset生成器核心逻辑

def generate_dataset(batch_size_in, dontChange, index_list_in, filenames):
    epoch = 0
    while True:
        epoch = epoch + 1
        raw_dataset = tf.data.TFRecordDataset(filenames)
        batch_size = batch_size_in
        index_list = []
        index_list = index_list_in
        cx = 1
        for index in index_list:
            tf.print("Epoch: {}, Batch: {}, Batches total: {}".format(epoch, cx, len(index_list_in)), summarize=32 *
            10 * 250, output_stream="file://logtest.txt")
            cx = cx + 1
            how_much_to_take = batch_size
            batch_of_records = raw_dataset.skip(index).take(how_much_to_take)
            max_number_of_tv_in_batch = 0
            batch = []

            # 省略数据处理逻辑...

            yield (tf.stack(input_batch), tf.stack(attention_mask_batch), tf.stack(padding_mask_batch)), tf.stack(output_batch)

原因分析

问题出在生成器的循环逻辑和TensorFlow的迭代机制上:

  1. 生成器用了无限循环while True,每轮循环对应一个epoch
  2. model.fit执行完steps_per_epoch个批次后,会额外尝试取一次数据(内部迭代检查逻辑),用来判断生成器是否还有后续数据
  3. 此时生成器刚好跑完for index in index_list的循环,直接进入下一轮while True,epoch自动加1,cx重置为1,于是生成了那个多余的Epoch:2,Batch:1批次
  4. 这个额外批次不会参与训练/验证,但会提前更新生成器的计数器,导致下一个真正的epoch批次计数从2开始

修改方案

核心思路是:让生成器在跑完一个epoch的所有批次后,停在循环的等待状态,不要提前进入下一轮epoch更新计数器。具体修改如下:

修改后的generate_dataset函数

def generate_dataset(batch_size_in, dontChange, index_list_in, filenames):
    epoch = 0
    while True:
        # 每轮epoch开始前,先重置批次计数器
        cx = 1
        # 重新加载TFRecord数据集(按需保留,若数据集无需每次重载可移到循环外)
        raw_dataset = tf.data.TFRecordDataset(filenames)
        batch_size = batch_size_in
        index_list = index_list_in
        
        for index in index_list:
            # 打印时用epoch+1,因为此时还没更新epoch,对应当前正在执行的轮次
            tf.print("Epoch: {}, Batch: {}, Batches total: {}".format(epoch + 1, cx, len(index_list_in)), 
                     summarize=32 * 10 * 250, output_stream="file://logtest.txt")
            cx = cx + 1
            how_much_to_take = batch_size
            batch_of_records = raw_dataset.skip(index).take(how_much_to_take)
            max_number_of_tv_in_batch = 0
            batch = []

            # 省略数据处理逻辑...

            yield (tf.stack(input_batch), tf.stack(attention_mask_batch), tf.stack(padding_mask_batch)), tf.stack(output_batch)
        # 只有当当前epoch的所有批次都已生成并返回后,才更新epoch编号
        epoch += 1

修改说明

  1. 调整epoch的更新时机:从while True循环开头移到for循环结束后,确保只有当一个epoch的所有批次都处理完,才会把epoch加1
  2. 日志打印用epoch+1,因为此时epoch还未更新,对应正在执行的当前轮次
  3. 当model.fit尝试取额外批次时,生成器会停在for循环结束、epoch +=1之前的位置,不会提前进入下一轮epoch的批次生成,也就不会出现多余的批次
  4. 每轮while True循环开头重置cx为1,保证每个epoch的批次计数都从1开始

另外,若TFRecord数据集无需每轮epoch重新加载,可把raw_dataset = tf.data.TFRecordDataset(filenames)移到while True循环外面,提升运行效率。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.12 01:53:27