Sklearn-源码解析-书-v1-0-二-
Sklearn 源码解析(书)v1.0(二)
-
显式参数列表提供了 IDE 自动补全、类型检查(mypy/pyright)、文档生成(Sphinx)的完整支持。
-
**kwargs会丢失参数名信息,导致文档生成器无法识别合法配置键,用户也无法获得类型提示。 -
10 个配置项数量固定且稳定,显式写入的维护成本极低,换来的是极佳的开发体验。
为什么 get_config() 返回浅拷贝而非深拷贝或只读代理?
-
配置值均为不可变类型(bool、int、str),浅拷贝已足够隔离。
-
深拷贝有性能开销;只读代理(
types.MappingProxyType)虽更优雅,但会破坏dict接口预期(如config['new_key'] = val会报错而非静默失败),增加用户认知负担。 -
浅拷贝是"最小惊讶原则"下的平衡点。
3.10 动手练习
-
追踪配置的生命周期
在 Python 交互环境中执行以下代码,观察配置在不同阶段的变化:
import sklearn from sklearn import get_config, set_config, config_context # 1. 查看初始配置 config1 = get_config() print('初始 working_memory:', config1['working_memory']) # 2. 修改配置后查看 set_config(working_memory=2048) config2 = get_config() print('修改后 working_memory:', config2['working_memory']) print('修改后 print_changed_only:', config2['print_changed_only']) # 3. 尝试修改 get_config 返回的字典 config2['working_memory'] = 9999 config3 = get_config() print('外部修改后实际配置:', config3['working_memory'])回答问题:
-
set_config(working_memory=2048)是否影响了print_changed_only?为什么? -
直接修改
get_config()返回的字典为什么不会改变实际配置? -
如果改用
_get_threadlocal_config()返回的字典进行修改会发生什么?(提示:阅读源码第28-36行)
-
-
验证 config_context 的恢复机制
设计一个实验验证 config_context 在异常情况下的恢复行为:
import sklearn from sklearn import get_config, config_context print('上下文外 assume_finite:', get_config()['assume_finite']) # 实验1:正常退出 with config_context(assume_finite=True): print('上下文内 assume_finite:', get_config()['assume_finite']) print('正常退出后 assume_finite:', get_config()['assume_finite']) # 实验2:异常退出 try: with config_context(assume_finite=True): print('异常前 assume_finite:', get_config()['assume_finite']) raise ValueError('测试异常') except ValueError: pass print('异常退出后 assume_finite:', get_config()['assume_finite']) # 实验3:嵌套上下文 with config_context(assume_finite=True): print('外层 assume_finite:', get_config()['assume_finite']) with config_context(assume_finite=False): print('内层 assume_finite:', get_config()['assume_finite']) print('内层退出后 assume_finite:', get_config()['assume_finite'])回答问题:
-
finally块在异常场景下是否执行了?结合源码第293-295行解释。 -
嵌套上下文的配置恢复顺序是什么?为什么内层退出后回到了外层配置而非全局配置?
-
-
探索线程局部存储的隔离性
编写多线程程序验证配置的线程隔离性:
import threading import time from sklearn import set_config, get_config def worker(thread_id, new_value, sleep_time): set_config(working_memory=new_value) time.sleep(sleep_time) print(f'线程{thread_id}: working_memory={get_config()["working_memory"]}') # 创建两个线程,分别设置不同的 working_memory 值 t1 = threading.Thread(target=worker, args=(1, 2048, 0.5)) t2 = threading.Thread(target=worker, args=(2, 4096, 0.1)) t1.start() t2.start() t1.join() t2.join() print(f'主线程: working_memory={get_config()["working_memory"]}')回答问题:
-
两个线程的配置是否相互影响?主线程的配置是否被修改?
-
结合
_get_threadlocal_config()源码解释为什么每个线程获得独立的配置副本。 -
如果移除
threading.local()改用普通全局变量,上述程序会有什么不同?
-
-
分析环境变量的读取时机
进行以下实验理解环境变量的"一次性读取"特性:
import os import sklearn # 实验1:正常导入 os.environ['SKLEARN_WORKING_MEMORY'] = '512' import sklearn._config as config_module print('首次导入 working_memory:', config_module._global_config['working_memory']) # 实验2:修改环境变量后重新导入 os.environ['SKLEARN_WORKING_MEMORY'] = '2048' import importlib importlib.reload(config_module) print('重载后 working_memory:', config_module._global_config['working_memory']) # 实验3:不重新导入,仅修改环境变量 os.environ['SKLEARN_WORKING_MEMORY'] = '4096' print('不重载的 working_memory:', config_module._global_config['working_memory'])回答问题:
-
环境变量在什么时机被读取?为什么修改环境变量后必须重新导入模块才能生效?
-
结合源码第14-27行,解释
int()和bool()类型转换的作用。 -
这种"导入时一次性读取"的设计有什么优缺点?
-
3.11 本章小结
这一章中我们学习/了解/讨论了 scikit-learn 全局配置体系的核心机制。首先,我们剖析了 _global_config 字典的模块级初始化,理解环境变量如何在导入时一次性注入"出厂设置";其次,我们深入 threading.local() 实现的线程局部存储,揭示 _get_threadlocal_config() 的懒初始化与写时隔离机制;接着,我们解读了 get_config() 返回浅拷贝的防御式设计,以及 set_config() 通过 None 语义实现的选择性微事务更新;然后,我们拆解了 config_context() 基于 @contextmanager 的快照-恢复模式,重点分析了 try/finally 保证的异常安全恢复与嵌套上下文的"俄罗斯套娃"语义;最后,我们梳理了配置项如何联动估计器行为——从 __repr__ 精简显示到 Transformer 输出格式切换,从分块算法的内存预算到元数据路由的渐进式迁移开关,再到 Array API 标准分派机制。
本章我们一起学习了以下概念:
| 概念 | 解释 |
|------|------|
| _global_config | 模块级默认配置字典,包含10项全局设置,在导入时从环境变量初始化 |
| threading.local() | 为每个线程提供独立的配置命名空间,避免多线程配置串扰 |
| _get_threadlocal_config() | 懒初始化线程本地配置,首次访问时从全局配置复制基线快照 |
| get_config() | 返回配置的浅拷贝,防止外部代码直接修改内部状态 |
| set_config() | 选择性更新配置,None 参数表示保持不变,支持10项独立微事务 |
| config_context() | 上下文管理器,进入时快照旧配置,退出时在 finally 中恢复 |
| _check_array_api_dispatch() | 唯一的 set_config 额外校验逻辑,验证 Array API 分派参数的合法性 |
| 嵌套 config_context | 俄罗斯套娃语义:内层退出恢复到外层配置,外层退出恢复全局配置 |
| print_changed_only | 控制估计器 repr 的精简显示,只打印非默认参数 |
| transform_output | 控制 Transformer 输出格式(default/pandas/polars) |
| working_memory | 内存预算限制,触发分块算法路径(单位MiB) |
| enable_metadata_routing | 元数据路由的新旧API兼容性开关,提供渐进式迁移路径 |
| array_api_dispatch | Array API 标准分派开关,启用后支持 CuPy/PyTorch/JAX 等后端 |
下一章中,我们将学习 scikit-learn 的异常与警告体系,探索如何通过语义化错误提示帮助用户快速定位和修复问题。
第 4 章 —— 异常与警告体系 —— 构筑“智能化的错误反馈系统”
4.1 学习目标
-
难度:★★★☆☆(3/5)
-
预备知识:Python 基础、面向对象编程与 Markdown/代码阅读基础
-
理解 scikit-learn 自定义异常与警告体系的设计动机与分层哲学
-
掌握 NotFittedError 双重继承(ValueError + AttributeError)的兼容性考量
-
了解元数据路由异常 UnsetMetadataPassedError 的精确报错机制
-
能区分不同类型警告(收敛、效率、数据转换、版本一致性等)的适用场景
-
能够阅读并扩展异常类,实现自定义错误反馈
-
理解异常类与警告类在继承层次与使用场景上的本质区别
4.2 生活类比
想象 scikit-learn 的异常与警告体系是一家医院的分诊与预警系统:NotFittedError = 手术前发现病人还没做术前检查,护士拦住说“先去做检查(fit)再来”;DataConversionWarning = 药房自动把药片碾碎成粉末,虽然能用但提醒你“剂型变了”;ConvergenceWarning = 康复训练做了 100 次但还没达到目标,医生说“可以出院,但建议继续练”;FitFailedWarning = 多科室会诊中某一科出了状况,总台提醒“某科室诊断失败,整体结论仍可用”;InconsistentVersionWarning = 用旧版病历本看新版医生,护士提醒“信息可能对不上,风险自负”;EstimatorCheckFailedWarning = 体格检查报告,详细记录每个未通过的项目和原因。异常是“硬性拦截”(不能继续),警告是“软性提醒”(可以继续但要注意),就像医院里“手术坚决不能做错”与“饮食建议仅供参考”的区别。
4.3 源码地图
sklearn/exceptions.py
├── all 列表(5-19行)
│ ├── 公共 API 声明:ConvergenceWarning、DataConversionWarning、NotFittedError 等 11 个类
│ └── 未导出项:InconsistentVersionWarning(暗示内部使用)
├── UnsetMetadataPassedError(22-40行)
│ ├── 继承 ValueError # 语义上属于输入值不合法
│ ├── docstring # 版本标注、参数说明
│ ├── init() # 仅关键字参数:message、unrequested_params、routed_params
│ └── 属性:unrequested_params、routed_params
├── NotFittedError(43-62行)
│ ├── 双重继承:ValueError, AttributeError
│ ├── docstring # 示例、版本迁移记录
├── ConvergenceWarning(65-68行)
│ ├── 继承 UserWarning
│ └── versionchanged 0.18 # 从 sklearn.utils 迁移
├── DataConversionWarning(70-91行)
│ ├── 继承 UserWarning
│ ├── 场景1:整数数组被隐式转换为浮点
│ ├── 场景2:请求无拷贝但实现需要拷贝
│ ├── 场景3:输入形状有歧义
│ └── versionchanged 0.18 # 从 sklearn.utils.validation 迁移
├── DataDimensionalityWarning(94-107行)
│ ├── 继承 UserWarning
│ ├── 场景:随机投影中目标维度高于原始维度
│ └── versionchanged 0.18 # 从 sklearn.utils 迁移
├── EfficiencyWarning(110-121行)
│ ├── 继承 UserWarning
│ ├── versionadded 0.18
│ └── 可被子类化为更具体的 Warning
├── FitFailedWarning(124-134行)
│ ├── 继承 RuntimeWarning # 运行时环境问题,而非用户代码错误
│ ├── 用于 GridSearchCV、RandomizedSearchCV、cross_val_score
│ └── versionchanged 0.18 # 从 sklearn.cross_validation 迁移
├── SkipTestWarning(137-143行)
│ ├── 继承 UserWarning
│ └── 可选依赖缺失时测试跳过而非失败
├── UndefinedMetricWarning(146-149行)
│ ├── 继承 UserWarning
│ ├── 指标在零除或空预测时未定义
│ └── versionchanged 0.18 # 从 sklearn.base 迁移
├── PositiveSpectrumWarning(152-160行)
│ ├── 继承 UserWarning
│ ├── PSD 矩阵出现显著负特征值时触发
│ └── versionadded 0.22
├── InconsistentVersionWarning(163-184行)
│ ├── 继承 UserWarning
│ ├── init() # 接收 estimator_name、current/old 版本号
│ ├── 属性:estimator_name、current_sklearn_version、original_sklearn_version
│ └── str() # 输出风险警告与官方文档链接
└── EstimatorCheckFailedWarning(187-215行)
├── 继承 UserWarning
├── docstring # 参数说明
├── init() # 构造参数:estimator、check_name、exception、status、expected_to_fail 等
├── 属性:estimator、check_name、exception、status、expected_to_fail、expected_to_fail_reason
├── repr() # 输出格式包含预期失败标记与异常信息
└── str() # 委托给 repr,保证 print 与交互式显示一致
4.4 异常体系总览 —— 认识这套“智能化的错误反馈系统”
4.4.1 为什么 scikit-learn 需要自定异常和警告?
通用异常(如 ValueError)无法传达机器学习特有的语义信息。当用户在未调用 fit 的情况下直接调用 predict 时,抛出通用的 ValueError("Invalid state") 远不如抛出 NotFittedError("This estimator is not fitted yet") 来得直观。自定义异常帮助用户快速定位是“模型未拟合”还是“数据格式问题”,而分层设计让用户可以精准捕获特定类型的错误,例如仅捕获收敛问题而不拦截数据转换警告。这正如医院分诊系统不能只挂牌“生病”,而要细分“内科”、“外科”、“急诊”,scikit-learn 的异常体系为机器学习工作流提供了专科级的错误分类。
4.4.2 all 列表的设计意图
我们来看 sklearn/exceptions.py 的开头部分,它声明了模块的公共 API 边界。
源码路径:sklearn/exceptions.py - 模块级(第 5-19 行)
# 第 4 章 —— 第 5-19 行
__all__ = [
"ConvergenceWarning", # 收敛警告:优化未收敛
"DataConversionWarning", # 数据转换警告:隐式类型转换
"DataDimensionalityWarning", # 维度警告:目标维度高于原始维度
"EfficiencyWarning", # 效率警告:计算路径非最优
"EstimatorCheckFailedWarning",# 检查失败警告:估计器合规性测试失败
"FitFailedWarning", # 拟合失败警告:交叉验证单折失败
"NotFittedError", # 未拟合异常:模型未训练就预测
"PositiveSpectrumWarning", # 谱警告:PSD 矩阵负特征值
"SkipTestWarning", # 跳过测试警告:可选依赖缺失
"UndefinedMetricWarning", # 未定义指标警告:零除或空预测
"UnsetMetadataPassedError", # 元数据路由异常:传参未被请求
]
这段代码定义了 exceptions 模块的公共导出列表。__all__ 明确声明了哪些类是供外部使用的公共 API,支持 from sklearn.exceptions import * 语法。列表按字母顺序排列,涵盖了从数据预处理到模型持久化的全生命周期。值得注意的是,InconsistentVersionWarning 故意未被包含在内,这暗示它可能被视为内部实现细节,用户需显式导入才能使用,体现了 API 暴露的最小化原则。
4.4.3 异常 vs 警告的分工哲学
异常与警告的核心区别在于程序能否继续执行。异常(Error)表示程序无法继续,必须由用户处理(如模型未拟合就预测);警告(Warning)表示程序可以继续,但需要提示用户注意潜在问题(如隐式数据类型转换)。这种区分让 API 在“严格”与“宽容” 之间取得平衡:严格拦截逻辑错误,宽容提示隐患。这对应医院里“手术坚决不能做错”(异常,必须处理)与“饮食建议仅供参考”(警告,可继续但需知情)的分工。
4.5 NotFittedError —— “未训练就上考场”的守门员(术前检查守门员)
4.5.1 为什么继承自 ValueError 和 AttributeError 两个基类?
NotFittedError 是 scikit-learn 中最常见的异常,它的类定义揭示了一个精妙的兼容性设计。
源码路径:sklearn/exceptions.py - NotFittedError(第 43-62 行)
# 第 4 章 —— 第 43-62 行
class NotFittedError(ValueError, AttributeError): # 双重继承:兼容两种捕获模式
"""Exception class to raise if estimator is used before fitting.
This class inherits from both ValueError and AttributeError to help with
exception handling and backward compatibility.
Examples
--------
>>> from sklearn.svm import LinearSVC
>>> from sklearn.exceptions import NotFittedError
>>> try:
... LinearSVC().predict([[1, 2], [2, 3], [3, 4]])
... except NotFittedError as e:
... print(repr(e))
NotFittedError("This LinearSVC instance is not fitted yet. Call 'fit' with
appropriate arguments before using this estimator."...)
.. versionchanged:: 0.18
Moved from sklearn.utils.validation.
"""
这段代码定义了 NotFittedError 类。它同时继承 ValueError 和 AttributeError,这是一个双重继承的经典案例。语义上,使用未拟合的估计器属于“错误的对象状态”,归 ValueError 管;但历史上,旧版本代码在未拟合时访问内部属性(如 coef_)会直接抛出 AttributeError。双重继承保证了 except ValueError 和 except AttributeError 两种捕获方式都能成功,完美实现了向后兼容。版本迁移记录显示它在 0.18 版本从 sklearn.utils.validation 迁移至此,体现了集中管理异常的架构演进。这正如医院要求“术前检查单”既属于“病历资料不全”(ValueError),又属于“必需检查项缺失”(AttributeError),两种分诊路径都能拦住未检查的病人。
4.6 UnsetMetadataPassedError —— 元数据路由的“精确报错机制”(用药单核对员)
4.6.1 UnsetMetadataPassedError 要解决什么问题?
随着 scikit-learn 1.3 引入元数据路由机制,用户可以精细控制 sample_weight 等元数据如何传递给子估计器。但如果用户传递了某个元数据参数,而估计器并未声明请求它,这往往意味着用户误解了 API 或拼写错误。此时必须明确报错,避免“静默忽略”导致难以调试的逻辑 bug。这正如药房核对员发现处方上写了“阿司匹林”但病人并未开具该药方,必须当面核对而非默默丢弃。
4.6.2 类定义与继承关系
源码路径:sklearn/exceptions.py - UnsetMetadataPassedError(第 22-40 行)
# 第 4 章 —— 第 22-40 行
class UnsetMetadataPassedError(ValueError): # 继承 ValueError:语义上属于输入值不合法
"""Exception class to raise if a metadata is passed which is not explicitly \
requested (metadata=True) or not requested (metadata=False).
.. versionadded:: 1.3
Parameters
----------
message : str
The message
unrequested_params : dict
A dictionary of parameters and their values which are provided but not
requested.
routed_params : dict
A dictionary of routed parameters.
"""
def __init__(self, *, message, unrequested_params, routed_params): # 仅关键字参数
super().__init__(message)
self.unrequested_params = unrequested_params # 记录“多传了什么”
self.routed_params = routed_params # 记录“实际接收了什么”
这段代码定义了 UnsetMetadataPassedError 类。它继承自 ValueError,因为这在语义上属于“输入值不合法”,保证现有 except ValueError 的代码不会漏掉这类错误。类文档字符串标注 versionadded:: 1.3,说明随元数据路由机制一同引入。构造函数使用了仅关键字参数(* 后的参数),强制调用者显式命名参数,如 UnsetMetadataPassedError(message="...", unrequested_params={...}, routed_params={...}),这消除了参数顺序错误的可能性。它接收三个参数:message 是人类可读的错误描述;unrequested_params 记录“多传了什么”参数及其值,直接指向用户的误操作;routed_params 记录“实际上被接收了什么”,供用户对比核对。两个属性作为实例变量保存,方便异常捕获后的程序化分析。
4.7 数据相关警告家族 —— 数据异常的“温柔提醒”(药房剂型变更通知、维度预警雷达)
4.7.1 DataConversionWarning:隐式转换的“透明化”
数据预处理阶段常发生隐式类型转换,scikit-learn 选择警告而非报错,让用户知情但不阻断流程。
源码路径:sklearn/exceptions.py - DataConversionWarning(第 70-91 行)
# 第 4 章 —— 第 70-91 行
class DataConversionWarning(UserWarning): # 继承 UserWarning:面向最终用户
"""Warning used to notify implicit data conversions happening in the code.
This warning occurs when some input data needs to be converted or
interpreted in a way that may not match the user's expectations.
For example, this warning may occur when the user
- passes an integer array to a function which expects float input and
will convert the input
- requests a non-copying operation, but a copy is required to meet the
implementation's data-type expectations;
- passes an input whose shape can be interpreted ambiguously.
.. versionchanged:: 0.18
Moved from sklearn.utils.validation.
"""
这段代码定义了 DataConversionWarning 类。文档字符串详细列举了三大典型触发场景:整数数组被隐式转为浮点、请求无拷贝但被迫拷贝、输入形状有歧义。这些场景的共同点是:操作合法且能继续,但结果可能不符合用户预期。继承 UserWarning 确保默认警告过滤器会显示它,且语义上明确面向最终用户而非开发者。这正如药房将药片碾碎成粉末给病人,虽然药效成分没变,但必须告知“剂型变了,吸收速度可能不同”。
4.7.2 DataDimensionalityWarning:维度问题的“预警雷达”
源码路径:sklearn/exceptions.py - DataDimensionalityWarning(第 94-107 行)
# 第 4 章 —— 第 94-107 行
class DataDimensionalityWarning(UserWarning): # 继承 UserWarning:用户可见
"""Custom warning to notify potential issues with data dimensionality.
For example, in random projection, this warning is raised when the
number of components, which quantifies the dimensionality of the target
projection space, is higher than the number of features, which quantifies
the dimensionality of the original source space, to imply that the
dimensionality of the problem will not be reduced.
.. versionchanged:: 0.18
Moved from sklearn.utils.
"""
这段代码定义了 DataDimensionalityWarning。典型场景是随机投影中用户设置 n_components > n_features,本意是“降维”却变成了“升维”。警告提示用户:操作合法,但语义可能违背初衷。同样继承 UserWarning,保证用户可见。这正如体检时医生发现“你想减重但食谱热量反而超标”,提醒方向可能反了。
4.8 收敛与效率警告 —— 算法运行状态的“仪表盘”(康复进度提醒、节能建议)
4.8.1 ConvergenceWarning:“算法还没跑够”的提示
迭代优化算法(如逻辑回归、SGD)可能在达到最大迭代次数前未收敛。
源码路径:sklearn/exceptions.py - ConvergenceWarning(第 65-68 行)
# 第 4 章 —— 第 65-68 行
class ConvergenceWarning(UserWarning): # 继承 UserWarning:用户可感知
"""Custom warning to capture convergence problems
.. versionchanged:: 0.18
Moved from sklearn.utils.
"""
这段代码定义了 ConvergenceWarning。它不是错误——模型已生成,系数已学到,只是可能没到最优解。警告告诉用户:“可以调大 max_iter 试试”或“检查数据缩放”。这正如康复医生说“训练 100 次还没达标,建议继续练”,病人可出院但知晓未完全康复。
4.8.2 EfficiencyWarning:计算效率的“节能提示”
源码路径:sklearn/exceptions.py - EfficiencyWarning(第 110-121 行)
# 第 4 章 —— 第 110-121 行
class EfficiencyWarning(UserWarning): # 继承 UserWarning:面向用户
"""Warning used to notify the user of inefficient computation.
This warning notifies the user that the efficiency may not be optimal due
to some reason which may be included as a part of the warning message.
This may be subclassed into a more specific Warning class.
.. versionadded:: 0.18
"""
这段代码定义了 EfficiencyWarning。设计为可被子类化,未来可细分出更具体的效率警告(如稀疏矩阵用稠密算法、重复计算等),体现了开放-封闭原则:对扩展开放,对修改封闭。这正如医院贴出“请勿重复挂号检查”的节能提示,后续可细化为“CT 重复拍摄”、“血液重复抽取”等具体子项。
4.8.3 FitFailedWarning:交叉验证中的“单折失败告警”(会诊单科室失败通知)
源码路径:sklearn/exceptions.py - FitFailedWarning(第 124-134 行)
# 第 4 章 —— 第 124-134 行
class FitFailedWarning(RuntimeWarning): # 继承 RuntimeWarning:运行时环境问题
"""Warning class used if there is an error while fitting the estimator.
This Warning is used in meta estimators GridSearchCV and RandomizedSearchCV
and the cross-validation helper function cross_val_score to warn when there
is an error while fitting the estimator.
.. versionchanged:: 0.18
Moved from sklearn.cross_validation.
"""
这段代码定义了 FitFailedWarning。关键区别:继承 RuntimeWarning 而非 UserWarning。RuntimeWarning 语义是“运行时环境问题”,暗示这是外部因素(如数据异常、数值溢出)导致的失败,而非用户代码错误。在 GridSearchCV 等元估计器中,单折失败不应中断整个搜索,但必须告知用户,配合 error_score 参数决定失败后的处理策略。这正如多科室会诊时某一科室设备故障,总台通报“神经内科诊断失败,其他科室结论有效”,不耽误整体会诊。
4.9 版本一致性警告 —— 模型持久化的“安全锁”(旧版病历风险告知)
4.9.1 InconsistentVersionWarning 的核心价值
模型持久化(pickle)是机器学习工程化的关键环节,但跨版本加载模型充满风险。
源码路径:sklearn/exceptions.py - InconsistentVersionWarning(第 163-184 行)
# 第 4 章 —— 第 163-184 行
class InconsistentVersionWarning(UserWarning): # 继承 UserWarning:用户决策
"""Warning raised when an estimator is unpickled with an inconsistent version.
Parameters
----------
estimator_name : str
Estimator name.
current_sklearn_version : str
Current scikit-learn version.
original_sklearn_version : str
Original scikit-learn version.
"""
def __init__(
self, *, estimator_name, current_sklearn_version, original_sklearn_version
): # 仅关键字参数:强制命名,避免顺序错
self.estimator_name = estimator_name
self.current_sklearn_version = current_sklearn_version
self.original_sklearn_version = original_sklearn_version
def __str__(self):
return (
f"Trying to unpickle estimator {self.estimator_name} from version"
f" {self.original_sklearn_version} when "
f"using version {self.current_sklearn_version}. This might lead to breaking"
" code or "
"invalid results. Use at your own risk. "
"For more info please refer to:\n"
"https://scikit-learn.org/stable/model_persistence.html"
"#security-maintainability-limitations"
)
这段代码实现了 InconsistentVersionWarning 的初始化。同样采用仅关键字参数,接收估计器名称、当前版本、原始版本三个关键信息,作为属性保存,方便程序化访问。
4.9.2 str 中的信息设计
源码路径:sklearn/exceptions.py - InconsistentVersionWarning.__str__()(第 178-184 行)
# 第 4 章 —— 第 178-184 行
def __str__(self):
return (
f"Trying to unpickle estimator {self.estimator_name} from version"
f" {self.original_sklearn_version} when "
f"using version {self.current_sklearn_version}. This might lead to breaking"
" code or "
"invalid results. Use at your own risk. "
"For more info please refer to:\n"
"https://scikit-learn.org/stable/model_persistence.html"
"#security-maintainability-limitations"
)
这段代码实现了 __str__ 方法,定义了 print(warning) 或交互式显示时的输出。它清晰地告知用户三个关键信息:估计器名、原始版本、当前版本;明确指出风险“可能导致代码崩溃或结果无效”;使用“Use at your own risk”这一法律术语,划清责任边界;并提供官方文档链接,引导用户深入了解模型持久化的安全限制。这体现了以用户为中心的错误信息设计:不仅说“错了”,还要说“哪里错、为何错、怎么办、去哪查”。这正如护士拿着旧版病历对医生说:“这病历是 5 年前版本的,现在诊疗规范变了,按旧版治风险自负,详见医院新版指南第 3 页”。
4.10 测试失败诊断警告 —— 估计器合规性检查的“体检报告”(详细体检单)
4.10.1 EstimatorCheckFailedWarning 的设计背景
scikit-learn 拥有数百个通用检查,验证估计器的 API 一致性、参数校验、标签正确性等。测试框架需要记录每个失败检查的完整上下文,以便开发者快速定位问题。
4.10.2 完整类定义与构造参数的结构化信息
源码路径:sklearn/exceptions.py - EstimatorCheckFailedWarning(第 187-215 行)
# 第 4 章 —— 第 187-215 行
class EstimatorCheckFailedWarning(UserWarning): # 继承 UserWarning:开发者可见
"""Warning raised when an estimator check from the common tests fails.
Parameters
----------
estimator : estimator object
Estimator instance for which the test failed.
check_name : str
Name of the check that failed.
exception : Exception
Exception raised by the failed check.
status : str
Status of the check.
expected_to_fail : bool
Whether the check was expected to fail.
expected_to_fail_reason : str
Reason for the expected failure.
"""
def __init__(
self,
*,
estimator,
check_name: str,
exception: Exception,
status: str,
expected_to_fail: bool,
expected_to_fail_reason: str,
): # 六个仅关键字参数,承载结构化信息
self.estimator = estimator
self.check_name = check_name
self.exception = exception
self.status = status
self.expected_to_fail = expected_to_fail
self.expected_to_fail_reason = expected_to_fail_reason
def __repr__(self):
expected_to_fail_str = (
f"Expected to fail: {self.expected_to_fail_reason}"
if self.expected_to_fail
else "Not expected to fail"
)
return (
f"Test {self.check_name} failed for estimator {self.estimator!r}.\n"
f"Expected to fail reason: {expected_to_fail_str}\n"
f"Exception: {self.exception}"
)
def __str__(self):
return self.__repr__()
这段代码实现了 EstimatorCheckFailedWarning 的完整定义。六个仅关键字参数承载结构化信息:estimator 是失败的具体估计器实例(方便复现);check_name 定位是哪一个通用检查失败;exception 保留底层异常对象及完整堆栈;status 记录检查状态;expected_to_fail 和 expected_to_fail_reason 区分“预期失败”(已知问题,标记 xfail)与“真实失败”(回归 bug)。这种结构化设计让自动化测试报告生成成为可能。
4.10.3 repr 与 str 的统一
源码路径:sklearn/exceptions.py - EstimatorCheckFailedWarning.__repr__()(第 212-220 行)和 __str__()(第 222-223 行)
# 第 4 章 —— 第 212-220 行
def __repr__(self):
expected_to_fail_str = (
f"Expected to fail: {self.expected_to_fail_reason}"
if self.expected_to_fail
else "Not expected to fail"
)
return (
f"Test {self.check_name} failed for estimator {self.estimator!r}.\n"
f"Expected to fail reason: {expected_to_fail_str}\n"
f"Exception: {self.exception}"
)
# 第 4 章 —— 第 222-223 行
def __str__(self):
return self.__repr__()
这段代码实现了 __repr__ 和 __str__。__repr__ 构建了结构化的多行输出:第一行指出哪个测试对哪个估计器失败;第二行通过三元表达式显示“预期失败原因”或“非预期失败”;第三行附带原始异常信息。__str__ 直接委托给 __repr__,保证 print(warning) 和交互式 REPL 显示完全一致。这种设计让开发者在测试日志中一眼就能看出:这是已知问题还是新引入的回归。这正如体检报告既要给医生看(__repr__ 详细),也要给病人看(__str__ 易懂),且内容完全一致。
4.11 特殊用途警告 —— 测试跳过与内存效率的“信号灯”(可选检查跳过单、指标无效防护、数值哨兵)
4.11.1 SkipTestWarning:测试基础设施的“跳过通知”(可选检查跳过单)
源码路径:sklearn/exceptions.py - SkipTestWarning(第 137-143 行)
# 第 4 章 —— 第 137-143 行
class SkipTestWarning(UserWarning): # 继承 UserWarning:测试框架内部信号
"""Warning class used to notify the user of a test that was skipped.
For example, one of the estimator checks requires a pandas import.
If the pandas package cannot be imported, the test will be skipped rather
than register as a failure.
"""
这段代码定义了 SkipTestWarning。它用于 scikit-learn 自身的测试框架(estimator checks),当可选依赖缺失时发出。区别于 pytest.skip:这是 scikit-learn 测试基础设施内部的信号机制,用于统计“跳过的检查项”,而非直接跳过 pytest 测试用例。这正如体检中心“因设备维护跳过 MRI 检查”记录在案,不算不合格,但报告上会注明。
4.11.2 UndefinedMetricWarning:指标计算的“无效防护”(指标无效防护)
源码路径:sklearn/exceptions.py - UndefinedMetricWarning(第 146-149 行)
# 第 4 章 —— 第 146-149 行
class UndefinedMetricWarning(UserWarning): # 继承 UserWarning:用户可见
"""Warning used when the metric is invalid
.. versionchanged:: 0.18
Moved from sklearn.base.
"""
这段代码定义了 UndefinedMetricWarning。典型场景:precision/recall 在零除或空预测时未定义。警告后返回 0.0 而非 NaN,让调用方明确知道“这个指标不可信,但不会炸掉下游计算”。这正如实验室报告“样本量不足无法计算浓度,报告 0.0 供流程通过,但标注不可信”。
4.11.3 PositiveSpectrumWarning:数值稳健性的“哨兵”(数值稳健性哨兵)
源码路径:sklearn/exceptions.py - PositiveSpectrumWarning(第 152-160 行)
# 第 4 章 —— 第 152-160 行
class PositiveSpectrumWarning(UserWarning): # 继承 UserWarning:用户决策
"""Warning raised when the eigenvalues of a PSD matrix have issues
This warning is typically raised by ``_check_psd_eigenvalues`` when the
eigenvalues of a positive semidefinite (PSD) matrix such as a gram matrix
(kernel) present significant negative eigenvalues, or bad conditioning i.e.
very small non-zero eigenvalues compared to the largest eigenvalue.
.. versionadded:: 0.22
"""
这段代码定义了 PositiveSpectrumWarning。PSD(半正定)矩阵理论上特征值应全非负,但浮点误差可能导致微小负值。当出现“显著负特征值”或“病态条件数”时触发,提示用户:数值结果可能不稳定,建议检查数据条件数或正则化参数。这正如心电图监测发现“ST 段异常”,提示心脏供血可能有问题,建议进一步检查。
4.12 设计中的取舍
对于 FitFailedWarning 选择继承 RuntimeWarning 而非 UserWarning,是因为它发生在交叉验证的单折拟合中,语义上属于运行时环境导致的问题(如奇异矩阵、内存不足、数据异常),而非用户调用 API 的方式错误。RuntimeWarning 面向运行时环境异常,UserWarning 面向用户操作不当,这种区分让上层应用可以选择性捕获:例如监控系统可能只关注 RuntimeWarning 以触发告警,而忽略 UserWarning。
ConvergenceWarning 和 EfficiencyWarning 虽都继承 UserWarning,但触发条件本质不同:ConvergenceWarning 由优化算法内部状态触发(迭代次数耗尽未收敛),关注“结果质量”;EfficiencyWarning 由计算路径选择触发(如用稠密算法处理稀疏数据、重复计算可缓存的中间结果),关注“资源利用率”。
需要独立的 SkipTestWarning 而不直接用 pytest.skip,是因为 SkipTestWarning 是 scikit-learn estimator checks 框架内部的信号,用于统计“哪些通用检查被跳过”,生成合规性报告;pytest.skip 直接跳过测试用例,不留痕迹。两者分工:pytest 管理测试执行流程,SkipTestWarning 管理检查项元数据。
4.13 动手练习
-
阅读异常类定义并分析继承结构
-
阅读
sklearn/exceptions.py第 43-62 行,理解 NotFittedError 的双重继承设计: -
- 为什么同时继承 ValueError 和 AttributeError?
-
- 在 Python 解释器中模拟:分别用
except ValueError和except AttributeError捕获 NotFittedError,验证两种方式都能成功。
- 在 Python 解释器中模拟:分别用
-
- 思考:如果只继承 ValueError,会破坏哪些旧代码的兼容性?
-
回答问题:
-
- all 列表中包含哪些异常/警告?未包含的 InconsistentVersionWarning 为什么被排除?
-
- 异常与警告在继承层次上的最大区别是什么?
-
-
编写自定义异常并集成到 scikit-learn 风格体系
-
假设你要为 scikit-learn 添加一个自定义警告:
MulticollinearityWarning,用于提示数据矩阵存在严重的多重共线性。 -
- 在
sklearn/exceptions.py中添加该类的完整定义,继承自 UserWarning,包含清晰的 docstring。
- 在
-
- 将其添加到 all 列表中。
-
- 编写一段测试代码:导入该警告并在数据条件数超过阈值时触发
warnings.warn(..., MulticollinearityWarning)。
- 编写一段测试代码:导入该警告并在数据条件数超过阈值时触发
-
注意:
-
- 保持与现有警告风格一致(versionadded 标注、简洁 docstring)
-
- 不要包含实际答案,只描述实现步骤
-
-
分析 EstimatorCheckFailedWarning 的信息结构
-
阅读
sklearn/exceptions.py第 187-215 行,理解 EstimatorCheckFailedWarning 的设计: -
- 构造函数的 6 个 keyword-only 参数分别承载什么信息?
-
- 为什么
__str__直接委托给__repr__?这样做的优缺点是什么?
- 为什么
-
- repr 的输出格式中,
expected_to_fail_str的三元表达式逻辑是什么?
- repr 的输出格式中,
-
回答问题:
-
- 如果 check 失败但 expected_to_fail=False,输出会如何呈现?
-
- 如果 expected_to_fail=True 且 expected_to_fail_reason="known issue #123",输出中的原因部分是什么?
-
-
比较不同警告的基类选择与使用场景
-
阅读
sklearn/exceptions.py全文,完成以下对比分析: -
- FitFailedWarning 为什么继承 RuntimeWarning 而非 UserWarning?
-
- ConvergenceWarning 和 EfficiencyWarning 都继承 UserWarning,如何区分它们的触发条件?
-
- SkipTestWarning 与 pytest.skip 的关系是什么?为什么需要独立的警告类?
-
回答问题:
-
- 在 GridSearchCV 中,某折拟合失败时会触发哪个警告?
-
- 如果用户在随机投影中设置 n_components 大于 n_features,会收到哪个警告?
-
- 如果 pickle 加载旧版模型,会收到哪个警告?
-
-
实战:为元数据路由添加更精确的错误信息
-
阅读
sklearn/exceptions.py第 22-40 行 UnsetMetadataPassedError 的实现: -
- 构造函数使用了仅关键字参数(
*),目的是什么?
- 构造函数使用了仅关键字参数(
-
self.unrequested_params和self.routed_params分别存储什么内容?
-
- 在 scikit-learn 源码中搜索
UnsetMetadataPassedError的使用位置(可使用 grep),理解它在元数据路由验证中的触发时机。
- 在 scikit-learn 源码中搜索
-
回答问题:
-
- 如果用户传递了 sample_weight 但估计器未声明请求,会触发什么异常?
-
- 异常消息中会包含哪些信息帮助用户调试?
-
4.14 本章小结
这一章中我们学习/了解/讨论了 scikit-learn 完整的异常与警告体系。首先,我们理解了异常与警告的分工哲学:异常是硬性拦截,警告是软性提醒。其次,我们深入剖析了 NotFittedError 的双重继承设计,理解其如何同时兼容 ValueError 和 AttributeError 捕获模式。接着,我们探讨了 UnsetMetadataPassedError 在元数据路由中的精确报错机制,特别是仅关键字参数和结构化错误信息的设计。然后,我们系统梳理了数据转换、维度、收敛、效率、拟合失败等警告家族的适用场景与基类选择考量。最后,我们分析了版本一致性警告的风险沟通设计,以及估计器合规检查失败警告的结构化体检报告格式。
本章我们一起学习了以下概念:
下表汇总了本章涉及的 12 个核心异常与警告类及其核心职责:
| 概念 | 解释 |
|------|------|
| NotFittedError | 双重继承 ValueError 和 AttributeError,预测前未拟合的统一守门员 |
| UnsetMetadataPassedError | 元数据路由中传了未请求的参数时精确报错,避免静默忽略 |
| DataConversionWarning | 隐式数据转换(整数→浮点、强制拷贝、形状歧义)的温柔提醒 |
| DataDimensionalityWarning | 维度问题预警,如随机投影中目标维度高于原始维度 |
| ConvergenceWarning | 优化算法达到 max_iter 未收敛,提示可调大迭代次数 |
| EfficiencyWarning | 计算效率未达最优的节能提示,可被子类化扩展 |
| FitFailedWarning | 交叉验证中单折拟合失败,继承 RuntimeWarning,配合 error_score 使用 |
| InconsistentVersionWarning | 跨版本 pickle 加载模型的风险安全锁,未在 all 中导出 |
| EstimatorCheckFailedWarning | 估计器合规性检查失败的结构化体检报告,str 委托给 repr |
| SkipTestWarning | 可选依赖缺失时测试跳过而非失败的通知信号 |
| UndefinedMetricWarning | 指标在零除或空预测时未定义,返回 0.0 而非 NaN |
| PositiveSpectrumWarning | PSD 矩阵出现显著负特征值的数值稳健性哨兵 |
下一章中,我们将学习 scikit-learn 的测试基础设施,了解全局随机种子控制、数据集获取策略、平台条件跳过和线程安全机制如何共同保证测试的可靠性与可重复性。
第 5 章 —— 测试基础设施 —— 搭建“质量保障的防线”
5.1 学习目标
-
难度:★★★☆☆(3/5)
-
预备知识:Python 基础、面向对象编程与 Markdown/代码阅读基础
-
理解pytest钩子函数(pytest_configure、pytest_generate_tests、pytest_collection_modifyitems、pytest_addoption)在测试会话生命周期中的执行时机与职责边界
-
掌握通过环境变量SKLEARN_TESTS_GLOBAL_RANDOM_SEED实现确定性随机测试的机制
-
理解_fetch_fixture装饰器如何实现网络受限环境下的数据集按需加载与离线跳过
-
掌握pytest-xdist并行测试中线程限制与数据下载的安全策略
-
理解pyplot fixture如何管理matplotlib图形生命周期
-
掌握hide_available_pandas通过monkeypatch模拟缺失依赖的技术
-
理解raccoon_face_or_skip对网络和pooch依赖的多重检测逻辑
-
掌握dt_config严格模式配置与scipy_doctest集成机制
-
能阅读并修改conftest.py中的全局测试配置逻辑
5.2 生活类比
想象scikit-learn的测试基础设施是一家精密仪器的质量检测实验室:pytest_configure 就是实验室开门的准备工作(切换无界面模式、限制并行设备数量、设置警报阈值、张贴实验守则),pytest_generate_tests 则是实验方案的参数化设计(同一实验在不同随机条件下重复验证),SKLEARN_TESTS_GLOBAL_RANDOM_SEED 环境变量就像随机数生成器的“种子库”(预设100种种子供选择),而 _fetch_fixture 装饰器则扮演仓库管理员的角色(需要时去取数据集,网络断了就跳过实验)。当使用 pytest-xdist 进行并行测试时,多个实验台需要协调工作,这时 pytest_collection_modifyitems 就成了质检调度中心——它负责统一预下载物料、标记不合格项、隔离实验环境,避免多个工作台同时操作同一资源导致冲突。pyplot fixture 确保绘图设备在每次实验前后都被正确关闭,就像实验室的“开关管理员”防止资源泄漏;hide_available_pandas 则允许我们模拟“缺少某种试剂”的场景,验证实验在受限环境下的表现;raccoon_face_or_skip 函数就像特殊标本管理员,会先检查网络连接和专用工具(pooch)是否就绪,只有两者都可用时才允许取用特殊数据集;最后,dt_config 的严格模式就像精密仪器的校准标准,能够区分 3.14 和 np.float64(3.14) 这样细微但关键的差异。模块级代码在导入时就像实验室主管在实验开始前统一配置设备、准备物料、标记已知问题,它会检测并行插件、校验pytest版本、注册所有数据集fixture,并启用dt_config的严格检查模式,确保整个测试套件在启动时就处于受控状态。
这一类比将贯穿全章:我们将把每个钩子、fixture 和模块级代码都映射到实验室的具体角色和流程中,帮助你建立直观的心智模型。
5.3 源码地图
5.4 全局测试配置与随机性控制 —— 启动确定性测试的“总闸门”
在机器学习测试中,随机性无处不在:数据采样、模型权重初始化、特征打乱顺序都依赖随机数生成器。如果每个测试文件都使用独立的随机种子,当我们遇到偶发失败时,就很难精准复现问题发生的条件。通过统一的环境变量 SKLEARN_TESTS_GLOBAL_RANDOM_SEED 控制全局随机种子,所有使用 global_random_seed fixture 的测试都将进入确定性框架,这为调试提供了可靠的基础。
pytest_configure 钩子正是这个框架的“总闸门”,它在测试会话启动时统一完成四项关键配置:设置 matplotlib 后端为无界面模式、根据并行工作进程数动态限制线程数量、启用警告升级为错误的机制,以及在必要时注册自定义测试标记。这些看似琐碎的操作实则为后续测试的可靠性奠定了基础——特别是在 CI 环境中,没有显示器的服务器上强制使用 'agg' 后端可以避免图形初始化失败,而根据物理核心数和 xdist worker 数量动态调整 OpenMP/BLAS 线程数则能有效防止资源过载。
当测试会话启动时,pytest_configure 会首先尝试将 matplotlib 后端设置为 'agg'(无界面模式),这在无显示器的服务器环境中尤为重要。随后,它根据可用的物理核心数和可能的 xdist worker 数量计算安全的线程限制值,最后通过 threadpool_limits 将这个值应用到 OpenMP 和 BLAS 库。如果环境变量 SKLEARN_WARNINGS_AS_ERRORS 被设置,则会动态修改 pytest 的警告过滤规则,将特定警告升级为错误;当 pytest_run_parallel 插件不可用时,还会注册三个自定义 marker(parallel_threads、thread_unsafe、iterations)以避免未知标记警告。
# 第 5 章 —— sklearn/conftest.py (215-247)
def pytest_configure(config):
# Use matplotlib agg backend during the tests including doctests
try:
import matplotlib
matplotlib.use("agg")
except ImportError:
pass
allowed_parallelism = joblib.cpu_count(only_physical_cores=True)
xdist_worker_count = environ.get("PYTEST_XDIST_WORKER_COUNT")
if xdist_worker_count is not None:
# Set the number of OpenMP and BLAS threads based on the number of workers
# xdist is using to prevent oversubscription.
allowed_parallelism = max(allowed_parallelism // int(xdist_worker_count), 1)
threadpool_limits(allowed_parallelism)
if environ.get("SKLEARN_WARNINGS_AS_ERRORS", "0") != "0":
# This seems like the only way to programmatically change the config
# filterwarnings. This was suggested in
# https://github.com/pytest-dev/pytest/issues/3311#issuecomment-373177592
for line in get_pytest_filterwarning_lines():
config.addinivalue_line("filterwarnings", line)
if not PARALLEL_RUN_AVAILABLE:
config.addinivalue_line(
"markers",
"parallel_threads(n): run the given test function in parallel "
"using `n` threads.",
)
config.addinivalue_line(
"markers",
"thread_unsafe: mark the test function as single-threaded",
)
config.addinivalue_line(
"markers",
"iterations(n): run the given test function `n` times in each thread",
)
config.addinivalue_line(
"markers",
"iterations(n): run the given test function `n` times in each thread",
)
这段代码定义了测试会话的启动配置,确保了 matplotlib 在无头环境下可用、线程资源不会过度订阅、警告可以被当作错误处理,并且在缺少并行插件时仍能通过自定义 marker 实现类似功能。通过在会话开始时统一处理这些全局性问题,我们为后续的测试执行创造了一个可预测、受控的环境。
代码解析:pytest_configure 是测试会话的“总闸门”,它按照优先级依次处理:1)图形后端设置(最先执行,影响后续所有绘图操作);2)线程资源限制(根据物理核心和 worker 数计算,防止 BLAS/OpenMP 过度订阅);3)警告升级为错误(仅当环境变量开启时);4)自定义 marker 注册(仅当并行插件不可用时)。这种分层配置体现了“先全局环境,再资源控制,后可选增强”的设计原则。
5.5 pytest_generate_tests 钩子 —— 随机种子的“参数化流水线”
pytest 为每个测试函数调用 pytest_generate_tests 钩子时,会传入一个 metafunc 对象,这个对象携带了当前测试函数所需的 fixture 信息。钩子的核心职责是检查其中是否需要 global_random_seed fixture,如果需要的话,就根据环境变量 SKLEARN_TESTS_GLOBAL_RANDOM_SEED 动态生成多组种子参数,并通过 metafunc.parametrize 注入到测试函数中。这种机制让我们能够用不同的随机种子重复运行同一测试,从而检验其对种子选择的敏感性——这正是确定性测试的核心价值所在。
环境变量支持三种解析模式:未设置时使用默认种子 [42] 保证日常开发的可复现性;设置为 'all' 时遍历 [0, 99] 的全部 100 个种子进行彻底验证;设置为具体值或区间(如 '10-20')时则精确控制测试范围。区间解析特别值得注意,它使用 range(int(start), int(stop) + 1) 来确保右端点被包含,随后会立即校验所有种子是否在有效范围 [0, 99] 内,越界时会抛出携带原始输入值的 ValueError 以便快速定位问题。由于 pytest-xdist 的 worker 是子进程,它们会自动继承父进程的环境变量,因此无需额外同步机制——每个 worker 都会独立读取同一环境变量并执行参数化,这正是注释中所强调的“依赖环境变量在 worker 子进程中可访问”这一设计前提的体现。
# 第 5 章 —— sklearn/conftest.py (176-213)
def pytest_generate_tests(metafunc):
"""Parametrization of global_random_seed fixture
based on the SKLEARN_TESTS_GLOBAL_RANDOM_SEED environment variable.
The goal of this fixture is to prevent tests that use it to be sensitive
to a specific seed value while still being deterministic by default.
See the documentation for the SKLEARN_TESTS_GLOBAL_RANDOM_SEED
variable for instructions on how to use this fixture.
https://scikit-learn.org/dev/computing/parallelism.html#sklearn-tests-global-random-seed
"""
# When using pytest-xdist this function is called in the xdist workers.
# We rely on SKLEARN_TESTS_GLOBAL_RANDOM_SEED environment variable which is
# set in before running pytest and is available in xdist workers since they
# are subprocesses.
RANDOM_SEED_RANGE = list(range(100)) # All seeds in [0, 99] should be valid.
random_seed_var = environ.get("SKLEARN_TESTS_GLOBAL_RANDOM_SEED")
default_random_seeds = [42]
if random_seed_var is None:
random_seeds = default_random_seeds
elif random_seed_var == "all":
random_seeds = RANDOM_SEED_RANGE
else:
if "-" in random_seed_var:
start, stop = random_seed_var.split("-")
random_seeds = list(range(int(start), int(stop) + 1))
else:
random_seeds = [int(random_seed_var)]
if min(random_seeds) < 0 or max(random_seeds) > 99:
raise ValueError(
"The value(s) of the environment variable "
"SKLEARN_TESTS_GLOBAL_RANDOM_SEED must be in the range [0, 99] "
f"(or 'all'), got: {random_seed_var}"
)
if "global_random_seed" in metafunc.fixturenames:
metafunc.parametrize("global_random_seed", random_seeds)
这段代码实现了基于环境变量的随机种子参数化,确保了测试既能默认使用固定种子获得可复现性,又能在需要时通过环境变量灵活控制测试范围。通过这种机制,我们既保证了日常调试的效率,又能在需要彻底验证时启用全种子遍历,真正做到了“既确定又灵活”。
代码解析:pytest_generate_tests 是随机种子的“参数化流水线”,它采用“按需参数化”策略——仅当测试函数声明需要 global_random_seed 时才生成参数。三种模式的解析逻辑清晰分离:默认模式保证开发效率,全种子模式支持 CI 彻底验证,区间模式允许精准复现特定失败区间。区间解析时的 +1 操作体现了对闭区间语义的坚持,而范围校验的即时失败则体现了“快速失败、精准定位”的工程哲学。
5.6 数据集 Fixture 工厂 —— _fetch_fixture 装饰器的“按需供应站”
获取真实世界数据集往往需要网络访问,这在受限环境或离线测试中会成为主要障碍。为了解决这个问题,scikit-learn 引入了环境变量 SKLEARN_SKIP_NETWORK_TESTS 的机制:当其值为默认的 '1' 时,所有需要网络的数据集测试将被自动跳过;只有当显式设置为 '0' 时,才会允许按需从远程仓库下载数据集。值得注意的是,这个环境变量是在模块导入时就被读取并固定的,而不是在每次 fixture 调用时动态检测,这避免了在测试执行过程中产生不一致的行为。
_fetch_fixture 装饰器正是围绕这一机制构建的三层嵌套结构。最外层函数 _fetch_fixture(f) 接收一个原始的数据集获取函数(如 fetch_20newsgroups),并返回一个包装过的 pytest fixture;中间层的 wrapped 函数负责将环境变量转换为 download_if_missing 参数注入到原始函数调用中,并精心处理可能的 OSError 异常;最内层则使用 pytest.fixture(lambda: wrapped) 将这个包装函数注册为真正的 pytest fixture。这种设计巧妙地将环境感知逻辑与 fixture 注册解耦,同时保持了对原始函数签名的透明支持。
当网络不可用且 download_if_missing 为 False 时,原始的 fetch 函数会抛出一个特定的 OSError,其消息恰好是 "Data not found and download_if_missing is False"。wrapped 函数会捕获这个异常,仅当错误消息精确匹配时才调用 pytest.skip 跳过测试;其他类型的文件系统错误(如磁盘损坏或权限问题)则会原样抛出,以避免掩盖真正的问题。为了保持代码的可调试性和文档友好性,装饰器还使用 @wraps(f) 保留了原始函数的 name 和 doc。值得注意的是,虽然大多数数据集 fixture 都是通过这个装饰器批量生成的,但 raccoon_face_fxt 是个例外——它直接使用 pytest.fixture(raccoon_face_or_skip) 包装,因为 raccoon_face_or_skip 已经是一个完整的、无需额外参数注入的函数。模块级代码通过一系列赋值语句(如 fetch_20newsgroups_fxt = _fetch_fixture(fetch_20newsgroups))完成了所有数据集 fixture 的注册,并同时维护了一个 dataset_fetchers 字典,该字典在后续的收集阶段调度中会被用到。
# 第 5 章 —— sklearn/conftest.py (92-108)
def _fetch_fixture(f):
"""Fetch dataset (download if missing and requested by environment)."""
download_if_missing = environ.get("SKLEARN_SKIP_NETWORK_TESTS", "1") == "0"
@wraps(f)
def wrapped(*args, **kwargs):
kwargs["download_if_missing"] = download_if_missing
try:
return f(*args, **kwargs)
except OSError as e:
if str(e) != "Data not found and `download_if_missing` is False":
raise
pytest.skip("test is enabled when SKLEARN_SKIP_NETWORK_TESTS=0")
return pytest.fixture(lambda: wrapped)
# 第 5 章 —— sklearn/conftest.py (51-62)
def raccoon_face_or_skip():
# SciPy requires network access to get data
run_network_tests = environ.get("SKLEARN_SKIP_NETWORK_TESTS", "1") == "0"
if not run_network_tests:
raise SkipTest("test is enabled when SKLEARN_SKIP_NETWORK_TESTS=0")
try:
import pooch # noqa: F401
except ImportError:
raise SkipTest("test requires pooch to be installed")
return face(gray=True)
# 第 5 章 —— sklearn/conftest.py (110-124)
# 第 5 章 —— Adds fixtures for fetching data
fetch_20newsgroups_fxt = _fetch_fixture(fetch_20newsgroups)
fetch_20newsgroups_vectorized_fxt = _fetch_fixture(fetch_20newsgroups_vectorized)
fetch_california_housing_fxt = _fetch_fixture(fetch_california_housing)
fetch_covtype_fxt = _fetch_fixture(fetch_covtype)
fetch_kddcup99_fxt = _fetch_fixture(fetch_kddcup99)
fetch_lfw_pairs_fxt = _fetch_fixture(fetch_lfw_pairs)
fetch_lfw_people_fxt = _fetch_fixture(fetch_lfw_people)
fetch_olivetti_faces_fxt = _fetch_fixture(fetch_olivetti_faces)
fetch_rcv1_fxt = _fetch_fixture(fetch_rcv1)
fetch_species_distributions_fxt = _fetch_fixture(fetch_species_distributions)
raccoon_face_fxt = pytest.fixture(raccoon_face_or_skip)
这三个代码块共同实现了数据集获取的完整流程:_fetch_fixture 装饰器负责将环境变量转换为下载行为并处理离线场景;raccoon_face_or_skip 函数提供了对特殊数据集(带网络和 pooch 依赖的灰度浣熊脸图像)的访问控制;而模块级代码则通过显式赋值将所有这些包装函数注册为 pytest 可用的 fixture。这种设计不仅确保了网络测试的安全可控,还使得离线测试变得简单直接——只需设置环境变量 SKLEARN_SKIP_NETWORK_TESTS=1,所有需要网络的测试就会被优雅地跳过,而不会因为缺失数据而失败。
代码解析:_fetch_fixture 采用“环境变量固化 + 参数注入 + 精准异常捕获”的三层设计。环境变量在模块导入时读取,保证了整个测试会话的一致性;download_if_missing 作为关键字参数注入,保持了对原始函数签名的零侵入;异常捕获仅匹配特定错误消息,体现了“只处理预期的离线情况,不掩盖真实文件系统错误”的防御性编程原则。raccoon_face_or_skip 则展示了“双重门禁”模式:先检网络开关,再检 pooch 依赖,两者缺一不可,这正是特殊数据集需要更严格准入条件的体现。
5.7 pytest_collection_modifyitems —— 收集阶段的数据集“预下载调度器”
在使用 pytest-xdist 进行并行测试时,如果允许每个工作进程独立下载数据集,就可能出现多个进程同时写入同一缓存文件的情况,这不仅会导致文件损坏,还可能因为竞态条件引发不可预测的错误。为了避免这种线程/进程不安全的行为,scikit-learn 将数据下载的时机提前到了测试收集阶段完成之后——此时虽然 pytest 已经完成了所有测试项的发现,但尚未开始真正的测试执行,因此可以由单点调度入口安全地完成所有必要的下载工作。
pytest_collection_modifyitems 钩子正是这个调度入口的实现。它会遍历所有已收集的测试项,识别出那些需要数据集 fixture 的测试(无论是通过 fixturenames 属性的普通测试,还是需要从 item.name 中解析函数名的 DoctestItem),并根据当前的网络开关状态(由环境变量 SKLEARN_SKIP_NETWORK_TESTS 决定)决定是将这些 fixture 加入待下载集合,还是直接为对应的测试项添加跳过标记。只有当网络测试被允许且我们在主进程或第一个 xdist worker(标识为 'gw0')时,才会真正触发数据集的下载过程;其他所有 worker 则会跳过这一步,直接使用已经下载好的缓存文件。
除了数据集下载调度,这个钩子还承担着多项环境适配工作:它会检测 ARM64 平台上的 GradientBoostingClassifier 已知失败并添加 xfail 标记(保留回归敏感性而非直接跳过)、根据 matplotlib、平台位数、操作系统以及 NumPy/SciPy 版本决定是否跳过 doctest,并通过将 dtest.globs 设置为空字典来强制 doctest 示例必须自包含所有必要的 import。值得注意的是,对于 contextmanager 类型的 doctest(如 sklearn._config.config_context),需要特别排除在跳过列表之外,这是为了规避 pytest 已知的 Bug #8796;最后,它还会在收集阶段检测 PIL 依赖,并为相关的图像处理测试添加跳过标记——这种在收集阶段而非导入时进行依赖检测的设计,避免了在 conftest 顶部引入额外的导入开销。
# 第 5 章 —— sklearn/conftest.py (130-174)
def pytest_collection_modifyitems(config, items):
"""Called after collect is completed.
Parameters
----------
config : pytest config
items : list of collected items
"""
run_network_tests = environ.get("SKLEARN_SKIP_NETWORK_TESTS", "1") == "0"
skip_network = pytest.mark.skip(
reason="test is enabled when SKLEARN_SKIP_NETWORK_TESTS=0"
)
# download datasets during collection to avoid thread unsafe behavior
# when running pytest in parallel with pytest-xdist
dataset_features_set = set(dataset_fetchers)
datasets_to_download = set()
for item in items:
if isinstance(item, DoctestItem) and "fetch_" in item.name:
fetcher_function_name = item.name.split(".")[-1]
dataset_fetchers_key = f"{fetcher_function_name}_fxt"
dataset_to_fetch = set([dataset_fetchers_key]) & dataset_features_set
elif not hasattr(item, "fixturenames"):
continue
else:
item_fixtures = set(item.fixturenames)
dataset_to_fetch = item_fixtures & dataset_features_set
if not dataset_to_fetch:
continue
if run_network_tests:
datasets_to_download |= dataset_to_fetch
else:
# network tests are skipped
item.add_marker(skip_network)
# Only download datasets on the first worker spawned by pytest-xdist
# to avoid thread unsafe behavior. If pytest-xdist is not used, we still
# download before tests run.
worker_id = environ.get("PYTEST_XDIST_WORKER", "gw0")
if worker_id == "gw0" and run_network_tests:
for name in datasets_to_download:
with suppress(SkipTest):
dataset_fetchers[name]()
for item in items:
# Known failure on with GradientBoostingClassifier on ARM64
if (
item.name.endswith("GradientBoostingClassifier")
and platform.machine() == "aarch64"
):
marker = pytest.mark.xfail(
reason=(
"know failure. See "
"https://github.com/scikit-learn/scikit-learn/issues/17797"
)
)
item.add_marker(marker)
skip_doctests = False
try:
import matplotlib # noqa: F401
except ImportError:
skip_doctests = True
reason = "matplotlib is required to run the doctests"
if _IS_32BIT:
reason = "doctest are only run when the default numpy int is 64 bits."
skip_doctests = True
elif sys.platform.startswith("win32"):
reason = (
"doctests are not run for Windows because numpy arrays "
"repr is inconsistent across platforms."
)
skip_doctests = True
if np_base_version < parse_version("2"):
# TODO: configure numpy to output scalar arrays as regular Python scalars
# once possible to improve readability of the tests docstrings.
# https://numpy.org/neps/nep-0051-scalar-representation.html#implementation
reason = "Due to NEP 51 numpy scalar repr has changed in numpy 2"
skip_doctests = True
if sp_version < parse_version("1.14"):
reason = "Scipy sparse matrix repr has changed in scipy 1.14"
skip_doctests = True
# Normally doctest has the entire module's scope. Here we set globs to an empty dict
# to remove the module's scope:
# https://docs.python.org/3/library/doctest.html#what-s-the-execution-context
for item in items:
if isinstance(item, DoctestItem):
item.dtest.globs = {}
if skip_doctests:
skip_marker = pytest.mark.skip(reason=reason)
for item in items:
if isinstance(item, DoctestItem):
# work-around an internal error with pytest if adding a skip
# mark to a doctest in a contextmanager, see
# https://github.com/pytest-dev/pytest/issues/8796 for more
# details.
if item.name != "sklearn._config.config_context":
item.add_marker(skip_marker)
try:
import PIL # noqa: F401
pillow_installed = True
except ImportError:
pillow_installed = False
if not pillow_installed:
skip_marker = pytest.mark.skip(reason="pillow (or PIL) not installed!")
for item in items:
if item.name in [
"sklearn.feature_extraction.image.PatchExtractor",
"sklearn.feature_extraction.image.extract_patches_2d",
]:
item.add_marker(skip_marker)
这段代码实现了测试收集完成后的统一调度,它不仅解决了并行测试中的数据下载安全问题,还通过多种环境检测确保了测试只在适合的平台上运行。通过在收集阶段而非执行阶段处理数据下载,我们避免了多 worker 同时写入同一文件的风险;通过精心设置的 skip/xfail 标记,我们既保留了对已知问题的敏感性,又避免了不必要的测试执行开销;而将 doctest 的 globs 重置为空字典这一看似微小的操作,却显著提高了文档测试的可移植性和可靠性。
代码解析:pytest_collection_modifyitems 是收集阶段的“总调度中心”,它遵循“先识别、再决策、后执行”的三步走策略。数据集下载调度通过“收集依赖集合 → 仅 gw0 下载 → 其他 worker 复用缓存”实现了并行安全;ARM64 xfail 使用平台检测精准定位已知问题;Doctest 条件跳过采用“任一条件不满足即跳过”的宽松策略;globs 重置为空字典强制示例自包含;PIL 检测延迟到收集阶段避免导入开销。每一个决策点都体现了对测试执行环境的深度理解和对边界情况的精细把控。
5.8 Doctest 条件跳过与作用域隔离 —— 文档测试的“环境安检员”
doctest 之所以能够嵌入在文档字符串中并随代码一起测试,是一项极大的便利,但它也带来了独特的挑战:文档中的示例可能依赖 matplotlib 绘图、假设特定的平台行为或依赖某些版本特有的 NumPy/SciPy 表现。为了确保 doctest 只在真正支持的环境中运行,scikit-learn 在 pytest_collection_modifyitems 中设置了四重条件门禁:matplotlib 必须已安装、平台不能是 32 位(因为 numpy 整数默认位数不同会影响数组表示)、不是 Windows 系统(由于 numpy 数组 repr 的平台差异)、以及 NumPy 和 SciPy 必须满足特定的最低版本要求(因为标量和稀疏矩阵的字符串表示在版本更新后会发生变化)。
当满足跳过条件时,我们不仅会跳过对应的 doctest,还会执行一次关键的作用域隔离操作:将每个 DoctestItem 的 dtest.globs 属性设置为空字典。这看似简单的一行代码实际上具有深远的意义——它强制 doctest 示例必须自包含所有必要的 import 语句,而不能依赖于被测试模块的全局命名空间。这样的设计显著提高了 doctest 的可移植性和可靠性,因为它消除了对外部上下文的隐式依赖,使得每个 doctest 示例都能在独立的环境中被正确理解和执行。
然而,直接对 contextmanager 类型的 doctest 添加跳过标记会触发 pytest 的内部错误(Bug #8796),因此需要特别排除像 sklearn._config.config_context 这样的特殊情况。这种 workaround 虽显得有些笨拙,但正是基于对 pytest 行为的深刻理解而做出的妥协——它让我们在保持大多数 doctest 安全跳过的同时,避免了对特殊构造的误伤。此外,钩子还会在收集阶段检测 PIL 库的可用性,并为依赖它的图像处理函数(如 PatchExtractor)添加跳过标记;选择在收集阶段而非导入时进行这个检测,是为了避免在 conftest 模块顶部引入不必要的导入开销,保持模块的轻量级特征。
# 第 5 章 —— sklearn/conftest.py (176-211)
skip_doctests = False
try:
import matplotlib # noqa: F401
except ImportError:
skip_doctests = True
reason = "matplotlib is required to run the doctests"
if _IS_32BIT:
reason = "doctest are only run when the default numpy int is 64 bits."
skip_doctests = True
elif sys.platform.startswith("win32"):
reason = (
"doctests are not run for Windows because numpy arrays "
"repr is inconsistent across platforms."
)
skip_doctests = True
if np_base_version < parse_version("2"):
# TODO: configure numpy to output scalar arrays as regular Python scalars
# once possible to improve readability of the tests docstrings.
# https://numpy.org/neps/nep-0051-scalar-representation.html#implementation
reason = "Due to NEP 51 numpy scalar repr has changed in numpy 2"
skip_doctests = True
if sp_version < parse_version("1.14"):
reason = "Scipy sparse matrix repr has changed in scipy 1.14"
skip_doctests = True
# Normally doctest has the entire module's scope. Here we set globs to an empty dict
# to remove the module's scope:
# https://docs.python.org/3/library/doctest.html#what-s-the-execution-context
for item in items:
if isinstance(item, DoctestItem):
item.dtest.globs = {}
if skip_doctests:
skip_marker = pytest.mark.skip(reason=reason)
for item in items:
if isinstance(item, DoctestItem):
# work-around an internal error with pytest if adding a skip
# mark to a doctest in a contextmanager, see
# https://github.com/pytest-dev/pytest/issues/8796 for more
# details.
if item.name != "sklearn._config.config_context":
item.add_marker(skip_marker)
try:
import PIL # noqa: F401
pillow_installed = True
except ImportError:
pillow_installed = False
if not pillow_installed:
skip_marker = pytest.mark.skip(reason="pillow (or PIL) not installed!")
for item in items:
if item.name in [
"sklearn.feature_extraction.image.PatchExtractor",
"sklearn.feature_extraction.image.extract_patches_2d",
]:
item.add_marker(skip_marker)
这段代码展示了 doctest 的环境适配与作用域管理逻辑,它不仅根据多种条件决定 doctest 是否应该运行,还通过重置 globs 来确保每个 doctest 示例都是自包含的。这种双重保障机制——先通过环境检测避免在不支持的平台上运行,再通过作用域隔离消除对外部上下文的依赖——正是 scikit-learn 能够在多样化环境中保持高质量文档测试的关键所在。
代码解析:Doctest 的环境安检采用“层层设防、任一不符即停”的策略:matplotlib 缺失、32位平台、Windows 系统、NumPy 版本过低、SciPy 版本过低——任一条件触发即全量跳过。这种看似激进的策略实则是对 doctest 脆弱性的妥善应对:文档示例极其敏感,任何微小的环境差异都可能导致字符串表示不一致从而产生假阳性失败。globs = {} 的作用域隔离则是从根源上消除了隐式依赖,让每个示例都成为可独立运行的最小单元。contextmanager 的特例排除则展示了工程实践中“不得不为”的妥协之美。
5.9 pyplot 与 hide_available_pandas Fixture —— 资源管理与依赖模拟的“双面手”
在使用 matplotlib 进行测试时,如果不妥善管理图形资源,很容易导致内存泄漏或后续测试受到干扰。pyplot fixture 正是为解决这个问题而设计的:它在每个测试函数执行前后都会调用 pyplot.close("all") 来确保没有遗留的图形对象占用资源。值得一提的是,这个 fixture 使用了 yield 而不是 return,这使得它成为一个生成器函数,从而能够自然地支持 teardown 逻辑——在 yield 语句之后的代码会在测试函数执行完成后统一运行,这正是实现“前置准备”和“后置清理”的关键手段。当 matplotlib 不可用时,fixture 会通过 pytest.importorskip 优雅地跳过测试,而不会抛出导致测试中断的 ImportError。
相比之下,hide_available_pandas fixture 则扮演着完全不同的角色——它允许测试者模拟“pandas 未安装”的环境,以验证 scikit-learn 在该依赖缺失时的行为。这个 fixture 的核心是通过 monkeypatch 替换 builtins.import 函数:当尝试导入名字精确等于 'pandas' 的模块时,它会抛出 ImportError;而对于所有其他导入请求(包括 pandas.core.frame 这样的子模块),则会转发到原始的 import 函数。这个精确匹配的设计并非偶然:一旦 pandas 顶层包导入失败,Python 的导入机制确保其子模块根本不会有机会被尝试加载,因此没有必要去匹配更具体的模块名。fixture 返回值为 None 是有意而为之的设计,因为测试函数只需通过在参数列表中声明 fixture 来激活其效果,无需实际使用其返回值。
这两个 fixture 虽然服务于不同的目的,但它们共同体现了 scikit-learn 测试基础设施的设计哲学:通过声明式的机制(简单地声明 fixture 就能获得对应能力)来管理复杂的环境状态,同时将资源生命周期的控制(如图形关闭)和依赖模拟(如 import 拦截)封装在可重用的组件中。当测试函数需要图形资源时,它只需声明 pyplot fixture;当需要验证在无 pandas 环境下的行为时,它只需声明 hide_available_pandas fixture;而在两者之间,monkeypatch fixture 负责在测试结束后自动还原被修改的状态,彻底释放测试作者从手动管理的负担。
# 第 5 章 —— sklearn/conftest.py (216-224)
@pytest.fixture(scope="function")
def pyplot():
"""Setup and teardown fixture for matplotlib.
This fixture checks if we can import matplotlib. If not, the tests will be
skipped. Otherwise, we close the figures before and after running the
functions.
Returns
-------
pyplot : module
The ``matplotlib.pyplot`` module.
"""
pyplot = pytest.importorskip("matplotlib.pyplot")
pyplot.close("all")
yield pyplot
pyplot.close("all")
# 第 5 章 —— sklearn/conftest.py (250-262)
@pytest.fixture
def hide_available_pandas(monkeypatch):
"""Pretend pandas was not installed."""
import_orig = builtins.__import__
def mocked_import(name, *args, **kwargs):
if name == "pandas":
raise ImportError()
return import_orig(name, *args, **kwargs)
monkeypatch.setattr(builtins, "__import__", mocked_import)
这两个 fixture 的实现展示了截然不同但同样精妙的资源管理策略:pyplot fixture 通过在 yield 前后分别清理图形资源来确保没有泄漏,而 hide_available_pandas fixture 则通过拦截特定的导入请求来模拟依赖缺失。前者关注的是“事后清理”(确保测试之间互不干扰),后者关注的是“事前准备”(创建受控的测试环境),二者共同支撑起了 scikit-learn 测试套件在各种环境下的可靠运行。
代码解析:pyplot fixture 体现了“前置清理 + yield 提供资源 + 后置清理”的完整资源生命周期管理模式。使用 function scope 确保每个测试函数都有独立的干净图形环境;importorskip 让缺少 matplotlib 的环境自动跳过而非报错。hide_available_pandas 则展示了“精准拦截 + 透传其他 + 自动还原”的依赖模拟模式:仅拦截顶层 'pandas' 导入利用了 Python 导入机制的级联失败特性,monkeypatch 自动还原消除了手动清理的负担,返回 None 强化了“声明即生效”的声明式使用模式。
5.10 全局 dtype、并行运行与 scipy_doctest 集成 —— 测试多样性的“三重保障”
在机器学习算法的测试中,数值精度往往是一个需要特别关注的维度。许多算法在 float32 下的表现与 float64 存在细微但重要的差别,因此 scikit-learn 提供了通过环境变量 SKLEARN_RUN_FLOAT32_TESTS 选择性启用 float32 测试的机制。global_dtype fixture 正是这个机制的核心实现:它使用 pytest.fixture(params=[...]) 声明了一个参数化 fixture,其中包含两个选项——np.float32(带有 skipif 标记)和 np.float64(默认运行)。当环境变量不为 '1' 时,float32 参数会被自动跳过;只有当显式设置为 '1' 时,才会同时运行 float32 和 float64 两套测试。这种参数级别的 skipif 标记(而非 fixture 级别)提供了更细粒度的控制,使得我们能够在同一个测试函数中仅对特定数据类型进行验证。
然而,float32 测试虽然能够揭示潜在的数值问题,但它也会显著增加测试时间,因此默认保持关闭状态是一个合理的折中方案。当需要进行彻底的精度验证时,开发者只需设置环境变量即可启用这套额外的验证层。fixture 使用 yield 风格而非 return 也是为未来可能的 teardown 逻辑预留空间——尽管目前没有实际的清理工作需要做,但这种设计为后续扩展提供了便利。
在并行测试的场景下,我们还需要考虑 pytest-run-parallel 这个可选插件的存在。当这个插件可用时,PARALLEL_RUN_AVAILABLE 会被设为 True,conftest 将不会注册任何并行相关的自定义标记,以避免与插件功能产生冲突;只有当插件不可用时,才会通过 pytest_addoption 注册 thread_unsafe_fixtures 配置项,这个配置项正是插件的功能替代品,让测试基础设施在插件缺失时仍能通过声明式的方式实现类似的线程隔离效果。模块导入时的版本检查则是一道重要的防线——它确保只有当满足最低 pytest 版本要求时,测试套件才会被加载,否则会直接抛出 ImportError 防止因版本不兼容而产生误导性的结果。
最后,我们来看看 scipy_doctest 的集成。当这个模块可用时,conftest 会从中导入 dt_config 对象,并在模块级代码中将其 strict_check 属性设置为 True。这个看似微小的配置其实有着重要的作用:它能够区分 doctest 中的 3.14(Python 原生浮点数)和 np.float64(3.14)(NumPy 标量)这样的细微差异。在默认情况下,由于 NumPy 标量的 repr 行为可能随版本变化,直接比较字符串表示可能导致假阳性;而启用严格模式后,我们就能确保 doctest 中的浮点数比较是真正意义上的数值相等,而不是依赖于特定的字符串表示形式。这种对细节的关注正是 scikit-learn 能够在长期演进中保持高质量标准的根源所在。
# 第 5 章 —— sklearn/conftest.py (80-84)
# 第 5 章 —— Global fixtures
@pytest.fixture(params=[pytest.param(np.float32, marks=_SKIP32_MARK), np.float64])
def global_dtype(request):
yield request.param
# 第 5 章 —— sklearn/conftest.py (200-203)
def pytest_addoption(parser, pluginmanager):
if not PARALLEL_RUN_AVAILABLE:
parser.addini("thread_unsafe_fixtures", "list of stuff")
# 第 5 章 —— sklearn/conftest.py (1-52)
# 第 5 章 —— ... (省略版权信息和导入语句)
if parse_version(pytest.__version__) < parse_version(PYTEST_MIN_VERSION):
raise ImportError(
f"Your version of pytest is too old. Got version {pytest.__version__}, you"
f" should have pytest >= {PYTEST_MIN_VERSION} installed."
)
# 第 5 章 —— sklearn/conftest.py (265-268)
if dt_config is not None:
# Strict mode to differentiate between 3.14 and np.float64(3.14)
dt_config.strict_check = True
# dt_config.rtol = 0.01
这三个代码块共同构成了测试基础设施的多重保障机制:global_dtype fixture 提供了数据类型的参数化测试能力,pytest_addoption 确保了在缺少并行插件时仍能通过配置项实现线程隔离,而模块级代码则负责版本检查、fixture 注册以及启用 scipy_doctest 的严格模式。通过这种分层设计,我们不仅能够在不同的数据精度下验证算法行为,还能在并行与串行环境下保持一致的测试体验,并且确保 doctest 在面对 NumPy 和 SciPy 演进时仍能提供可靠的验证。
代码解析:global_dtype 采用“参数级 skipif + yield 风格”的设计,实现了“默认 float64,可选 float32”的精度测试矩阵,参数级标记比 fixture 级标记粒度更细,允许单测试函数内混合精度验证。pytest_addoption 体现了“插件优先、配置兜底”的兼容策略:有插件用插件,无插件用 ini 配置,避免功能重复与冲突。模块级代码的导入时版本检查是“快速失败”原则的体现——尽早发现环境不匹配,避免后续难以诊断的怪异错误。dt_config.strict_check = True 则解决了 NumPy 标量 repr 版本演进导致的 doctest 脆弱性问题,将字符串比较升级为数值比较,是长期维护质量的关键一招。
5.11 设计中的取舍
在构建这套测试基础设施时,scikit-learn 团队在多个关键决策点上进行了深思熟虑的权衡。
首先是随机种子控制机制的选择。为什么不直接在代码中硬编码固定种子,而是设计环境变量驱动的三模式参数化方案?固定种子虽然实现简单,但无法支持“遍历所有种子”这种彻底验证需求。环境变量方法让我们既能保持日常开发的默认可复现性(使用种子 42),又能在需要时通过设置 SKLEARN_TESTS_GLOBAL_RANDOM_SEED=all 来运行 100 种种子的全面检查,甚至可以精确指定区间如 '10-20' 来复现特定失败场景。这种灵活性是固定方案无法提供的,而实现复杂度的增加主要集中在 pytest_generate_tests 的解析逻辑中,这是一次性投入、长期受益的设计。
其次是 _fetch_fixture 的三层嵌套结构。相比于直接在 fixture 中读取环境变量,装饰器工厂模式将“环境变量读取时机固化在模块导入期”这一关键决策显式化了——download_if_missing 在装饰器定义时就被确定,不会随测试执行时的环境变化而变化,这保证了整个测试会话的一致性。同时,@wraps(f) 保留原函数元数据、lambda 包装适配 pytest fixture 协议、精准捕获特定 OSError 而透传其他异常,这些细节处理共同构成了一个既安全又易用的数据集获取抽象。虽然增加了理解门槛,但它解决了“离线环境优雅跳过”、“在线环境按需下载”、“异常精准分类”三个核心问题,收益远超成本。
再次是 pytest_collection_modifyitems 中的数据预下载调度。将下载时机从测试执行期前移到收集完成后,虽然增加了收集阶段的逻辑复杂度(需要遍历所有 item、区分普通测试与 Doctest、处理 gw0 worker 单点下载),但它成功避免了 pytest-xdist 并行测试中的数据下载竞态问题——多进程同时写入同一缓存文件会导致文件损坏和不可预测错误,这是正确性层面的硬性约束。在正确性与实现简洁性的权衡中,scikit-learn 选择了正确性,并通过清晰的代码结构和详细注释控制了认知负载。
最后是 doctest 的作用域隔离与条件跳过策略。将 dtest.globs 设为空字典强制示例自包含,虽然要求文档编写者在每个示例中重复 import,但消除了隐式依赖带来的不稳定性,这是对“显式优于隐式”原则的坚持。四重条件门禁(matplotlib、32位、Windows、NumPy/SciPy 版本)看似严苛,实则是对 doctest 极度敏感特性的务实应对:与其在不支持的环境中产生噪音失败,不如在收集阶段就静默跳过。contextmanager 的特例排除则是对 pytest 已知 Bug 的务实规避——完美是良好的敌人,在工程实践中有时必须接受局部的不完美以换取整体的稳健。
5.12 动手练习
-
阅读 pytest_generate_tests 的种子解析逻辑:尝试设置
SKLEARN_TESTS_GLOBAL_RANDOM_SEED=5-10运行测试,观察参数化生成的种子数量;再尝试设置SKLEARN_TESTS_GLOBAL_RANDOM_SEED=150观察报错信息。 -
分析 _fetch_fixture 的三层嵌套结构:编写一个简化版的装饰器工厂,接收一个函数并返回一个 pytest fixture,其中包含环境变量控制的行为开关和异常捕获跳过逻辑。
-
理解 pytest_collection_modifyitems 的下载调度:在使用 pytest-xdist 并行运行测试时(
pytest -n 4),观察控制台输出中数据集下载只发生在 gw0 worker 中的现象;尝试设置SKLEARN_SKIP_NETWORK_TESTS=1观察网络测试被标记跳过。 -
探索 pyplot 与 hide_available_pandas 的实现细节:编写一个使用 pyplot fixture 的测试,在测试中创建多个图形,验证测试前后图形数量均为 0;编写一个使用 hide_available_pandas 的测试,验证在 fixture 作用域内
import pandas抛出 ImportError 而import numpy正常工作。 -
分析 raccoon_face_or_skip 与 dt_config 的多重检测逻辑:阅读 scipy.datasets.face 的源码理解其下载机制;查阅 scipy_doctest 文档理解 strict_check 的具体比较行为差异。
5.13 本章小结
本章我们深入探讨了 scikit-learn 的测试基础设施,理解了它如何通过一系列精心设计的 pytest 钩子和 fixture 来构建可靠、可重复的测试环境。我们首先从测试会话的启动配置开始,看到 pytest_configure 如何统一设置 matplotlib 后端、限制线程数、管理警告和注册自定义标记;然后深入研究了 pytest_generate_tests 如何通过环境变量实现随机种子的参数化,支持默认、全部和区间三种模式来控制测试的确定性程度;接着我们探究了 _fetch_fixture 装饰器如何优雅地处理网络受限环境下的数据集获取,通过三层嵌套结构实现按需下载和离线跳过;我们还分析了 pytest_collection_modifyitems 如何在测试收集阶段充当调度中心,安全地预下载数据集以避免并行测试中的资源冲突,同时处理平台相关的 xfail/skip 标记和 doctest 作用域隔离;随后我们考察了 pyplot fixture 如何通过 yield 机制管理 matplotlib 图形的完整生命周期,以及 hide_available_pandas 如何通过 monkeypatch 拦截导入来模拟依赖缺失场景;最后我们理解了 global_dtype fixture 如何提供 float32/float64 的参数化测试,以及 dt_config 严格模式如何确保 doctest 中的浮点数比较不受 NumPy 标量表示变化的影响。
这一章中我们学习了 pytest 钩子函数在测试会话生命周期中的作用,掌握了通过环境变量实现确定性随机测试的机制,理解了数据集获取的按需供应策略,探究了并行测试中的线程安全和数据下载调度,学会了 matplotlib 资源管理和依赖模拟的技术,以及掌握了特殊数据集访问控制和 doctest 环境适配的方法。
以下是本章核心概念对照表,供快速查阅与复习:
| 概念 | 解释 |
|------|------|
| pytest_configure(config) | 测试会话启动钩子,统一配置 matplotlib 后端、线程数限制、警告升级机制和自定义 marker 注册 |
| pytest_generate_tests(metafunc) | 基于环境变量动态参数化 global_random_seed,支持默认/全部/区间三种种子模式 |
| _fetch_fixture(f) | 装饰器工厂,将数据集 fetch 函数包装为支持网络跳过和按需下载的 pytest fixture |
| pytest_collection_modifyitems(config, items) | 收集完成后统一调度:预下载数据集、添加平台相关 xfail/skip 标记、隔离 doctest 作用域、检测 PIL 依赖 |
| global_dtype(request) | 参数化 fixture,使用 yield 风格,默认 float64,通过环境变量选择性启用 float32 测试 |
| pyplot() | matplotlib 测试 fixture,在测试前后关闭所有图形,导入失败时跳过测试 |
| hide_available_pandas(monkeypatch) | 通过 monkeypatch 替换 builtins.import,模拟 pandas 未安装场景 |
| pytest_addoption(parser, pluginmanager) | 在 pytest_run_parallel 插件不可用时注册 thread_unsafe_fixtures 配置项 |
| raccoon_face_or_skip() | 检查网络访问和 pooch 依赖,不可用时跳过测试,否则返回灰度浣熊脸图像 |
| dataset_fetchers | 数据集 fixture 名称到 fetch 函数的映射字典,用于收集阶段的预下载调度 |
| __main__(模块级全局代码) | 模块导入时执行:检测并行插件、校验 pytest 版本、注册所有数据集 fixture、配置 dt_config 严格模式 |
| threadpool_limits | 根据物理核心数和 xdist worker 数动态限制 OpenMP/BLAS 线程,防止过载 |
| dt_config.strict_check | 配置 scipy_doctest 的严格检查模式,区分 Python 标量与 NumPy 标量的 repr 差异 |
下一章中,我们将学习构建系统与发行商定制 —— 解析“从源码到安装包的旅程”。
第 6 章 —— 构建系统与发行商定制 —— 解析“从源码到安装包的旅程”
6.1 学习目标
-
难度:★★★☆☆(3/5)
-
预备知识:Python 基础、面向对象编程与 Markdown/代码阅读基础
-
理解 Tempita 模板引擎在构建系统中的代码生成机制,掌握从 .tp 模板到源文件的转换流程
-
掌握构建脚本的命令行参数解析与输入验证设计,理解扩展名白名单和快速失败策略
-
了解版本号动态提取的实现原理,理解如何避免硬编码版本带来的维护问题
-
理解发行商初始化钩子的设计意图,掌握源码包与发行包工程分离的实践
-
能独立阅读并修改构建工具链中的模板处理脚本
6.2 生活类比
想象构建系统是一家定制印刷厂的排版车间:当出版一本需要针对不同地区定制内容的书籍时,印刷厂不会为每个地区单独设计一套印版,而是采用智能的模板化方案。Tempita 模板就像带有可变插槽的活字印刷版(比如 .pyx.tp 文件中到处可见的 {{variable}} 占位符),这些插槽专门留给将要注入的个性化信息——比如目标平台的特殊路径、当前构建的版本号或者编译器特定的宏定义。而 process_tempita() 函数则是经验丰富的排版工人,他们熟练地将从订单系统中调取的具体数据(就像从数据库里查到的“北京分校需要用 simd 指令”、“版本号是 1.5.0”)一一填入那些预留的空位,完成最终可用于印刷的定制印版。车间入口还有严格的质检员——扩展名验证机制,它就像门卫只放行带有 .tp 后缀的模板文件,决不允许有人误把已经生成好的 .pyx 源码当作输入模板送进来,这样既防止了源码被意外覆盖,又保证了每次构建都从纯净的模板开始。与此同时,版本号提取器默默地在仓库里工作,它不依赖于已经印好的书籍封面,而是直接去总账本(sklearn/__init__.py)里抄录当前批次的编号,这避免了版本号需要在多个地方同步修改的维护噩梦。至于发行商初始化钩子 _distributor_init.py,它就像为不同地区的分销商预留的“本地化改装区”:在标准出版社印好的书本基础上,北京发行商可能需要在这空白页里加入当地的教学大纲对应表,而柏林分销商则可能在此处植入符合欧盟隐私法规的数据处理声明——所有这些定制都不会触及核心教材内容本身,实现了源码的纯净与平台定制之间的完美平衡。
6.3 源码地图
sklearn/_build_utils/tempita.py
├── process_tempita(fromfile, outfile=None) # 核心模板处理函数
│ ├── open(fromfile, 'r', encoding='utf-8') # 读取 .tp 模板文件
│ ├── tempita.Template(template_content) # 解析模板语法({{variable}})
│ ├── template.substitute() # 替换占位符生成最终内容
│ └── open(outfile, 'w', encoding='utf-8') # 写入输出文件
├── main() # 命令行入口
│ ├── argparse.ArgumentParser() # 定义三个参数:infile、--outdir、--ignore
│ ├── args.infile.endswith('.tp') # 扩展名白名单验证
│ ├── os.path.join(os.getcwd(), args.outdir) # 输出目录绝对路径计算
│ ├── os.path.splitext(os.path.split(...)) # 去除 .tp 后缀生成输出文件名
│ └── process_tempita(args.infile, outfile) # 调用模板处理函数
sklearn/_build_utils/version.py
├── sklearn_init = os.path.join(...) # 定位 sklearn/init.py
├── data = open(sklearn_init).readlines() # 逐行读取包入口文件
├── version_line = next(line for line in data if line.startswith('version')) # 定位版本行
├── version = version_line.strip().split(' = ')[1].replace('"', '').replace("'", '') # 清洗提取版本号
└── print(version) # 输出版本号到标准输出
sklearn/_distributor_init.py
├── 模块文档字符串 # 说明发行商可注入自定代码的用途
└── 空实现 # 标准源码包中不含任何业务代码
sklearn/_build_utils/init.py
└── 空文件 # 包标记,使目录成为 Python 包
6.4 Tempita 模板引擎 —— 代码生成的“活字印刷机”
在 scikit-learn 的构建流程中,我们经常需要根据不同的编译环境生成略有差异的源文件。例如,某些平台可能需要在 Cython 文件中加入特定的编译器指令,或者根据是否启用调试模式来包含不同的调试代码。如果我们为每种组合都维护一套几乎相同的 .pyx 或 .c 文件,不仅工作量巨大,而且一旦核心逻辑需要更新,就必须在所有副本中同步修改——这无疑是噩梦般的维护负担。为了解决这个问题,构建系统引入了模板引擎的概念:我们只需维护一套带占位符的模板文件(如 foo.pyx.tp),在构建时根据实际环境动态填充这些占位符,从而生成目标源文件(如 foo.pyx)。这种方式保证了源码的唯一性,同时又能灵活应对平台差异。
为什么需要模板引擎?
-
Cython 编译前需要根据配置生成不同版本的源文件
-
避免手工维护多份几乎相同的 .pyx/.c 代码
-
构建时通过模板变量注入平台相关或版本相关的差异
process_tempita 的核心三步流程
-
读取 .tp 模板文件的原始文本
-
调用
tempita.Template解析模板语法(如{{variable}}) -
执行
substitute()将占位符替换为实际值并写入输出文件
文件扩展名的约定
-
输入必须是
.c.tp或.pyx.tp结尾 -
输出文件名自动去除
.tp后缀,保持与源码相同的基名
为什么用 Cython 内置的 Tempita 而非独立依赖?
-
Cython 本身已经依赖 Tempita,无需引入额外包
-
保持构建依赖的最小化与确定性
源码路径:sklearn/_build_utils/tempita.py - process_tempita()(16-29行)
// 逐行注释解释
def process_tempita(fromfile, outfile=None):
"""Process tempita templated file and write out the result.
The template file is expected to end in `.c.tp` or `.pyx.tp`:
E.g. processing `template.c.in` generates `template.c`.
"""
with open(fromfile, "r", encoding="utf-8") as f: # 以 UTF-8 编码读取模板文件
template_content = f.read()
template = tempita.Template(template_content) # 解析模板语法,识别 {{variable}} 占位符
content = template.substitute() # 执行替换,生成最终内容
with open(outfile, "w", encoding="utf-8") as f: # 以 UTF-8 编码写入输出文件
f.write(content)
这段代码定义了模板处理的核心逻辑。它首先以 UTF-8 编码读取输入的模板文件,将其内容存入 template_content 变量;接着使用 Cython 自带的 tempita.Template 类解析这一文本内容,其中识别出所有形如 {{variable}} 的占位符;随后调用 substitute() 方法用实际值替换这些占位符,得到最终要写出的内容;最后将结果写入指定的输出文件。整个过程封装在文件操作的上下文管理器中,确保即使出现异常也能安全关闭文件句柄。
Mermaid 流程图
6.5 命令行参数解析 —— 构建脚本的“安全准入闸门”
tempita.py 的 main() 函数是脚本的命令行入口,它承担着参数解析、输入校验、路径计算与调度核心逻辑的职责。其核心类型定义体现在 argparse.ArgumentParser 所构建的参数模型上:infile 为位置参数,指定待处理的模板文件路径;--outdir 为必选的命名参数,指定输出目录;--ignore 为可选的哑元参数,供构建系统(如 Meson)声明伪依赖以触发重构,自身不参与实际处理逻辑。
以下逐行解析关键函数 main() 的实现细节:
def main():
parser = argparse.ArgumentParser() # 创建参数解析器
parser.add_argument("infile", type=str, help="Path to the input file") # 必填:输入模板文件路径
parser.add_argument("-o", "--outdir", type=str, help="Path to the output directory") # 必填:输出目录
parser.add_argument( # 可选:哑元参数,用于构建系统依赖追踪
"-i",
"--ignore",
type=str,
help=(
"An ignored input - may be useful to add a "
"dependency between custom targets"
),
)
args = parser.parse_args() # 解析命令行参数
if not args.infile.endswith(".tp"): # 白名单验证:仅允许 .tp 后缀
raise ValueError(f"Unexpected extension: {args.infile}")
if not args.outdir: # 快速失败:缺失 --outdir 立即报错
raise ValueError("Missing `--outdir` argument to tempita.py")
outdir_abs = os.path.join(os.getcwd(), args.outdir) # 将输出目录转为绝对路径
outfile = os.path.join( # 构建输出文件完整路径
outdir_abs, os.path.splitext(os.path.split(args.infile)[1])[0] # 去除 .tp 后缀
)
process_tempita(args.infile, outfile) # 调用核心处理函数
输入验证的设计哲学体现在两个关键检查点上。首先是扩展名白名单机制:if not args.infile.endswith(".tp") 确保脚本只接受以 .tp 结尾的文件作为输入,这不仅是一种约定,更是一道安全闸门——防止已生成的源文件(如 foo.pyx)被误作为模板输入而导致源码被意外覆盖。其次是快速失败策略:对 --outdir 的缺失采取立即抛出 ValueError 的方式,而不是尝试使用当前目录作为后备。这种“宁可错杀一千,不可放走一”的态度在构建系统中尤为重要,因为一个静默的错误可能导致整批发行包都带着看不见的缺陷。
输出路径的计算逻辑蕴含着对跨平台兼容性的考虑。它先通过 os.path.split(args.infile)[1] 提取输入文件的纯文件名(去除目录部分),然后用 os.path.splitext 去掉这个文件名的 .tp 后缀,接着将得到的基名(例如从 src/foo.pyx.tp 中得到的 foo.pyx)与之前计算得到的绝对输出目录路径拼接起来。这种做法确保了无论输入文件是以绝对路径还是相对路径给出,输出文件总是能正确放置在用户指定的目录下,同时保持与源码约定一致的命名方式。最后,main() 调用 process_tempita 完成实际的模板处理工作,实现了参数解析、验证和路径计算与核心模板逻辑的解耦。
Mermaid 流程图
6.6 版本号提取器 —— 从包入口“抽取灵魂数字”
在软件发布过程中,版本号是一个贯穿始终的重要信息:它需要出现在发行包的元数据中,可能需要写入生成的二进制文件,更重要的是,它必须和源码树中所标记的版本完全一致。如果我们在构建脚本、setup.py 以及 __init__.py 里都各自维护一份版本号拷贝,那么一旦需要更新版本,就必须记住去所有地方改动——而一旦遗漏,就可能导致用户通过 pip install 得到的包版本声明和实际运行时检测到的版本不一致,这无疑会给调试带来噩梦。scikit-learn 采取了更为聪明的做法:将版本号的唯一来源定位在 sklearn/__init__.py 中的 __version__ 变量,构建过程中通过专门的脚本直接从这个权威来源读取版本号,这样无论是开发者还是自动化构建流程,都只需要在一个地方维护版本信息。
为什么版本号要动态提取而不是硬编码?
-
避免多处维护版本号导致的不一致
-
以
sklearn/__init__.py中的__version__为唯一权威来源
实现方式的巧妙之处
-
无需 import sklearn 包体(避免触发重量级初始化)
-
直接以文本模式读取
__init__.py并逐行扫描 -
用
startswith("__version__")定位版本行,避免复杂正则
字符串清洗流程
-
strip()去除首尾空白 -
split(" = ")[1]取出等号右侧内容 -
连续
replace去除双引号和单引号,兼容两种定义风格
输出用途
-
供构建工具链(如 Meson、CI 脚本)读取当前版本
-
用于生成 wheel/sdist 的版本元数据
源码路径:sklearn/_build_utils/version.py - 全局执行逻辑(9-16行)
// 逐行注释解释
sklearn_init = os.path.join(os.path.dirname(__file__), "../__init__.py") # 定位 sklearn/__init__.py
data = open(sklearn_init).readlines() # 逐行读取文件内容
version_line = next(line for line in data if line.startswith("__version__")) # 查找以 __version__ 开头的行
version = version_line.strip().split(" = ")[1].replace('"', "").replace("'", "") # 清洗提取版本号
print(version) # 输出到标准输出
这段代码实现了版本号的动态提取。它首先构建指向 sklearn/__init__.py 的相对路径(相对于当前脚本所在目录),然后以文本模式逐行读取这个文件内容。值得注意的是,它故意避免了直接 import sklearn——这样做虽然能最直接获得 __version__,但会触发 scikit-learn 包的完整初始化过程,在构建环境中可能既不必要又带来不必要的开销。取而代之的是纯文本扫描:它逐行检查每一行是否以 __version__ 开头,一旦找到匹配的行就立即停止搜索(得益于 next() 函数的惰性求值)。得到版本赋值语句后,它依次执行三步清洗:先用 strip() 去除可能的首尾空白字符,然后以 = (注意前后各有一个空格)为分隔符切割字符串并取第二部分(即等号右侧),最后连续调用两次 replace 方法去除可能包裹的双引号或单引号,这样无论版本号是在 __init__.py 中以 __version__ = "1.5.0" 还是 __version__ = '1.5.0' 的形式出现,都能正确提取出纯净的版本号字符串。最终结果通过 print() 输出到标准输出,便于其他构建脚本通过管道或命令替换捕获使用。
Mermaid 流程图
6.7 发行商初始化钩子 —— 下游定制的“预留接口”
在开源软件的生态中,同一份源码往往会被不同的组织重新打包发行以适应他们的特定需求:例如,Anaconda 可能希望在其 conda-forge 渠道下的 scikit-learn 包中预置某些数据科学常用插件的初始化代码;Debian 维护者可能需要在这个包中加入符合该发行版政策的特定路径处理;而官方 Windows wheel 的制作者则可能希望在导入时自动解压并优先加载打包好的 DLL 文件以确保 Cython 编译扩展能正确运行。如果这些平台特定的需求直接混入主源码树,不仅会污染源码的纯净度,而且每次上游更新都需要小心翼翼地合并这些分支更改——这显然不是一个可持续的模式。
为了解决这个问题,scikit-learn 在其构建和发行流程中预留了一个专门的定制接口:_distributor_init.py 文件。这个文件的设计理念非常简单却十分 powerful:在标准源码包中,这个文件被刻意保持为空(除了必要的文档头和许可证声明之外),但它在包的 __init__.py 中被非常早地导入——事实上,它是导入列表中的第一项。这意味着当用户在他们的环境中执行 import sklearn 时,实际上第一件被加载和运行的代码就是这个发行商初始化文件。下游发行商如果需要注入任何平台特定的逻辑,他们只需要在自己重新打包的过程中替换掉这个原本为空的文件,放入他们自己的初始化代码;由于这个替换发生在源码树之外(属于发行过程的一部分),上游的核心业务代码保持完全不变,因而当上游发布新版本时,下游只需要用他们维护的定制文件去覆盖对应位置即可,无须处理任何复杂的合并冲突。
为什么需要 _distributor_init.py?
-
允许各发行版(conda-forge、Debian、Windows wheel 等)注入平台特定逻辑
-
例如:Windows 上优先加载打包的 DLL 文件
-
检查硬件要求(如 CPU 指令集支持)
设计约束与安全边界
-
标准源码包中该文件为空,仅保留文档说明
-
下游发行商可以直接替换此文件而不影响业务代码
-
文件在
__init__.py最早期被导入,确保初始化逻辑先行
与构建流程的联动
-
build_tools/github/vendor.py会动态改写此文件以嵌入 DLL 路径 -
实现了“源码保持纯净,发布时定制”的工程分离
源码路径:sklearn/_distributor_init.py - 全局模块定义(1-11行)
"""Distributor init file
Distributors: you can add custom code here to support particular distributions
of scikit-learn.
For example, this is a good place to put any checks for hardware requirements.
The scikit-learn standard source distribution will not put code in this file,
so you can safely replace this file with your own version.
"""
# 第 6 章 —— Authors: The scikit-learn developers
# 第 6 章 —— SPDX-License-Identifier: BSD-3-Clause
这段代码看起来似乎“几乎什么都没做”,但正是这种有意的留白赋予了它巨大的灵活性。文件开头是一个详尽的模块文档字符串,清晰地解释了这个文件的目的:它是为下游发行商预留的定制空间,他们可以在这里放置任何需要在 scikit-learn 导入时尽早执行的代码——无论是检测硬件能力、设置环境变量,还是如前所述的在 Windows 平台上优先加载本地打包的 DLL 文件以避免“找不到模块”的错误。文档还特别强调了在标准源码发行中这个文件保持为空的事实,并明确鼓励发行商“安全地”用他们自己的版本替换它。代码主体其实只有版权声明和 SPDX 许可证标识符——这些是法律上必要的元数据,不承载任何功能逻辑。正是这种“几乎为空却又必不可少”的设计,使得 _distributor_init.py 成为连接上游源码纯净与下游平台定制之间最优雅的桥梁:上游开发者可以专注于改进机器学习算法而无需考虑打包细节;下游发行商则拥有了一个干净、well-defined 的入口点来注入他们的平台特定逻辑,而两者之间不需要协调任何复杂的接口或担心未来的兼容性问题。
Mermaid 流程图
6.8 _build_utils 包入口 —— 模块发现的“隐形路标”
在 Python 中,一个目录只有在包含 __init__.py 文件时才会被认作是一个包,从而才能被其他模块通过 import 语句导入。这个看似微小的文件实际上承担着重要的角色:它不仅标志着“这里是一个可导入的单元”,而且在被导入时会执行其中的代码(哪怕是空文件也会执行,只不过没有实际效果)。对于 scikit-learn 的构建工具链来说,_build_utils 目录下零散着几个实用脚本(如 tempita.py 和 version.py),这些脚本在构建过程中需要相互调用——例如,主构建配置可能需要先调用 version.py 得到当前版本,再根据这个版本去处理某些模板文件。如果 _build_utils 目录缺少 __init__.py,那么尽管其中的 .py 文件依然存在并且可以直接通过文件路径被执行,但其他构建脚本就无法通过诸如 from sklearn._build_utils import tempita 这样的优雅导入语句来使用这些工具了,这既破坏了代码的组织性,也增加了出错的可能性。
空白 __init__.py 的作用
-
将
_build_utils目录标记为 Python 包,使构建工具能正确导入其子模块 -
本身不导出任何公共符号,属于内部基础设施模块
下划线前缀的命名约定
-
_build_utils以单下划线开头,语义上表示“私有模块” -
对外不承诺 API 稳定性,可随版本自由重构
构建系统的模块发现机制
-
Meson 或其他构建工具通过文件路径遍历来确定可执行脚本
-
该目录下的
tempita.py和version.py均可作为独立脚本运行
源码路径:sklearn/_build_utils/__init__.py - 全局模块定义(空文件)
这个文件确实完全空白——没有任何代码,甚至没有注释。但正是这种“有意的虚无”完成了它的使命:通过仅仅存在(而不包含任何内容),它向 Python 解释器宣告 _build_utils 目录应该被视为一个可导入的包。值得注意的是,这个目录名以单下划线开头,遵循了 Python 中的一个约定俗成的命名惯例:单下划线前缀通常被用来表示“内部使用”的或“私有的”模块,暗示这个包虽然可以被导入,但其 API 并不打算作为公开的、稳定的对外契约——换句话说,scikit-learn 的开发团队保留了在未来版本中根据需要重构甚至彻底移除这个目录内部结构的自由,而不会因此破坏对外界承诺的功能。尽管被标记为“私有”,但 _build_utils 目录下的脚本如 tempita.py 和 version.py 在构建过程中确实经常需要被其他工具(如 Meson 构建配置或 CI 脚本)直接调用;好在这些脚本本身被设计成可以独立运行的(它们都带有 if __name__ == "__main__": 主入口检查),所以即使不通过包导入机制,它们也能够被当作独立的可执行文件来使用,这为构建系统提供了很好的灵活性:既能享受包导入带来的命名空间整洁(在包内部互相引用时),又不失通过直接脚本路径调用的直接性和可调试性。
Mermaid 架构图
6.9 设计中的取舍
在设计构建系统时,我们总是在不同的目标之间寻找平衡点。对于模板引擎的选择,scikit-learn 采用了 Cython 内置的 Tempita 而非更通用的 Jinja2。虽然 Jinja2 功能更强大,支持复杂的控制流和过滤器,但会引入一个额外的运行时依赖;而 scikit-learn 的构建系统已经依赖于 Cython 来编译核心扩展,Cython 本身就捆绑了 Tempita 作为其模板处理器,因此直接使用这个内置方案不仅能避免增加新依赖,还能保证构建环境中的一致性:只要能编译 Cython 扩展,就一定能运行 Tempita。这种“就地取材”的思路极大地简化了构建流程的准备工作,特别是在 CI 环境中,减少一种依赖意味着少了一层可能出错的地方。在版本号提取的实现上,代码选择了逐行文本扫描而非正则表达式。乍看之下,正则似乎能更“优雅”地一行搞定,但实际使用中这种方法更易出错:版本号可能被注释追随,或者等号两侧的空格数量不固定;相比之下,逐行扫描配合精确的字符串切割和清洗步骤,虽然看起来略显啰嗦,但其实更易于调试和维护——每一步都有明确的目的和可预期的结果,当出现问题时也更容易定位是哪一步的处理逻辑出了偏差。至于发行商初始化钩子 _distributor_init.py 为何保持为空而不是放置一些通用的默认实现,这其实是一种有意识的“反抽象”设计:如果我们在这个文件里放置了某些默认行为(比如总是尝试加载某个通用的辅助库),那么下游发行商想要覆盖或修改这种行为时就不得不面对“如何优雅地覆盖上游默认实现”这个额外的复杂度;而保持完全为空则意味着下游完全掌控权——他们想放什么就放什么,不需要考虑如何与上游默认行为进行交互或覆盖,真正实现了“零干扰”的定制能力。
6.10 动手练习
-
阅读 Tempita 模板处理流程
-
追踪版本号提取逻辑
-
理解发行商初始化钩子
-
模拟模板处理脚本
-
修改版本号提取器
6.11 本章小结
在这一章中,我们深入探究了 scikit-learn 构建系统的内部工作机制,从模板驱动的代码生成到版本信息的动态提取,再到为下游发行商预留的定制接口。我们首先了解了如何通过 Tempita 模板引擎实现源码的纯净与平台定制之间的平衡,接着考察了构建脚本如何通过严格的参数验证和路径计算确保操作的安全性,然后学习了版本号提取器如何避免硬编码带来的维护噩梦,最后理解了那个看似空白却关键的 _distributor_init.py 文件如何实现源码包与发行包的工程分离。
这一章中我们学习/了解/讨论了构建系统的核心组成部分。首先我们理解了 Tempita 模板引擎如何通过三步流程将带占位符的模板文件转换为目标源码;其次我们掌握了命令行参数解析中的白名单验证和快速失败策略如何保证操作安全;接着我们知道了版本号动态提取如何以 __init__.py 为权威来源避免多处维护;然后我们探究了发行商初始化钩子如何为下游定制预留接口而不污染源码;最后我们认识了空的 __init__.py 如何使 _build_utils 目录成为可导入的包。
本章我们一起学习了以下概念:
| 概念 | 解释 |
|------|------|
| process_tempita() | 模板处理核心函数,读取 .tp 文件、解析 Tempita 语法并输出最终源码 |
| tempita.Template | Cython 内置的模板引擎,解析 {{variable}} 占位符语法 |
| .tp 扩展名验证 | 输入文件必须 .tp 结尾,白名单机制防止意外覆盖源码 |
| --outdir 快速失败 | 输出目录缺失时立即抛 ValueError,避免静默错误 |
| 版本号动态提取 | 从 init.py 文本读取 version,避免多处维护版本号 |
| _distributor_init.py | 下游发行商注入平台特定逻辑的预留接口,源码包中为空 |
| _build_utils 包标记 | 空 init.py 将目录标记为 Python 包,支持模块导入 |
| 工程分离设计 | 源码保持纯净,发布时通过 vendor.py 动态改写发行商文件 |
下一章中,我们将学习实验性功能启用机制 —— 打开“未来特性的传送门”。
第 7 章 —— 实验性功能启用机制 —— 打开“未来特性的传送门”
7.1 学习目标
-
难度:★★★☆☆(3/5)
-
预备知识:Python 基础、面向对象编程与 Markdown/代码阅读基础
-
理解 sklearn.experimental 包的设计目的与风险警示机制
-
掌握 setattr 动态注入与 all 更新实现实验性功能启用的原理
-
了解弃用启用器的演变路径与向后兼容策略
-
能独立为新实验性估计器编写启用模块
7.2 生活类比
想象你走进一家科技公司的创新实验室,门口有明显的警示标志:“此区域展示的产品仍在测试阶段,性能和接口可能随时调整,使用需自担风险。”实验室里摆放着最新的原型机器学习估计器——它们可能明天就成为主力产品,也可能因为表现不佳被悄悄撤下。为了让好奇的用户能够提前体验这些前沿功能,同时又不干扰正式产品展厅的秩序,公司设计了一套门禁系统:想要进入实验室的访客必须主动出示专门的入场凭证(显式导入 enable_* 模块),凭证一旦被使用,系统就会在后台悄悄将新产品的样品放到展厅的展示架上(通过 setattr 动态注入到目标模块),并更新导览手册(更新 all 列表),让游客知道刚刚多了哪些新展品。入口处还挂着详细的试用须知(模块文档字符串的 warning),明确告知这些产品不享受常规的退役流程。而当某个原型最终成熟转正时,原来的门禁卡并不会立刻被回收——而是,系统保留该文件但改为只发出一条友好的提示:“此产品已正式上线,无需再使用门禁卡”,这样既避免了突然断裂老用户的工作流,又为彻底退出留出了缓冲期。正是这种“创新探索”与“API 稳定”之间的平衡艺术,让 scikit-learn 能够在保持核心框架可靠性的同时,持续向用户输送最前沿的算法尝鲜。
7.3 源码地图
7.4 实验性功能包入口 —— 通往“未来特性传送门”的门厅
scikit-learn 的实验性功能通过一个顶层包 sklearn.experimental 进行组织和隔离。这个包的核心职责不是直接提供功能,而是作为一个明确的边界:凡是在这里定义的特性,都标注着“尚未稳定”的身份证。
我们先来看看这个包的大门长什么样子——也就是 __init__.py 文件。它非常简洁,只有模块文档字符串和一些作者信息,但这几段文字承担着重要的风险提示作用。
7.4.1 核心类型定义
源码路径:sklearn/experimental/__init__.py - __main__(1-16行)
"""Importable modules that enable the use of experimental features or estimators.
.. warning::
The features and estimators that are experimental aren't subject to
deprecation cycles. Use them at your own risks!
"""
# 第 7 章 —— Authors: The scikit-learn developers
# 第 7 章 —— SPDX-License-Identifier: BSD-3-Clause
7.4.2 逐行解析关键逻辑
这段文档字符串使用了 Sphinx 的 warning 指令,在生成的文档中会以显眼的警告框形式呈现。它的意思非常直白:实验性功能不受常规的弃用政策约束——这意味着它们可能在任何小版本中被修改、移除或彻底改写接口,而不会像普通 API 那样先发布弃用警告再过渡一段时间。使用它们的用户需要自行承担这种不确定性。顶层包本身不导出任何具体的估计器或函数,它的存在只是为了提供一个统一的命名空间,让所有实验性启用器都有明确的归属(如 sklearn.experimental.enable_halving_search_cv),并通过其文档字符串一次性向所有潜在用户传达风险信息。
7.5 逐次减半搜索启用器 —— 打开“超参数寻优的加速门”
现在我们进入具体的启用器机制。以 enable_halving_search_cv.py 为例,这个文件的作用是在用户显式导入它之后,把两个实验性的超参数搜索算法——HalvingGridSearchCV 和 HalvingRandomSearchCV——注入到 sklearn.model_selection 模块中,使得用户之后可以像使用普通估计器一样从 sklearn.model_selection 导入它们。
我们逐行解析这个启用器是如何完成这个“魔法”的。
7.5.1 核心类型定义
源码路径:sklearn/experimental/enable_halving_search_cv.py - __main__(17-37行)
"""Enables Successive Halving search-estimators
The API and results of these estimators might change without any deprecation
cycle.
Importing this file dynamically sets the
:class:`~sklearn.model_selection.HalvingRandomSearchCV` and
:class:`~sklearn.model_selection.HalvingGridSearchCV` as attributes of the
`model_selection` module::
>>> # explicitly require this experimental feature
>>> from sklearn.experimental import enable_halving_search_cv # noqa
>>> # now you can import normally from model_selection
>>> from sklearn.model_selection import HalvingRandomSearchCV
>>> from sklearn.model_selection import HalvingGridSearchCV
The ``# noqa`` comment comment can be removed: it just tells linters like
flake8 to ignore the import, which appears as unused.
"""
# 第 7 章 —— Authors: The scikit-learn developers
# 第 7 章 —— SPDX-License-Identifier: BSD-3-Clause
from sklearn import model_selection # 导入目标模块 model_selection,获取模块对象引用
from sklearn.model_selection._search_successive_halving import ( # 从内部实现模块导入两个实验性估计器类
HalvingGridSearchCV,
HalvingRandomSearchCV,
)
# 第 7 章 —— use settattr to avoid mypy errors when monkeypatching
setattr(model_selection, "HalvingRandomSearchCV", HalvingRandomSearchCV) # 动态将 HalvingRandomSearchCV 绑定为 model_selection 模块的属性
setattr(model_selection, "HalvingGridSearchCV", HalvingGridSearchCV) # 动态将 HalvingGridSearchCV 绑定为 model_selection 模块的属性
model_selection.__all__ += ["HalvingRandomSearchCV", "HalvingGridSearchCV"] # 原地更新 __all__ 列表,使星号导入也能导出这两个类
7.5.2 逐行解析关键逻辑
让我们拆解这段代码到底做了什么:
-
导入目标模块:
from sklearn import model_selection这一行获取了sklearn.model_selection模块对象的引用,并将其绑定到局部变量model_selection上。这是我们后续要“植入”新功能的地方。 -
导入实现:接下来的两行从内部实现路径
sklearn.model_selection._search_successive_halving导入了实际的估计器类HalvingGridSearchCV和HalvingRandomSearchCV。注意,这些类目前仅存在于这个内部模块中,对普通用户是不可见的。 -
动态注入属性:这里是核心操作。
setattr(model_selection, "HalvingRandomSearchCV", HalvingRandomSearchCV)这行代码的作用是:在model_selection模块对象的属性字典中,设置(或覆盖)一个名为"HalvingRandomSearchCV"的属性,它的值正是我们刚才导入的HalvingRandomSearchCV类。第二行对HalvingGridSearchCV做了同样的事情。效果就是:这两个类如今可以通过model_selection.HalvingRandomSearchCV和model_selection.HalvingGridSearchCV访问到了。 -
更新公开接口:最后一行
model_selection.__all__ += ["HalvingRandomSearchCV", "HalvingGridSearchCV"]则确保了当用户使用from sklearn.model_selection import *时,也能导入这两个新加入的类。__all__列表定义了星号导入 (*) 时哪些名字应该被导出。通过使用+=进行原地修改,我们在不创建新列表的前提下更新了这个列表,这保持了对__all__现有引用的兼容性——如果其他地方已经持有了对这个列表的引用,它们会自动看到更新后的内容。
代码块后总结:这段代码定义了一个显式导入的副作用:通过导入此文件,用户可以在不修改任何源码的前提下,使得两个实验性的超参数搜索估计器在 sklearn.model_selection 模块中可用,同时更新了该模块的公开接口列表以支持星号导入。
你可能会好奇,那个看似多余的 # noqa 注释是做什么的?它其实是给代码检查工具(如 flake8)的一个提示。由于这两个导入语句在文件中后面没有被直接引用(它们的作用纯粹是通过导入动作本身来执行 setattr),静态检查工具会默认将它们标记为“未使用的导入”。而 # noqa 正是告诉检查器:“虽然这行代码看起来没用,但它是有意图的——它的副作用正是我们想要的,请不要报错。”如果你真的把 # noqa 删除,flake8 会给出如下警告:F401 'sklearn.model_selection._search_successive_halving imported but unused'。
7.6 迭代插补启用器 —— 点亮“缺失值填充的进阶技能”
与一次注入两个估计器的 Halving 启用器不同,enable_iterative_imputer.py 展示了另一种常见模式:只启用单个实验性估计器。这个文件的结构非常相似,但更侧重于说明“单一职责”原则——一个启用文件专注于激活一个功能模块。
让我们同样逐行看看它是如何工作的。
7.6.1 核心类型定义
源码路径:sklearn/experimental/enable_iterative_imputer.py - __main__(14-28行)
"""Enables IterativeImputer
The API and results of this estimator might change without any deprecation
cycle.
Importing this file dynamically sets :class:`~sklearn.impute.IterativeImputer`
as an attribute of the impute module::
>>> # explicitly require this experimental feature
>>> from sklearn.experimental import enable_iterative_imputer # noqa
>>> # now you can import normally from impute
>>> from sklearn.impute import IterativeImputer
"""
# 第 7 章 —— Authors: The scikit-learn developers
# 第 7 章 —— SPDX-License-Identifier: BSD-3-Clause
from sklearn import impute # 导入目标模块 impute,获取模块对象引用
from sklearn.impute._iterative import IterativeImputer # 从内部实现模块导入实验性估计器类
# 第 7 章 —— use settattr to avoid mypy errors when monkeypatching
setattr(impute, "IterativeImputer", IterativeImputer) # 动态将 IterativeImputer 绑定为 impute 模块的属性
impute.__all__ += ["IterativeImputer"] # 原地更新 __all__ 列表,使星号导入也能导出该类
7.6.2 逐行解析关键逻辑
这段代码的逻辑与 Halving 启用器几乎是一模一样的,只有一个细微但重要的区别在于导入策略。这里我们没有直接从 sklearn.impute._iterative 导入 IterativeImputer 类然后去设置 sklearn.impute 模块的属性(虽然这样也能工作),而是采用了两步走的方法:
-
首先,
from sklearn import impute导入了整个sklearn.impute模块,并将其绑定到变量impute上。 -
然后,
from sklearn.impute._iterative import IterativeImputer导入了实际的实现类。 -
接下来,
setattr(impute, "IterativeImputer", IterativeImputer)将这个类设置为impute模块(也就是sklearn.impute)的一个属性。 -
最后,
impute.__all__ += ["IterativeImputer"]更新了该模块的公开导出列表。
这样做的好处是:它明确地区分了“导入实现”(来自 _iterative)和“作用于目标模块”(即 impute)。这避免了潜在的命名空间混淆——特别是在复杂的重构过程中,能够确保我们总是在正确的模块对象上进行操作。从用户角度看,效果是完全相同的:导入此文件后,他们就可以通过常规路径 from sklearn.impute import IterativeImputer 获得这个实验性的缺失值填充器。
代码块后总结:这段代码定义了一个显式导入的副作用:通过导入此文件,用户可以使得实验性的 IterativeImputer 估计器在 sklearn.impute 模块中可用,并且该估计器也会被包含在星号导入 (from sklearn.impute import *) 的范围内。
7.7 直方图梯度提升启用器的退役 —— 观察“功能稳定化的墓碑”
到目前为止,我们看到的两个启用器都处于“激活实验性功能”的状态。但软件世界里,实验性功能要么被证明不够好被淘汰,要么经过足够时间的考验后变得稳定并“转正”为正式 API。enable_hist_gradient_boosting.py 正是一个很好的例子,展示了一个功能从实验性走向稳定的完整演变路径——以及 scikit-learn 如何通过保留一个“无操作”(no-op)文件来维护向后兼容性。
让我们看看这个如今只剩下警告的文件。
7.7.1 核心类型定义
源码路径:sklearn/experimental/enable_hist_gradient_boosting.py - __main__(9-23行)
"""This is now a no-op and can be safely removed from your code.
It used to enable the use of
:class:`~sklearn.ensemble.HistGradientBoostingClassifier` and
:class:`~sklearn.ensemble.HistGradientBoostingRegressor` when they were still
:term:`experimental`, but these estimators are now stable and can be imported
normally from `sklearn.ensemble`.
"""
# 第 7 章 —— Authors: The scikit-learn developers
# 第 7 章 —— SPDX-License-Identifier: BSD-3-Clause
# 第 7 章 —— Don't remove this file, we don't want to break users code just because the
# 第 7 章 —— feature isn't experimental anymore.
import warnings # 导入标准库 warnings 模块用于发出警告
warnings.warn( # 在模块导入时立即发出 UserWarning 级别警告
"Since version 1.0, "
"it is not needed to import enable_hist_gradient_boosting anymore. "
"HistGradientBoostingClassifier and HistGradientBoostingRegressor are now "
"stable and can be normally imported from sklearn.ensemble."
)
7.7.2 逐行解析关键逻辑
这段代码已经完全失去了之前两个启用器的“植入”能力。它不再导入任何内部模块,也不再使用 setattr 去修改任何目标模块的属性,也不再更新任何 __all__ 列表。它所做的只有两件事:
-
导入 Python 的标准
warnings模块。 -
在模块被导入时(即此时),立即发出一个
UserWarning级别的警告。
警告信息的内容非常友好和清晰:它告诉用户,“自从 1.0 版本起,你就不再需要导入这个文件了;那些曾经是实验性的梯度提升估计器现在已经很稳定,你可以直接从 sklearn.ensemble 正常导入它们。”换句话说,这个文件现在的唯一目的就是作为一个“墓碑”——它站在那里,不执行任何实际的启用逻辑,但只要用户的旧代码里还有 from sklearn.experimental import enable_hist_gradient_boosting 这行语句,它就会温和地提醒对方:“嘿,你可以安全地删掉这行了,功能已经就在主菜单里了。”
代码块后总结:这段代码实现了一个“安全退役”机制:当一个曾经实验性的功能变得稳定时,其对应的启用器文件被保留但改为仅发出弃用警告,引导用户迁移到稳定的导入路径,而不会立即破坏依赖旧启用器的用户代码。
你可能会问:既然功能已经稳定,为什么不直接删除这个文件?这样岂不是更干净?答案在于对用户迁移成本的考虑。想象一下,如果有成千上万的用户脚本里还有这行导入语句,突然删除这个文件会让他们的代码因为 ModuleNotFoundError 而立刻崩溃。通过保留这个文件并只发出警告,scikit-learn 给了用户一个缓冲期:他们可以在下次维护代码时,看到警告后知道该删掉这行了,而不会因为一个突如其来的错误被迫紧急修复。注释中那句“Don't remove this file”正是对维护者的嘱托:即使看起来这个文件“毫无用处”,它也在为 API 的平滑演化发挥作用。
7.8 设计中的取舍
在实验性功能的设计中,scikit-learn 团队做出了一系列有意识的取舍,以在“鼓励创新”和“保护用户”之间找到平衡。
为什么不用静态导入在 model_selection/__init__.py 中直接写入实验性估计器?如果直接在主模块的 __init__.py 中导入实验性类,那么每一次用户导入 sklearn.model_selection(即便他们根本不打算使用这些实验性功能),都会不必要地加载和初始化这些可能尚不成熟的代码路径。这会增加主模块的启动开销和内存占用。更重要的是,它会让实验性API“无感”地出现在用户面前,使用户难以区分哪些是稳定的承诺,哪些是可能随时变化的尝鲜品。通过显式导入机制,只有当用户真的觉得自己准备好承担实验性API的风险时,才会付出导入启用器的微小成本来获取这些功能。
这种设计的 trade-off(权衡)是什么?主要的权衡在于:实验性功能的可发现性 versus 主命名空间的纯净性和加载效率。与一些将实验性功能直接放在主命名空间(即使被警告包裹)的库不同,scikit-learn 选择了让实验性功能完全“隐形”,除非用户主动去解锁它们。这确实意味着用户需要多知道一步(“我需要去导入那个奇怪的 experimental 模块才能用 Halving搜索”),但换来的是:主导入如 from sklearn.model_selection import * 永远不会不经意地把一个不稳定的东西塞进用户的命名空间;用户的 IDE 自动补全也不会被实验性的、可能命名古怪的类名污染;最重要的是,它为最终的“转正”提供了一个清晰的里程碑——当一个功能准备好时,只需把它从 _search_successive_halving 移到模块的主导出路径中,用户的代码不需要改动(他们已经是从 sklearn.model_selection 导入的),而曾经的启用器文件则可以体面地退役为一个温和的提示器。
7.9 动手练习
-
追踪实验性功能的启用流程
-
在 Python 交互环境中依次执行以下命令,观察每一步的命名空间变化:
import sklearn.model_selection as ms print(hasattr(ms, 'HalvingGridSearchCV')) # False from sklearn.experimental import enable_halving_search_cv print(hasattr(ms, 'HalvingGridSearchCV')) # True print('HalvingGridSearchCV' in ms.__all__) # True -
回答问题:
-
为什么 import 一个看起来未使用的模块会产生副作用?
-
如果去掉
# noqa注释,flake8 会给出什么警告?
-
-
-
为新的实验性估计器编写启用模块
-
假设你需要为 sklearn 编写一个新的实验性估计器
MyExperimentalEstimator,它位于
sklearn/cluster/_my_experimental.py。请写出
sklearn/experimental/enable_my_experimental.py的完整代码,包括:模块文档字符串、导入语句、setattr 注入、all 更新。
注意参考
enable_iterative_imputer.py的单估计器注入模式。
-
-
分析弃用启用器的兼容性策略
-
阅读
enable_hist_gradient_boosting.py源码,思考:-
既然功能已转正,为什么不直接删除这个文件?
-
如果用户代码中仍然存在
from sklearn.experimental import enable_hist_gradient_boosting,会发生什么?警告信息会在什么时机触发?
-
设计一个实验:如何在测试中验证该警告确实被触发?
(提示:使用 pytest.warns(UserWarning, match="..."))
-
-
-
对比三种启用器的实现差异
-
对比
enable_halving_search_cv.py、enable_iterative_imputer.py与enable_hist_gradient_boosting.py三个文件:-
它们分别处于实验性功能生命周期的哪个阶段?
-
enable_hist_gradient_boosting.py缺少了哪些代码元素(如 setattr、all 更新)? -
这种代码演变体现了怎样的 API 设计哲学?
-
-
7.10 本章小结
在这一章中,我们深入探讨了 scikit-learn 如何通过 sklearn.experimental 包及其配套的 enable_* 启用器模块,来管理实验性功能的生命周期。我们首先理解了顶层包的作用在于通过醒目的文档警告,向用户明确传达实验性功能不受弃用周期保护、使用需自担风险的设计意图。然后,我们逐行解析了两种激活实验性估计器的核心机制:通过 setattr 动态地将实现类注入到目标模块(如 model_selection 或 impute)的属性字典中,并同步更新该模块的 __all__ 列表,以确保这些新加入的功能既能通过常规导入路径获得,也能被星号导入 (import *) 捕获。我们还特别关注了工程细节,如 # noqa 注释的用途——它告知静态检查工具那些表面未使用但实际有副作用(执行注入)的导入语句是有意图的。最后,我们观察了一个功能从实验性到稳定的完整演变案例:enable_hist_gradient_boosting.py 已经退化为仅发出 UserWarning 的无操作文件,其存在的目的不是提供功能,而是为用户代码的平滑迁移留出缓冲期,避免因突然删除而破坏旧代码,这充分体现了 scikit-learn 对 API 稳定性和向后兼容性的承诺。
本章我们一起学习了以下概念:
| 概念 | 解释 |
|------|------|
| experimental 包 | 为未稳定特性提供试验田,通过显式导入激活,不受弃用周期约束 |
| setattr 动态注入 | 在运行时将实验性估计器绑定到目标模块的属性字典,避免静态导入开销 |
| all 同步更新 | 原地追加估计器名称,确保 star import 也能导出实验性功能 |
| # noqa 注释 | 告知 linter 该导入语句有副作用,虽然看起来未使用 |
| enable_halving_search_cv | 动态注入 HalvingGridSearchCV 与 HalvingRandomSearchCV 两个估计器 |
| enable_iterative_imputer | 动态注入 IterativeImputer 单个估计器,体现单一职责原则 |
| enable_hist_gradient_boosting | 已退化为 no-op,仅发出弃用警告,保留文件以兼容旧用户代码 |
| 弃用路径设计 | 先保留文件并警告,未来版本再移除,为迁移预留缓冲期 |
下一章中,我们将学习 scikit-learn 的主构建系统与发行商定制机制,理解项目如何从源码编译成可分发的安装包,以及下游发行商如何通过初始化钩子进行自定义。
第 8 章 —— scikit‑learn 概览 —— 走进“机器学习算法百宝箱”
8.1 学习目标
-
难度:★★★☆☆(3/5)
-
预备知识:Python 基础、面向对象编程与 Markdown/代码阅读基础
-
了解 scikit‑learn 的工程结构与核心机制
-
掌握依赖管理、子包 re‑export、克隆语义、元数据路由、文档校验与线程安全等关键实现细节
-
能够快速定位源码并理解每个模块的职责与设计动机
8.2 动手练习
-
在本地克隆 scikit‑learn 仓库,运行
pytest sklearn/tests/test_min_dependencies_readme.py验证 README 与_min_dependencies.py的一致性。 -
修改
sklearn/linear_model/__init__.py中的__all__列表,删除LinearRegression,然后运行python -c "import sklearn.linear_model; print(hasattr(sklearn.linear_model, 'LinearRegression'))"确认公开 API 的变化。 -
实现一个简单的元估计器,继承
MetaEstimatorMixin,并在其fit方法中使用MetadataRouter路由sample_weight到子估计器,验证路由是否生效。 -
使用
numpydoc.validate检查自定义估计器的文档字符串,确保参数说明与函数签名一致。 -
在多线程环境中使用
config_context临时修改assume_finite配置,确认配置在各线程之间互不干扰。
8.3 章节概述
在本章节中,我们把 scikit‑learn 的核心工程结构、质量防线与元数据路由机制逐层拆解,每个解析单元均标注源码路径(src/...),并配以 Mermaid 流程图 说明内部工作流。所有代码块均采用 逐行注释(# ← 注释),随后紧跟解释性文字,帮助读者快速定位每一行的职责。
关键概念表(表前已提供引导语)
以下表格列出了本章将涉及的关键概念及其实现路径。
| 概念 | 代码实现路径 | 关键作用 |
|------|--------------|----------|
| README 徽章矩阵 | README.md (行 1‑35) | CI / 发行状态展示 |
| 最小依赖声明 | README.md (行 36‑46) → sklearn/_min_dependencies.py | 统一依赖管理 |
| 子包 __init__ 重导出 | sklearn/linear_model/__init__.py 等 | 公共 API 统一入口 |
| BaseEstimator.clone 语义 | sklearn/tests/test_base.py: test_clone* | 深拷贝安全 |
| 元估计器路由矩阵 | sklearn/tests/test_metaestimators_metadata_routing.py: METAESTIMATORS | 配置驱动元数据路由 |
| MetadataRouter 与 process_routing | sklearn/utils/metadata_routing.py | 元数据请求/消费框架 |
| numpydoc 文档校验 | sklearn/tests/test_docstring_parameters*.py、sklearn/tests/test_docstrings.py | 文档一致性保证 |
| 线程安全配置 & OpenMP 检浫 | sklearn/tests/test_config.py、sklearn/tests/test_build.py | 运行时与编译时安全 |
8.4 README 与最小依赖的一致性
8.4.1 徽章矩阵(src/README.md)
源码路径:README.md
.. -*- mode: rst -*-
|Azure| |Codecov| |CircleCI| |Nightly wheels| |Ruff| |PythonVersion| ...
# ← CI / 代码质量徽章
# 每个徽章均通过 ReST `.. |Name| image:: URL` 声明
解释:这些徽章在 CI 通过、代码覆盖率、构建状态等维度提供即时可视化,类似药品包装上的质量认证标识,帮助用户快速判断项目健康度。
生活类比:就像在购买食品时,我们会查看包装上的“有机认证”“无糖标签”等图标来快速判断产品质量,scikit‑learn 的徽章矩阵也是如此——它通过一系列自动化检测的结果,让用户一眼看出项目是否健康、构建是否通过、测试覆盖率是否达标。
8.4.2 最小依赖声明(src/README.md 行 36‑46)
源码路径:README.md
.. |PythonMinVersion| replace:: 3.11 # ← Python 最低版本
.. |NumPyMinVersion| replace:: 1.24.1 # ← NumPy 最低版本
.. |SciPyMinVersion| replace:: 1.10.0 # ← SciPy 最低版本
.. |JoblibMinVersion| replace:: 1.3.0 # ← joblib 最低版本
.. |ThreadpoolctlMinVersion| replace:: 3.2.0 # ← threadpoolctl 最低版本
.. |MatplotlibMinVersion| replace:: 3.6.1 # ← Matplotlib 最低版本(可选)
.. |Scikit-ImageMinVersion| replace:: 0.22.0 # ← scikit‑image 最低版本(可选)
.. |PandasMinVersion| replace:: 1.5.0 # ← pandas 最低版本(可选)
.. |SeabornMinVersion| replace:: 0.13.0 # ← seaborn 最低版本(可选)
.. |PytestMinVersion| replace:: 7.1.2 # ← pytest 最低版本(测试用)
.. |PlotlyMinVersion| replace:: 5.18.0 # ← Plotly 最低版本(可选)
解释:使用 Sphinx
replace指令集中维护版本号,test_min_dependencies_readme.py会自动比对sklearn/_min_dependencies.py中的dependent_packages,任何不一致即抛异常,防止文档与实际打包不匹配。
生活类比:这就像食品标签上明确标注的“保质期”或“配料表”——如果标签上写着“保质期至2025年12月”,但实际产品已经过期,消费者会感到被误导。同样,如果 README 声明需要 NumPy 1.24.1,但实际打包时依赖了 1.23.0,用户在安装时可能会遇到意外错误,因而保持一致性至关重要。
8.4.3 一致性检测(src/sklearn/tests/test_min_dependencies_readme.py)
源码路径:sklearn/tests/test_min_dependencies_readme.py
# 第 8 章 —— src: sklearn/tests/test_min_dependencies_readme.py (12‑30行)
pattern = re.compile( # ← 正则匹配 “|PackageMinVersion| replace:: X.Y.Z”
r"\.\. \|"
r"([A-Za-z-]+)"
r"MinVersion\| replace::"
r"( [0-9]+\.[0-9]+(\.[0-9]+)?)"
)
解释:遍历 README 每行,提取包名与版本;随后与
_min_dependencies.py中的声明比较,确保 单一真理源(single source of truth)。
生活类比:想象你在管理一个团队的项目文档,所有成员都必须参考同一份“需求规格说明书”。如果有人根据过时的版本开发功能,而另一个人又根据最新版本修改接口,最终整合时必然会出现冲突。这个测试就是确保所有人都在用同一份规格说明书工作的机制。
8.4.4 流程图:README → 依赖校验
8.5 子包 __init__ 重导出(re‑export)机制
8.5.1 线性模型子包(src/sklearn/linear_model/__init__.py)
源码路径:sklearn/linear_model/__init__.py
"""A variety of linear models.""" # ← 模块文档字符串
# 第 8 章 —— Authors: The scikit-learn developers
# 第 8 章 —— SPDX-License-Identifier: BSD-3-Clause
# 第 8 章 —— 公开 API 通过内部实现模块导入后加入 __all__
from sklearn.linear_model._base import LinearRegression # ← 基础回归模型
from sklearn.linear_model._bayes import ARDRegression, BayesianRidge
from sklearn.linear_model._coordinate_descent import ( # ← Lasso / ElasticNet 系列
ElasticNet, ElasticNetCV, Lasso, LassoCV,
MultiTaskElasticNet, MultiTaskElasticNetCV,
MultiTaskLasso, MultiTaskLassoCV, enet_path, lasso_path,
)
# 第 8 章 —— …(其余私有模块同理)
__all__ = [ # ← 公开符号白名单
"ARDRegression","BayesianRidge","ElasticNet","ElasticNetCV",
"LinearRegression","LogisticRegression","LogisticRegressionCV",
# …(共 70+ 符号)
]
解释:通过
from … import …把实现隐藏在私有子模块,随后使用__all__明确向用户暴露的公共 API。这样做的好处是内部实现可自由重构,而 用户只依赖于稳定的入口。
生活类比:这就像一个餐厅的菜单(
__all__)只列出菜品名称,而厨房的具体做法(私有模块如_base.py、_coordinate_descent.py)对食客是透明的。厨师可以改进炖煮技术或换用新食材,但只要菜品名称和味道保持不变,食客就不会察觉到后厨的变化。
8.5.2 支持向量机子包(src/sklearn/svm/__init__.py)
源码路径:sklearn/svm/__init__.py
"""Support vector machine algorithms."""
from sklearn.svm._bounds import l1_min_c # ← 边界工具
from sklearn.svm._classes import ( # ← 主估计器
SVC, SVR, LinearSVC, LinearSVR, NuSVC, NuSVR, OneClassSVM,
)
__all__ = [ # ← 公开符号
"SVC","SVR","LinearSVC","LinearSVR",
"NuSVC","NuSVR","OneClassSVM","l1_min_c",
]
解释:同样采用 re‑export,保持子包内部实现(如
_bounds、_classes)的封装性。
生活类比:相当于连锁咖啡店只对外公布“拿铁”“卡布奇诺”等饮品名称,而每家店内部的咖啡豆烘焙程度、奶泡打发技术可以因地制宜,只要最终产品符合标准,顾客就不会感知到内部差异。
8.5.3 神经网络子包(src/sklearn/neural_network/__init__.py)
源码路径:sklearn/neural_network/__init__.py
"""Models based on neural networks."""
from sklearn.neural_network._multilayer_perceptron import MLPClassifier, MLPRegressor
from sklearn.neural_network._rbm import BernoulliRBM
__all__ = ["BernoulliRBM","MLPClassifier","MLPRegressor"]
解释:仅 3 个公共类,结构极简,便于快速定位实现代码。
生活类比:这就像一个只有三种口味的冰淇淋店——虽然种类少,但每种口味的配方都清晰可见,顾客可以快速做出选择,而店主也能专注于在这三种味道上做到极致。
8.5.4 流程图:子包导出
8.6 BaseEstimator 与克隆语义(src/sklearn/tests/test_base.py)
8.6.1 测试类与克隆(test_clone*)
源码路径:sklearn/tests/test_base.py
def test_clone(): # ← 基础克隆测试
# 创建 selector, 调用 sklearn.clone → 深拷贝
selector = SelectFpr(f_classif, alpha=0.1) # ← 实例化
new_selector = clone(selector) # ← 调用 clone
assert selector is not new_selector # ← 检查对象不相等
assert selector.get_params() == new_selector.get_params()# ← 参数保持一致
# 对含数组的参数进行克隆
selector = SelectFpr(f_classif, alpha=np.zeros((10, 2)))
new_selector = clone(selector)
assert selector is not new_selector
解释:
clone通过BaseEstimator.get_params读取构造参数,再用同样的参数创建新实例,不复制实例属性(如训练好的权重),确保 无副作用。
生活类比:这就像复制一份食谱卡片——你只复制了原料清单和步骤(参数),但不会复制厨师已经做好的菜肴(属性)。这样,每个人都可以从同一份食谱出发,独立烹饪,互不干扰。
8.6.2 错误路径(test_clone_buggy)
源码路径:sklearn/tests/test_base.py
def test_clone_buggy():
buggy = Buggy() # ← 未正确设置参数的 Estimator
buggy.a = 2 # ← 手动添加属性
with pytest.raises(RuntimeError):
clone(buggy) # ← 触发 “参数不匹配” 错误
varg_est = VargEstimator() # ← 带 *args 的非法构造
with pytest.raises(RuntimeError):
clone(varg_est)
解释:
clone对 不符合 BaseEstimator 约定(如使用*args、未设置self.__dict__)的对象抛出RuntimeError,防止 沉默复制不完整实例。
生活类比:如果你试图复制一个厨师的菜谱,但发现他 secretly 用了“秘密酱料"(未在参数中声明的属性),那么复制出的菜谱可能无法还原原菜的味道。这时直接报错比悄悄生成一个可能失败的复制品更安全。
8.6.3 __sklearn_clone__ 协议(test_clone_protocol)
源码路径:sklearn/tests/test_base.py
def test_clone_protocol():
class FrozenEstimator(BaseEstimator):
def __init__(self, fitted_estimator):
self.fitted_estimator = fitted_estimator
def __sklearn_clone__(self):
return self # ← 返回自身,实现 no‑op
pca = PCA().fit(np.array([[-1, -1], [-2, -1], [-3, -2]]))
frozen = FrozenEstimator(pca)
clone_frozen = clone(frozen)
assert clone_frozen is frozen # ← clone 不再创建新对象
解释:实现
__sklearn_clone__可以 显式声明不可克隆(例如保持共享状态的对象),在 模型持久化或管线复制 时非常有用。
生活类比:想象一个共享的花园灌溉系统——如果你试图“复制”这个系统,实际上你并不需要再建一套管道,因为原系统本来就是设计为共享使用的。此时返回自身(
self)就是最合理的做法,避免不必要的重复建设。
8.6.4 流程图:clone 与 __sklearn_clone__
8.7 元估计器方法委托与数据验证(src/sklearn/tests/test_metaestimators.py)
8.7.1 方法委托检查(test_metaestimator_delegation)
源码路径:sklearn/tests/test_metaestimators.py
def test_metaestimator_delegation():
class SubEstimator(BaseEstimator):
def fit(self, X, y=None): # ← 简单 fit
self.coef_ = np.arange(X.shape[1])
self.classes_ = []
return True
@hides # ← 隐藏特定方法
def predict(self, X):
self._check_fit()
return np.ones(X.shape[0])
# …其余方法同理
for delegator_data in DELEGATING_METAESTIMATORS:
delegator = delegator_data.construct(SubEstimator())
# 未 fit 前,各方法均不可调用 → NotFittedError
for method in methods:
with pytest.raises(NotFittedError):
getattr(delegator, method)(X)
# fit 后,应当委托到子估计器
delegator.fit(X, y)
for method in methods:
getattr(delegator, method)(X) # ← 通过子估计器执行
解释:元估计器(如
Pipeline、GridSearchCV)仅在子估计器实现对应方法时才暴露;否则保持 API 清晰,避免用户调用不存在的属性。
生活类比:这就像一个智能家居控制面板——如果你家没有安装空调,面板上就不会出现“调节温度”的按钮;只有当空调被安装并连接后,该按钮才会出现并可用。这样可以避免用户误以为设备具备某项功能而导致困惑。
8.7.2 数据验证委托(test_meta_estimators_delegate_data_validation)
源码路径:sklearn/tests/test_metaestimators.py
def test_meta_estimators_delegate_data_validation(estimator):
estimator = clone(estimator) # ← 确保线程安全
X = np.random.choice(np.array(["aa","bb","cc"]), size=30)
y = np.random.randn(30) if is_regressor(estimator) else np.random.randint(3, size=30)
estimator.fit(X, y) # ← 子估计器负责 data validation
assert not hasattr(estimator, "n_features_in_") # ← 元估计器不自行记录特征数
解释:元估计器不直接调用
check_array,而是 把验证责任下放 给子估计器,避免 重复检查 与 特征维度冲突。
生活类比:相当于总经理不自己审查每份报告的格式和内容,而是要求各部门负责人在提交前自行确保报告符合模板。这样既避免了总经理成为瓶颈,又确保了每份报告都经过了专业把关。
8.7.3 流程图:元估计器委托
8.8 元数据路由矩阵与工具层(src/sklearn/tests/test_metaestimators_metadata_routing.py、metadata_routing_common.py、test_metadata_routing.py)
8.8.1 METAESTIMATORS 配置矩阵(摘录)
源码路径:sklearn/tests/test_metaestimators_metadata_routing.py
METAESTIMATORS = [
{
"metaestimator": MultiOutputRegressor, # ← 元回归器
"estimator_name": "estimator",
"estimator": "regressor",
"X": X, "y": y_multi,
"estimator_routing_methods": ["fit", "partial_fit"],
},
{
"metaestimator": BaggingClassifier, # ← 案例:不转发元数据
"estimator_name": "estimator",
"estimator": "classifier",
"preserves_metadata": False,
"estimator_routing_methods": [
("fit", ["metadata"]), "predict", "predict_proba"
],
"method_mapping": {"predict": ["predict","predict_proba"]},
},
# …共 40+ 条配置
]
解释:每条记录描述 元估计器、子估计器类型、需要路由的元数据(
sample_weight、metadata)以及是否 保留原始元数据。测试会自动遍历该矩阵,检查 请求、消费、冲突 的完整性。
生活类比:这就像在一个快递中心工作,每个包裹都有不同的“特殊标识”(如易碎、冷链、加急)。系统需要根据这些标识自动将包裹送到对应的处理线路——易碎品去缓冲区,冷链品去冷藏仓,加急品走快速通道。如果标识填错或处理线路不匹配,包裹可能会被误送或损坏。
8.8.2 路由请求生成与冲突检测(filter_metadata_in_routing_methods 与 set_requests)
源码路径:sklearn/tests/test_metaestimators_metadata_routing.py
def filter_metadata_in_routing_methods(estimator_routing_methods):
# 将 ["fit", ("fit", ["sample_weight"])] 转换为统一 dict
res = {}
for spec in estimator_routing_methods:
if isinstance(spec, str):
method, metadata = spec, ["sample_weight","metadata"]
else:
method, metadata = spec
res[method] = metadata
return res # ← 统一格式
def set_requests(obj, *, method_mapping, methods, metadata_name, value=True):
# 为子估计器的每个 method 调用 set_{callee}_request
for caller in methods:
for callee in method_mapping.get(caller, [caller]):
getattr(obj, f"set_{callee}_request")(**{metadata_name: value})
解释:显式请求是必需的;若用户在
fit时传递sample_weight,但子估计器未声明需求,则会抛出UnsetMetadataPassedError,防止 隐式沉默。
生活类比:想象你在酒店前台登记入住时,如果你悄悄把一只宠物带进房间但没告知前台,而酒店政策明确要求必须申报并可能收取额外费用,那么你可能会因违规被罚款或被要求退房。相比之下,如果你主动申报并获得批准,整个过程就会透明且合规。
8.8.3 MetadataRouter 与 MetadataRequest(metadata_routing_common.py)
源码路径:sklearn/tests/metadata_routing_common.py
class MetadataRequest:
def __init__(self, owner):
self.owner = owner
self.fit = MethodMetadataRequest(owner, "fit") # ← 每个 method 的请求容器
# …同理定义 partial_fit、predict、score 等
class MetadataRouter:
def __init__(self, owner):
self.owner = owner
self._self_request = None
def add_self_request(self, obj):
# 将自身请求拷贝进去,保证后续修改不影响原对象
self._self_request = deepcopy(get_routing_for_object(obj))
return self
def add(self, **kwargs):
# 为每个子对象注册路由信息(对象、方法映射)
for name, estimator in kwargs.items():
# …内部维护 MethodMapping 与路由请求
return self

浙公网安备 33010602011771号