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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.01 13:43:12