ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

scikit-learn 回调 API 开发指南:为自定义估计器实现回调支持与开发回调

scikit-learn 回调 API 开发指南:为自定义估计器实现回调支持与开发回调 scikit-learn 回调 API 开发指南为自定义估计器实现回调支持与开发回调【免费下载链接】scikit-learnscikit-learn: machine learning in Python项目地址: https://gitcode.com/gh_mirrors/sc/scikit-learnscikit-learn 提供了一套完整的 :term:callback回调API让内置与第三方估计器都能以统一方式注册回调在fit的各个阶段执行自定义逻辑。本文面向希望实现回调或为估计器添加回调支持的开发者系统讲解sklearn.callback模块中的回调协议、任务树模型、上下文管理与回调传播机制并给出可直接运行的估计器与回调最小实现。本文对应的官方开发文档为 callbacks.rst其完整内容由 callback_support.rst为估计器实现回调支持与 developing_callbacks.rst开发回调两部分组成。若你只想了解回调的使用方式请参阅 用户指南本文聚焦开发视角。回调 API 全景协议、钩子与任务树在 scikit-learn 中回调是以类形式实现的对象必须遵循一个基于 Pythontyping.Protocol的协议。协议要求回调类实现一组特定的方法——称为回调钩子hooks——这些钩子会在估计器或元估计器拟合过程的特定时刻被调用setup(estimator, context)仅在fit开始时调用一次负责初始化回调例如分配资源on_fit_task_begin(estimator, context, *, X, y, metadata, fitted_estimator)在每个任务开始时调用on_fit_task_end(estimator, context, *, X, y, metadata, fitted_estimator) - bool在每个任务结束时调用返回布尔值表示是否请求中断拟合teardown(estimator, context)仅在fit结束时调用一次负责清理回调例如释放资源。整个框架的公开入口集中在sklearn/callback/__init__.py导出了FitCallback、AutoPropagatedCallback、CallbackContext、CallbackSupportMixin、ProgressBar、ScoringMonitor、ScoringMonitorLog与with_callbacks等对象。其中任务task是理解整个框架的关键概念。一个任务是由估计器定义的任意工作单元——通常是学习算法的一次迭代也可以是流水线Pipeline的一个步骤、交叉验证的一折等。由于任务可以分解为子任务任务之间自然形成树状结构根任务就是整个fit。回调上下文CallbackContext正是围绕这棵任务树组织的每个上下文实例对应一个任务保存其父/子上下文引用并在fit过程中动态构建。FitCallback 协议回调的契约所有与 scikit-learn 估计器兼容的回调都必须实现 sklearn/callback/_base.py 中定义的FitCallback协议class FitCallback(Protocol): def setup(self, estimator, context) - None: ... def on_fit_task_begin( self, estimator, context, *, XNone, yNone, metadataNone, fitted_estimatorNone ) - None: ... def on_fit_task_end( self, estimator, context, *, XNone, yNone, metadataNone, fitted_estimatorNone ) - bool: ... def teardown(self, estimator, context) - None: ...协议中定义的方法即为回调钩子它们在估计器拟合过程中的调用时机如下setup与teardown各只调用一次分别位于fit的开始与结束。它们处理回调的建立与拆除例如分配和释放资源。on_fit_task_begin与on_fit_task_end在fit过程中每个任务的开头与结尾被调用。FitCallback协议在源码中以runtime_checkable标注见 sklearn/callback/_base.py这意味着isinstance运行时检查可用于校验回调实例。校验逻辑位于 sklearn/callback/_callback_support.py 的validate_callbacks函数它会检查回调必须是FitCallback的实例协议对类本身也会返回True因此额外排除了类对象每个钩子的位置参数必须恰好是[estimator, context]setup/teardown除了这两个位置参数外不能再有其他参数on_fit_task_begin/on_fit_task_end中仅允许X、y、metadata、fitted_estimator这几个关键字参数。钩子签名中的关键字参数约定在回调的具体实现中钩子应只声明实际使用到的关键字参数。参数在签名中的存在即表明该钩子需要这个参数这允许回调框架避免计算任何已注册回调都不使用的值——这是性能设计的关键一环。警告这些可选参数必须定义为仅限关键字keyword-only。如果不是 keyword-only参数值将不会被提供给钩子。同时即便钩子请求了某个参数估计器也可能因为在该任务上无法产出该值而不提供它。因此钩子的实现不应假设每次调用都能收到每个可选参数的值而应根据实际情况调整行为。钩子接收的强制位置参数是调用回调的估计器实例与CallbackContext对象后者持有足够唯一定位当前任务的上文信息以公开属性暴露。关于estimator参数需要注意钩子收到的估计器实例处于调用钩子那一刻的状态因此它不一定是完全拟合的teardown钩子除外。回调不应依赖该实例直接进行predict、transform等操作而应优先使用fitted_estimator当可用时。通过 on_fit_task_end 中断 fiton_fit_task_end钩子返回一个布尔值当为True时表示请求估计器在该任务处停止fit过程。需要说明的是不打算支持中断的估计器会忽略这个请求并继续下一个任务。任务与 CallbackContext 的任务树模型为估计器添加回调支持本质上就是三件事启用回调注册、将fit表达为一棵任务树、在每个任务的开始与结束处调用回调。为此sklearn.callback模块提供了三个核心工具CallbackSupportMixin启用回调注册并在fit开始时初始化回调处理CallbackContext代表任务是fit期间管理回调的核心对象with_callbacks保证fit结束时回调被正确拆除。任务的树状结构在文档中有两个典型示例。第一个是KMeans它有两层嵌套循环外层由n_init控制内层由max_iter控制因此任务树如下KMeans fit (root) ├── init 0 │ ├── iter 0 │ ├── iter 1 │ ├── ... │ └── iter n ├── init 1 │ ├── iter 0 │ ├── ... │ └── iter n └── init 2 ├── iter 0 ├── ... └── iter n其中每个最内层的iter j任务对应一次对完整数据集计算标签与中心的迭代。注册在 KMeans 估计器上的回调因此会在fit任务、每个外层init i任务和每个内层iter j任务的开始与结束处被调用。按约定出于性能与估计器间一致性考虑scikit-learn 估计器任务树的最内层任务即叶子对应对完整输入数据的操作增量式估计器则对应 batch。第二个示例是元估计器场景。当估计器是元估计器时一个任务叶子通常对应拟合一个子估计器。此时该叶子任务与子估计器的根任务实际上代表同一个任务两者会被合并为单个任务元估计器与子估计器的任务树被组合成一棵完整任务树。例如Pipeline的任务树如下Pipeline fit (root) ├── step 0 | StandardScaler fit │ └── insert StandardScaler task tree here └── step 1 | LogisticRegression fit └── insert LogisticRegression task tree here在源码层面这一合并机制实现在 sklearn/callback/_callback_context.py 的CallbackContext._merge_with方法中子估计器的根上下文在创建时检测到_parent_callback_ctx属性由元估计器通过propagate_callback_context注入便将自身与元估计器的叶子上下文合并——继承其父节点、任务 id 与root_uuid并记录合并来源source_estimator_name/source_task_name。CallbackContext的关键公开属性包括task_name任务名、task_id在兄弟任务中唯一标识、max_subtasks子任务最大数量0表示叶子None表示未知、sequential_subtasks子任务是否顺序执行、estimator_name、parent父上下文、root_uuid同一任务树共享的 UUID、init_time上下文初始化时间UTC等。CallbackSupportMixin让估计器支持回调的第一步要让一个估计器支持回调它必须继承CallbackSupportMixin该类见 sklearn/callback/_callback_support.py暴露以下方法set_callbacks(*callbacks)公开方法由用户调用以在估计器上注册回调。它先通过validate_callbacks校验回调合法性再调用_set_callbacks存储回调列表存入_skl_callbacks属性。_init_callback_context(...)应在fit开头调用创建对应整个fit任务的根CallbackContext并设置估计器上注册的回调即调用其setup钩子。注意虽然下划线前缀表明_init_callback_context是内部用途、不应出现在面向终端用户的自动补全建议中但它是对构建第三方估计器的开发者开放的应被视为公开 API 契约的一部分。_init_callback_context的签名与参数如下def _init_callback_context( self, task_namefit, task_id0, max_subtasks0, sequential_subtasksTrue ):task_name根任务名称默认fittask_id根任务标识默认0max_subtasks根任务可以拥有的子任务最大数量0表示叶子None表示事先未知默认0sequential_subtasks根上下文的子任务是否顺序执行。为True时通过subcontext创建的子上下文会自动获得从 0 开始的连续整数task_id。该方法内部调用CallbackContext._from_estimator创建根上下文并逐个为回调调用setup钩子。需要注意从元估计器传播而来的自动传播回调AutoPropagatedCallback类型且带有_parent_callback_ctx不会在此重复调用setup——setup只由最外层估计器调用一次。CallbackContextfit 期间管理回调的核心对象CallbackContext对象负责在正确时机调用回调。它们跟踪估计器的各个任务每个任务对应一个上下文实例并捕获fit执行过程中任务树的整体结构。重要回调上下文不应直接实例化而必须通过CallbackSupportMixin._init_callback_context创建根上下文或CallbackContext.subcontext创建子上下文来创建。subcontext创建子任务上下文def subcontext( self, task_name, task_idNone, max_subtasks0, sequential_subtasksTrue ):task_name子任务名称默认空字符串task_id子任务标识必须与兄弟任务不同为None时自动取下一个可用整数默认Nonemax_subtasks该子任务可拥有的子任务最大数量0表示叶子None表示未知默认0sequential_subtasks该子任务的子任务是否顺序执行默认True。注意两条约束源码在 sklearn/callback/_callback_context.py 中强制执行若当前上下文sequential_subtasksTrue则task_id必须为None反之若为False则task_id必须显式提供。重复的task_id或超出max_subtasks的子任务都会抛出ValueError。call_on_fit_task_begin / call_on_fit_task_end触发钩子这两个方法必须分别在上下文所对应任务的开始与结束处调用def call_on_fit_task_begin( self, *, estimator, XNone, yNone, metadataNone, reconstruction_attributesNone ) - None: ... def call_on_fit_task_end( self, *, estimator, XNone, yNone, metadataNone, reconstruction_attributesNone ) - bool: ...正如方法名所示它们会调用估计器上注册回调的on_fit_task_begin/on_fit_task_end钩子。除隐式传给已注册回调的回调上下文外这些方法的关键字参数用于传递拟合过程在某个任务上的附加状态信息。并非每次调用都需要提供所有参数——估计器应提供自己在该任务上能够产出的所有值回调则根据给定任务实际提供的值调整自身行为。reconstruction_attributes重建可评估的拟合估计器当调用call_on_fit_task_begin/end时估计器在该任务处的状态很可能不完整无法进行predict、transform等操作。reconstruction_attributes参数期望一个字典包含设置到估计器上以补齐缺失状态的属性从而保证估计器如同 fit 在此任务处停止那样可以立即predict、transform等。回调上下文会复制估计器在该任务时的状态设置这些重建属性并把得到的估计器作为fitted_estimator传给回调。这一逻辑实现在_from_reconstruction_attributes函数中sklearn/callback/_callback_context.py它用copy.copy浅拷贝估计器并设置属性避免修改真实估计器。如果不需要额外属性即可使估计器就绪应传入空字典而非保持默认值否则回调上下文不会向回调传递fitted_estimator。关键字参数的惰性求值对于这些关键字参数中的每一个都可以传入一个无参可调用对象callable替代实际值。当某个回调需要该参数时回调上下文才会求值该 callable 并转发返回值。这一机制实现了参数的惰性求值避免在没有回调需要某参数值的情况下进行潜在的高成本计算。为阻止性能退化估计器应惰性传递计算代价高昂的量典型做法是传入lambda: ...参见后文的SimpleKMeans完整示例。中断 fitcall_on_fit_task_end返回一个布尔值若任何回调在该任务结束时发出停止fit的信号则返回True否则返回False。该返回值可用于中断当前层级的迭代例如实现早停。propagate_callback_context组合估计器与元估计器的任务树propagate_callback_context(sub_estimator)是CallbackContext的一个上下文管理器源码见 sklearn/callback/_callback_context.py用于在估计器组合例如GridSearchCV作用于LogisticRegression中把各估计器的上下文树组合成一棵以顶层估计器fit为根的单一上下文树。它应在元估计器中、对应拟合子估计器这个任务的上下文上使用。该任务既是元估计器的叶子任务又是子估计器的根任务因此两个对应的上下文会在组合树中被合并为单个上下文。除合并任务树外propagate_callback_context还承担两个职责传播自动传播回调把元估计器上的auto-propagated回调即实现了AutoPropagatedCallback协议的回调传播给子估计器使它们也会在子估计器的任务处被调用。传播深度受回调的max_propagation_depth限制。退出时清理在with块退出时清除传播到子估计器上的回调保证拟合后的子估计器不再持有任何本地注册的回调并删除注入的_parent_callback_ctx属性。此外有两个重要的边界行为若子估计器自身已经注册了自动传播回调会抛出TypeError提示应把这类回调直接注册在元估计器上若子估计器不支持回调没有set_callbacks方法会发出UserWarning警告并跳过传播该警告在根上下文上只会对同一对估计器组合发出一次。with_callbacks 装饰器保证回调正确拆除对于实现回调支持的第三方估计器fit方法应使用with_callbacks装饰器sklearn/callback/_callback_support.py。该装饰器在try/finally块中运行被装饰方法保证回调的teardown钩子总是被调用——即使fit因错误异常退出。with_callbacks内部使用的是callback_management_context上下文管理器sklearn/callback/_callback_support.py其 finally 分支会对_skl_callbacks_to_teardown中记录的回调逐个调用teardown清理_skl_callbacks_to_teardown与_callback_fit_ctx属性若teardown自身抛出异常单个异常直接抛出多个异常聚合为ExceptionGroup抛出。注意只有未被元估计器传播的回调会被拆除传播来的回调由元估计器统一管理。对于scikit-learn 内置估计器sklearn.base._fit_context装饰器已经承担了回调拆除工作因此内置估计器不应再使用with_callbacks。最小实现为自定义估计器添加回调支持结合 callback_support.rst 的示例与 examples/callbacks/plot_callback_support.py 的完整演示一个典型的回调支持实现如下from sklearn.callback import CallbackSupportMixin, with_callbacks class MyEstimator(CallbackSupportMixin): def __init__(self, max_iter): self.max_iter max_iter with_callbacks def fit(self, X, y): callback_ctx self._init_callback_context(max_subtasksself.max_iter) callback_ctx.call_on_fit_task_begin(estimatorself, XX, yy) for i in range(self.max_iter): subcontext callback_ctx.subcontext(task_nameiteration) subcontext.call_on_fit_task_begin(estimatorself, XX, yy) # Do something if subcontext.call_on_fit_task_end(estimatorself, XX, yy): break callback_ctx.call_on_fit_task_end(estimatorself, XX, yy) return self实现步骤可以总结为继承CallbackSupportMixin同时继承BaseEstimator以获得 scikit-learn 估计器基础能力用with_callbacks装饰fit在fit开头调用_init_callback_context(max_subtasks...)创建根上下文同时触发回调的setup对根任务调用call_on_fit_task_begin(estimatorself, ...)对每个子任务用subcontext(...)创建子上下文并分别调用其call_on_fit_task_begin/call_on_fit_task_end依据call_on_fit_task_end的返回值决定是否中断if ... : break在fit返回前对根上下文调用call_on_fit_task_end回调的teardown由装饰器自动完成。在 examples/callbacks/plot_callback_support.py 中SimpleKMeans给出了更贴近实战的完整版本它在每次迭代结束调用call_on_fit_task_end时传入reconstruction_attributeslambda: {cluster_centers_: self.cluster_centers_}通过惰性求值把当前质心属性交给框架使回调获得可预测/可转换的fitted_estimator注册回调只需一行estimator SimpleKMeans(random_staterng) callback ProgressBar() estimator.set_callbacks(callback) estimator.fit(X)元估计器的回调支持与传播同一个示例脚本还展示了元估计器SimpleGridSearch的实现要点它与普通估计器的核心差异在于必须通过propagate_callback_context把回调传播给子估计器with inner_subcontext.propagate_callback_context(cloned_estimator): inner_subcontext.call_on_fit_task_begin(estimatorcaller, XX_train, yy_train) cloned_estimator.fit(X_train, y_train) scores_per_fold.append(score_func(cloned_estimator, X_test, y_test)) inner_subcontext.call_on_fit_task_end(estimatorcaller, XX_train, yy_train)值得注意的是在并行化元估计器中所有外层子上下文必须在并行之前一次性创建以避免创建过程中的竞争条件示例脚本对此有明确注释。另外一个支持回调的元估计器可以搭配不支持回调的子估计器使用——此时传播会发出警告并在子估计器中被忽略。开发自定义回调从最小示例开始developing_callbacks.rst 给出了一个每次被调用都会打印消息的简单回调class MyCallback: def setup(self, estimator, context): print(fSetup hook is being called in the {context.task_name} task.) def teardown(self, estimator, context): print(fTeardown hook is being called in the {context.task_name} task.) def on_fit_task_begin(self, estimator, context, *, XNone): msg f{context.task_name} task is starting. if X is not None: msg f With training data of shape {X.shape}. print(msg) def on_fit_task_end( self, estimator, context, *, XNone, yNone, fitted_estimatorNone ): msg f{context.task_name} task is ending. mean_squared_error ((y - fitted_estimator.predict(X))**2).mean() msg f With a mean squared error of {mean_squared_error}. print(msg)这个示例展示了开发回调时的几个要点通过context.task_name等上下文属性识别当前任务实现在任务树中定位仅声明钩子实际使用的关键字参数X、y、fitted_estimator让框架可以跳过无用的计算on_fit_task_end使用fitted_estimator进行predict等评估操作而不是依赖可能不完整的estimator实例。回调开发与使用方式的完整测试与行为约束可以参考 sklearn/callback/tests 下的测试套件例如test_callback_support.py、test_callback_context.py、test_progressbar.py与test_scoring_monitor.py。自动传播回调AutoPropagatedCallback所谓auto-propagated自动传播回调是指期望从元估计器传播到其子估计器的回调。这类回调必须实现AutoPropagatedCallback协议——它是FitCallback协议的扩展源码在 sklearn/callback/_base.pyclass AutoPropagatedCallback(FitCallback, Protocol): property def max_propagation_depth(self) - int | None: ...与只在注册它的估计器任务上被调用的普通回调不同自动传播回调会在估计器组合中所有估计器的任务上被调用直至最大传播深度max_propagation_depth 0不传播到子估计器只在顶层估计器的任务上被调用max_propagation_depth None传播到所有嵌套层级的子估计器其他整数值传播到最多 N 层嵌套估计器。自动传播回调应注册在顶层估计器上。如果顶层估计器不支持回调可以注册在子估计器上通常也能工作只是可能无法发挥全部能力。注意自动传播回调的setup与teardown钩子同样只在最外层估计器的fit开始与结束时调用一次不会为任何子估计器重复调用。回调共享状态跨多次 fit 累积数据由于回调注册所在的估计器可能被元估计器克隆并多次拟合回调应表现得如同同一个回调实例被注册到了多个估计器上。因此setup/teardown不应重置回调的状态on_fit_task_begin/on_fit_task_end应跨所有 fit 累积数据。原因在于若在setup/teardown中重置状态会丢失之前或并发 fit 中收集的信息例如嵌套交叉验证中外层 fit 与内层 fit 的监控数据。内置回调ProgressBar 与 ScoringMonitor除框架基础设施外sklearn.callback还提供了两个开箱即用的回调它们的源码分别位于 sklearn/callback/_progressbar.py 与 sklearn/callback/_scoring_monitor.py实现细节可作为开发高质量回调的参考范本。ProgressBar迭代进度条ProgressBar为估计器的每个迭代步骤显示进度条其唯一构造参数为max_propagation_depth显示进度条的嵌套估计器层级最大深度int默认10表示只显示最外层估计器的进度None表示显示所有层级。它是一个自动传播回调传播深度默认 1因此在Pipeline、网格搜索等元估计器组合中会自动传播到子估计器。使用它需要安装rich库check_rich_support会校验。从源码看其实现采用主进程监听 rich 监控线程的传输机制setup时打开监听器并启动RichProgressMonitor线程任务事件通过队列异步转发从而避免在整个任务树上进行昂贵的对象传递进度显示本身按任务树动态创建嵌套的 rich 任务非叶子任务都有对应进度条。ScoringMonitor逐任务记录评分ScoringMonitor在估计器的每个迭代步骤用指定的 scorer 在训练数据上计算评分并记录日志可通过get_logs检索scoring监控模型所用的评分方法。单个评分可以用字符串如neg_mean_squared_error或返回单值的 callable多个评分可以用字符串列表/元组、返回字典的 callable或指标名 - callable的字典。内部通过check_scoring统一转换为多指标评分器。get_logs(selectmost_recent, include_lineageFalse)返回ScoringMonitorLogselectmost_recent或ScoringMonitorLog列表selectall。include_lineageTrue时为每个记录评分的任务的祖先任务补充额外行便于恢复完整上下文data列包含task_id_path、parent_task_id_path、estimator_name、task_name、task_id、sequential_subtasks以及每个评分名列。ScoringMonitorLog的data是 dict 列表data_as_pandas属性可将其转为 pandas DataFrame需要 pandas 已安装。日志按run最外层fit调用分组若注册回调的估计器被包裹在元估计器中一个 run 对应最外层元估计器的一次fit否则对应估计器自身的一次fit。run_id与root_uuid一一对应可跨任务定位单次拟合的全部评分记录。其配套示例见 examples/callbacks/plot_scoring_monitor.py。仓库中的真实应用当前仓库中已有多个内置模块实际应用了这套回调框架可直接作为学习参考sklearn/linear_model/_logistic.py逻辑回归LogisticRegression在拟合中支持回调sklearn/model_selection/_search.py网格搜索与随机搜索作为元估计器使用propagate_callback_context向子估计器传播回调sklearn/pipeline.pyPipeline组合各步骤的子估计器任务树sklearn/preprocessing/_data.py部分预处理器如StandardScaler、MinMaxScaler等在partial_fit类场景下使用回调。对应的测试用例如sklearn/model_selection/tests/test_search.py、sklearn/linear_model/tests/test_logistic.py、sklearn/preprocessing/tests/test_data.py覆盖了回调注册、传播、中断与状态清理等行为是验证实现正确性的最佳参照。小结scikit-learn 的回调 API 将在拟合过程中插入自定义逻辑这一需求抽象为一棵任务树 上下文管理的完整框架开发者通过实现FitCallback协议创建回调通过继承CallbackSupportMixin与调用CallbackContext的钩子触发方法让估计器支持回调借助with_callbacks保证生命周期安全再通过propagate_callback_context让回调在元估计器组合中自动传播。掌握这套机制后无论是为第三方估计器接入进度监控、评分记录还是实现自定义早停、日志等逻辑都可以在不侵入核心算法代码的前提下完成。更多细节可继续阅读 callback_support.rst、developing_callbacks.rst 以及完整示例 plot_callback_support.py 和 plot_scoring_monitor.py。【免费下载链接】scikit-learnscikit-learn: machine learning in Python项目地址: https://gitcode.com/gh_mirrors/sc/scikit-learn创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表