【问题标题】:How to access history objects when cross validating keras wrapper estimators in sklearn?在sklearn中交叉验证keras包装器估计器时如何访问历史对象?
【发布时间】:2018-09-23 18:08:45
【问题描述】:

我想查看交叉验证中每个拆分的丢失/错误进展。 keras.wrappers.scikit_learn.KerasClassifier 的 fit 方法返回一个带有我想要的数据的 history 对象,但是在 sklearn.model_selection.cross_validate 变体方法中运行它时无法访问它。

如何访问每个拆分中每个时期的历史对象?

【问题讨论】:

  • cross_validate 克隆提供的模型以适应每个折叠。所以你不能得到那个。您需要滚动自己的交叉验证代码来查看每个折叠拆分的历史记录。
  • @Alex 你有没有找到更好的方法?

标签: scikit-learn keras


【解决方案1】:

您或许可以使用CSVLogger 回调来访问完整的历史记录。设置 CSVLogger 回调很简单,它会在您指定的任何文件名中记录:{epoch, acc, loss, val_acc, val_loss}。

在我的代码中,我做了类似的事情:

keras_classifier.fit(X, y, groups=None, 
    callbacks=[keras.callbacks.CSVLogger(filename, append=True)])

设置append=True 应确保所有拆分的所有数据都包含在文件中。

需要考虑的事项:

  • 我不确定这是否适用于 n_jobs=-1(用于在多个处理器上分配处理),但如果您运行单线程,它应该可以工作。
  • 确保在运行分类器之前(或在初始化期间)删除该文件,以避免无限期地附加到该文件。

【讨论】:

    猜你喜欢
    • 2019-06-11
    • 2020-09-13
    • 2013-12-18
    • 2014-11-24
    • 2020-07-07
    • 2017-10-23
    • 2021-01-17
    • 2012-12-31
    • 2018-08-16
    相关资源
    最近更新 更多