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

调用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])

关键修正点说明

  1. 补全钩子调用:必须按顺序调用on_train_begin→on_epoch_begin→on_train_batch_begin→on_train_batch_end→on_epoch_end→on_train_end,确保回调内部状态正确初始化
  2. 传递损失到logs:train_on_batch返回当前批次的损失值,需要将其存入logs['loss'],才能让自定义回调读取到
  3. 初始化CallbackList参数:通过clist.set_params传入训练的基本参数,确保回调逻辑正常运行

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.18 15:07:45