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的迭代机制上:
- 生成器用了无限循环
while True,每轮循环对应一个epoch model.fit执行完steps_per_epoch个批次后,会额外尝试取一次数据(内部迭代检查逻辑),用来判断生成器是否还有后续数据- 此时生成器刚好跑完
for index in index_list的循环,直接进入下一轮while True,epoch自动加1,cx重置为1,于是生成了那个多余的Epoch:2,Batch:1批次 - 这个额外批次不会参与训练/验证,但会提前更新生成器的计数器,导致下一个真正的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
修改说明
- 调整
epoch的更新时机:从while True循环开头移到for循环结束后,确保只有当一个epoch的所有批次都处理完,才会把epoch加1 - 日志打印用
epoch+1,因为此时epoch还未更新,对应正在执行的当前轮次 - 当
model.fit尝试取额外批次时,生成器会停在for循环结束、epoch +=1之前的位置,不会提前进入下一轮epoch的批次生成,也就不会出现多余的批次 - 每轮
while True循环开头重置cx为1,保证每个epoch的批次计数都从1开始
另外,若TFRecord数据集无需每轮epoch重新加载,可把raw_dataset = tf.data.TFRecordDataset(filenames)移到while True循环外面,提升运行效率。
内容的提问来源于stack exchange,提问作者bievjucs
相关产品推荐
相关产品推荐

