scikit-learn 开发者 API 重构:构建稳定的第三方开发生态
在机器学习开源生态中,第三方开发者依赖 scikit-learn 构建自定义 Estimator 是常态。然而,scikit-learn 长期以来存在一个痛点:其公共 API (Public API) 对向后兼容性有着近乎苛刻的标准(通常需提前一年通知),而私有 API (Private API) 则因内部迭代频繁变动,导致基于私有 API 开发的第三方工具极易断裂。
为了解决这一矛盾,scikit-learn 团队决定在公共 API 与私有 API 之间,构建一个全新的开发者 API (Developer API)。这一中间层旨在提供比公共 API 更丰富的功能,同时保持比私有 API 更高的稳定性,通常只需一个发布周期(Release Cycle)的警告即可进行变更。
核心突破:测试基础设施与标签系统升级
在 1.6 版本中,团队重点重构了测试基础设施和 Estimator 标签系统,使开发者能更清晰地定义模型行为。
1. 新的标签系统 (__sklearn_tags__)
旧式的私有标签(如 _more_tags)已被标准化为公开的 __sklearn_tags__ 方法。开发者现在可以通过继承 BaseEstimator 和 ClassifierMixin,并实现 __sklearn_tags__ 来动态定义模型属性。
from sklearn.base import BaseEstimator, ClassifierMixin
class MyEstimator(ClassifierMixin, BaseEstimator):
def __sklearn_tags__(self):
tags = super().__sklearn_tags__()
# 设置模型为非确定性
tags.non_deterministic = True
return tags
2. 测试跳过机制的现代化
原有的 _xfail_checks 标签已被废弃。现在,开发者需要直接通过 check_estimator 或 parametrize_with_checks 函数传递 expected_failed_checks 参数,以明确告知测试框架哪些检查预期会失败。
from sklearn.utils.estimator_checks import check_estimator, parametrize_with_checks
CHECKS_EXPECTED_TO_FAIL = {
"check_to_skip_name": "this check is known to fail"
}
def test_with_check_estimator():
check_estimator(MyEstimator(), expected_failed_checks=CHECKS_EXPECTED_TO_FAIL)
实用工具:sklearn_compat 包
由于测试规则的变更,部分旧代码可能需要调整。为了平滑过渡,团队推出了 sklearn_compat 包。开发者可以选择将其作为依赖项安装,或将单文件 vendoring 到自己的项目中,以自动处理 API 变更带来的兼容性冲突。
开发者价值总结
此次更新不仅解决了 API 稳定性问题,还显著提升了测试效率。通过标准化的标签系统,开发者可以更精确地控制模型的测试行为,减少因版本升级导致的代码重构成本。
“我们一直在努力创建一个介于公共 API 和私有 API 之间的开发者 API,旨在保持其稳定性,并在必要时引入一个发布周期的警告。” —— scikit-learn 开发团队