使用Scikit-learn包的set_callbacks
【摘要】 你提到的 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)