TensorFlow Keras自定义Sequence首个批次被调用两次的原因问询
自定义Keras Sequence的__getitem__首批次被调用两次的原因解析
问题现象
在开发Keras深度学习模型的数据加载器时,自定义keras.utils.Sequence的__getitem__方法中添加了打印语句,负责返回第n个批次。测试发现执行model.fit/model.predict时,首个批次(Batch 0)的打印语句触发了两次,但on_batch_end回调函数仅按预期触发一次。即使不重新初始化模型,重复运行fit/predict该现象仍存在,怀疑是模型构建时的初始化操作导致,但未找到相关资料。
代码示例
import numpy as np from keras.callbacks import Callback from tensorflow import keras class MyBatchCallback(Callback): def on_batch_end(self, batch, logs=None): # Your code to be executed at the end of each batch print(f"At the end of batch {batch}") # Create a custom data sequence class by inheriting from keras.utils.Sequence class MyDataSequence(keras.utils.Sequence): def __init__(self, x_set, y_set, batch_size): self.x, self.y = x_set, y_set self.batch_size = batch_size def __len__(self): return int(np.ceil(len(self.x) / self.batch_size)) def __getitem__(self, idx): print(f"------------ Batch {idx}") start_idx = idx * self.batch_size end_idx = (idx + 1) * self.batch_size batch_x = self.x[start_idx:end_idx] batch_y = self.y[start_idx:end_idx] return np.array(batch_x), np.array(batch_y) # Example usage # Generate some dummy data x_train = np.random.random((100, 32)) y_train = keras.utils.to_categorical( np.random.randint(10, size=(100, 1)), num_classes=10 ) # Set batch size batch_size = 32 # Create an instance of your custom data sequence data_sequence = MyDataSequence(x_train, y_train, batch_size) # Create a simple model for illustration purposes model = keras.models.Sequential() model.add(keras.layers.Dense(64, activation="relu", input_dim=32)) model.add(keras.layers.Dense(10, activation="softmax")) # Compile the model model.compile(optimizer="adam", loss="categorical_crossentropy", metrics=["accuracy"]) # Train the model using the data sequence model.fit(data_sequence, verbose=0, shuffle=False, callbacks=[MyBatchCallback()])
控制台输出
------------ Batch 0 2024-01-24 09:35:11.573324: W tensorflow/core/platform/profile_utils/cpu_utils.cc:128] Failed to get CPU frequency: 0 Hz 2024-01-24 09:35:11.706627: I tensorflow/core/grappler/optimizers/custom_graph_optimizer_registry.cc:113] Plugin optimizer for device_type GPU is enabled. ------------ Batch 0 ------------ Batch 1 At the end of batch 0 ------------ Batch 2 At the end of batch 1 ------------ Batch 3 At the end of batch 2 At the end of batch 3
原因解析
这是TensorFlow/Keras的正常内部行为,核心原因是训练前的输入形状验证与模型构建步骤:
- 正式训练开始前,Keras会调用一次数据加载器的首个批次,目的是推断输入数据的形状、类型等元信息,确保模型的输入层与数据匹配,完成最终的模型构建(即使显式定义了输入维度,部分TF版本仍会执行该检查)。
- 这次调用属于框架内部的验证操作,不会计入正式的训练批次循环,因此不会触发
on_batch_end回调——只有正式训练时的批次迭代才会触发回调逻辑。 - 每次调用
fit/predict时,Keras都会执行该输入验证步骤,因此即使模型已初始化完成,重复运行仍会出现首批次被调用两次的现象。
如果需要避免额外的打印输出,可以在__getitem__中添加判断逻辑(比如记录首次调用并跳过打印),但不建议修改数据返回的核心逻辑,这是框架确保模型与数据兼容性的必要步骤。
内容的提问来源于stack exchange,提问作者Steph Pepito
相关产品推荐
相关产品推荐

