调用train_on_batch时on_train_batch_end触发TypeError错误求助
问题解决思路
错误原因分析
报错TypeError: unsupported operand type(s) for -: 'float' and 'NoneType'是因为CallbackList内部的批次计时变量未初始化:你只调用了clist.on_train_batch_end,但没有调用对应的clist.on_train_batch_begin来启动计时,导致self._batch_start_time为None,无法计算时间差。
同时代码还存在两个问题:
model.train_on_batch的返回值是当前批次的损失,但你没有把这个值存入logs字典,导致自定义的ShowBatchLoss回调拿不到损失数据- 缺少训练流程的钩子调用(比如
on_train_begin、on_epoch_begin),不符合Keras回调的执行规范
修复方案
方案1:补全CallbackList的完整调用流程
按照Keras的回调生命周期,补全所有必要的钩子调用,同时将train_on_batch的损失存入logs:
import tensorflow as tf from keras.layers import * from keras.models import Model from keras.optimizers import SGD from keras.losses import SparseCategoricalCrossentropy import numpy as np from keras.callbacks import Callback, CallbackList batch_end_loss = list() class ShowBatchLoss(tf.keras.callbacks.Callback): def on_train_batch_end(self, batch, logs=None): if 'loss' in logs: batch_end_loss.append(logs['loss']) print(f"Batch {batch} loss: {logs['loss']:.4f}") callbacks = [ShowBatchLoss()] clist = CallbackList(callbacks=callbacks) # 构建模型 inputs = Input(shape=(784,), name="digits") x1 = Dense(64, activation="relu")(inputs) x2 = Dense(64, activation="relu")(x1) outputs = Dense(10, name="predictions")(x2) model = Model(inputs=inputs, outputs=outputs) optimizer = SGD(learning_rate=1e-3) loss_fn = SparseCategoricalCrossentropy(from_logits=True) model.compile(optimizer, loss_fn) # 准备数据 batch_size = 64 (x_train, y_train), (x_test, y_test) = tf.keras.datasets.mnist.load_data() x_train = np.reshape(x_train, (-1, 784)) x_test = np.reshape(x_test, (-1, 784)) x_val = x_train[-10000:] y_val = y_train[-10000:] x_train = x_train[:-10000] y_train = y_train[:-10000] train_dataset = tf.data.Dataset.from_tensor_slices((x_train, y_train)) train_dataset = train_dataset.shuffle(buffer_size=1024).batch(batch_size) val_dataset = tf.data.Dataset.from_tensor_slices((x_val, y_val)) val_dataset = val_dataset.batch(batch_size) # 初始化回调列表 clist.set_model(model) clist.set_params({ 'epochs': 2, 'steps': len(train_dataset), 'verbose': 1, 'do_validation': False, 'metrics': [], }) epochs = 2 clist.on_train_begin() # 启动训练钩子 for epoch in range(epochs): print("\nStart of epoch %d" % (epoch,)) clist.on_epoch_begin(epoch) # 启动epoch钩子 logs = {} for step, (x_batch_train, y_batch_train) in enumerate(train_dataset): clist.on_train_batch_begin(step, logs) # 启动批次计时 # 获取当前批次损失并存入logs batch_loss = model.train_on_batch(x_batch_train, y_batch_train) logs['loss'] = batch_loss clist.on_train_batch_end(step, logs) # 结束批次回调 clist.on_epoch_end(epoch, logs) # 结束epoch钩子 clist.on_train_end() # 结束训练钩子 # 查看收集的损失 print("\nCollected batch losses:", batch_end_loss[:5])
方案2:简化写法,直接调用自定义回调方法
如果不需要CallbackList的额外功能(比如计时、多回调管理),可以直接实例化自定义回调,手动调用其on_train_batch_end方法,更简洁:
# 其他代码(模型构建、数据准备)和上面一致 show_loss_callback = ShowBatchLoss() show_loss_callback.set_model(model) epochs = 2 for epoch in range(epochs): print("\nStart of epoch %d" % (epoch,)) for step, (x_batch_train, y_batch_train) in enumerate(train_dataset): batch_loss = model.train_on_batch(x_batch_train, y_batch_train) # 手动传入包含损失的logs show_loss_callback.on_train_batch_end(step, {'loss': batch_loss}) print("\nCollected batch losses:", batch_end_loss[:5])
关键修正点说明
- 补全钩子调用:必须按顺序调用
on_train_begin→on_epoch_begin→on_train_batch_begin→on_train_batch_end→on_epoch_end→on_train_end,确保回调内部状态正确初始化 - 传递损失到logs:
train_on_batch返回当前批次的损失值,需要将其存入logs['loss'],才能让自定义回调读取到 - 初始化CallbackList参数:通过
clist.set_params传入训练的基本参数,确保回调逻辑正常运行
内容的提问来源于stack exchange,提问作者ailauli69
相关产品推荐
相关产品推荐

