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

自定义Keras回调统计模型训练总时长异常,求排查与修复方案

问题

需要使用relu、tanh、sigmoid三种激活函数,将模型各训练n次以开展性能统计评估,其中一项评估内容是各模型的总训练时长。为此编写了如下自定义Keras回调:

class TimingCallback(keras.callbacks.Callback):
    def __init__(self, iteracao, func, logs=None):
        self.iteracao = iteracao
        self.func = func
        self.starttime = None
        if logs is None:
            logs = {'relu': {}, 'tanh': {}, 'sigmoid': {}}
        self.logs = logs

    def on_train_begin(self, epoch, logs={}):
        self.starttime = time.time()

    def on_train_end(self, epoch, logs={}):
        self.logs[self.func].update({self.iteracao: (time.time() - self.starttime)/60})

并通过嵌套循环执行训练:

n = 10
dic_loss_train = {'relu': {}, 'tanh': {}, 'sigmoid': {}}
for funcao in dic_loss_train:
    for i in range(n):
        model = Sequential()
        model.add(...)
        model.compile(...)
        tempo_callback = TimingCallback(i, funcao)
        cp = ModelCheckpoint(...)
        earlystop = EarlyStopping(...)
        treinando = model.fit(x_train, y_train, batch_size=32,
                              epochs=5, verbose=2,
                              callbacks=[cp, earlystop, tempo_callback],
                              validation_split=val_size, shuffle=False)

预期得到包含各激活函数每次训练时长的字典,但实际仅得到:

{'relu': {}, 'tanh': {}, 'sigmoid': {9: 0.0273}}

疑问:on_train_end中的update方法存在问题,错误原因是什么?该如何修复?

错误原因
  • 日志字典不共享:每次创建TimingCallback实例时,若未传入logs参数,都会生成一个全新的内部字典。前9次训练的时长数据都存在各自实例的私有字典中,最终只有最后一次实例的字典被保留,导致结果缺失。
  • 未关联外部目标字典:循环中用于存储结果的dic_loss_train从未传入回调,回调操作的是自身内部的字典,而非外部的目标字典,所以外部字典无法收集到所有训练数据。
  • 回调方法参数错误:on_train_begin和on_train_end的参数写错了,Keras规定这两个方法的参数是logs而非epoch,虽然可能不影响计时,但属于规范错误,可能导致方法触发异常。
修复方案

1. 修正回调类

移除内部默认创建字典的逻辑,强制传入外部目标字典,同时修正回调方法的参数:

class TimingCallback(keras.callbacks.Callback):
    def __init__(self, iteracao, func, logs):
        self.iteracao = iteracao
        self.func = func
        self.starttime = None
        self.logs = logs  # 直接使用外部传入的字典

    def on_train_begin(self, logs=None):
        self.starttime = time.time()

    def on_train_end(self, logs=None):
        elapsed_time = (time.time() - self.starttime) / 60
        self.logs[self.func][self.iteracao] = elapsed_time  # 直接赋值,替代update更简洁

2. 修改循环中的回调实例化

创建TimingCallback时,将外部的dic_loss_train传入logs参数,让所有回调实例共享同一个字典:

n = 10
dic_loss_train = {'relu': {}, 'tanh': {}, 'sigmoid': {}}
for funcao in dic_loss_train:
    for i in range(n):
        model = Sequential()
        model.add(...)
        model.compile(...)
        # 传入外部字典作为日志存储容器
        tempo_callback = TimingCallback(i, funcao, logs=dic_loss_train)
        cp = ModelCheckpoint(...)
        earlystop = EarlyStopping(...)
        treinando = model.fit(x_train, y_train, batch_size=32,
                              epochs=5, verbose=2,
                              callbacks=[cp, earlystop, tempo_callback],
                              validation_split=val_size, shuffle=False)

修改完成后,dic_loss_train会正确收集到每个激活函数下10次训练的时长数据。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.19 03:00:52