使用Scikit-learn包的set_callbacks

举报
yd_37369233 发表于 2026/09/10 19:28:06 2026/09/10
【摘要】 你提到的 Swith_callbacks 应该是拼写上的小笔误,在 Scikit-learn 中与之最相关的功能是回调(Callbacks)机制,其核心方法为 set_callbacks,而装饰器 with_callbacks 则是供开发者使用的。该功能在 Scikit-learn 1.9 版本中作为实验性特性引入。 回调机制概述回调是可以在估计器上注册的对象,用于在 fit 过程中的关键步...

回调机制概述

回调是可以在估计器上注册的对象,用于在 fit 过程中的关键步骤(开始/结束)插入自定义逻辑,例如监控进度或计算指标,而无需修改底层学习算法。目前仅部分估计器支持回调,可通过 Scikit-learn 官方文档查看兼容估计器列表。

核心用法:set_callbacks 方法

支持回调的估计器会提供 set_callbacks 方法,用于注册一个或多个回调对象。注册后,回调会在 fit 调用期间被触发。例如,LogisticRegression 已支持该功能:

from sklearn.linear_model import LogisticRegression
from sklearn.callback import ProgressBar

model = LogisticRegression(max_iter=1000)
model.set_callbacks(ProgressBar())   # 注册进度条回调
model.fit(X, y)

内置回调

Scikit-learn 1.9 提供了两个开箱即用的回调类:

回调类 用途
ProgressBar 显示拟合过程的进度条
ScoringMonitor 计算并记录评分指标,可通过 get_logs().data_as_pandas 获取日志

组合使用示例:

from sklearn.callback import ProgressBar, ScoringMonitor
from sklearn.linear_model import LogisticRegression

scoring_monitor = ScoringMonitor(scoring="d2_log_loss_score")
logreg = LogisticRegression(solver="lbfgs")
logreg.set_callbacks(scoring_monitor, ProgressBar())
logreg.fit(X, y)

# 获取评分日志
log = scoring_monitor.get_logs().data_as_pandas
print(log[["task_name", "task_id", "d2_log_loss_score"]])

自定义回调

要创建自定义回调,需实现 FitCallback 协议,主要包含四个钩子方法:

钩子方法 调用时机
setup(estimator, context) 拟合开始时调用一次
on_fit_task_begin(estimator, context, **kwargs) 每个任务开始时调用
on_fit_task_end(estimator, context, **kwargs) 每个任务结束时调用,返回 True 可请求提前停止拟合
teardown(estimator, context) 拟合结束时调用一次,用于资源清理

自定义回调的最小示例:

from sklearn.callback import FitCallback

class MyCallback:
    def setup(self, estimator, context):
        print("拟合开始")

    def on_fit_task_begin(self, estimator, context, **kwargs):
        pass

    def on_fit_task_end(self, estimator, context, **kwargs):
        return False  # 不中断

    def teardown(self, estimator, context):
        print("拟合结束")

# 使用
model.set_callbacks(MyCallback())

为估计器添加回调支持(开发者)

如果你在开发自定义估计器,需让它继承 CallbackSupportMixin,并使用 with_callbacks 装饰器修饰 fit 方法,以保证回调的清理钩子始终被调用。

from sklearn.callback import CallbackSupportMixin, with_callbacks

class MyEstimator(CallbackSupportMixin, BaseEstimator):
    @with_callbacks
    def fit(self, X, y=None):
        self._init_callback_context()   # 在 fit 开头初始化回调上下文
        # ... 训练逻辑 ...
        return self

注意事项

  • 回调 API 目前是实验性的,未来可能在不经过常规弃用周期的情况下发生变化。
  • on_fit_task_end 钩子返回 True 可请求中断拟合,但并非所有估计器都支持中断,不支持中断的估计器会忽略该请求。
  • 回调钩子中的可选参数必须使用仅关键字参数(keyword-only)定义,否则框架不会传入相应值。

如果你实际想询问的是某个特定拼写的函数或类,可以补充一下来源,我再帮你确认。

【版权声明】本文为华为云社区用户原创内容,未经允许不得转载,如需转载请自行联系原作者进行授权。如果您发现本社区中有涉嫌抄袭的内容,欢迎发送邮件进行举报,并提供相关证据,一经查实,本社区将立刻删除涉嫌侵权内容,举报邮箱: cloudbbs@huaweicloud.com
  • 点赞
  • 收藏
  • 关注作者

评论(0

0/1000
抱歉,系统识别当前为高风险访问,暂不支持该操作

全部回复

上滑加载中

设置昵称

在此一键设置昵称,即可参与社区互动!

*长度不超过10个汉字或20个英文字符,设置后3个月内不可修改。

*长度不超过10个汉字或20个英文字符,设置后3个月内不可修改。