自定义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
相关产品推荐
相关产品推荐

