Sklearn-源码解析-书-v1-0-二十-

Sklearn 源码解析(书)v1.0(二十)

__all__ 明确列出了 当用户执行 from sklearn.model_selection import * 时应当导出的符号,从而隐藏内部实现细节。即便 Halving* 属于实验特性,它们也被列入 __all__,因为 __getattr__ 会在实际访问时提供引导信息,保持文档与代码的一致性。

42.3.5 动态属性访问钩子(实验特性守卫)

def __getattr__(name):
    if name in {"HalvingGridSearchCV", "HalvingRandomSearchCV"}:
        raise ImportError(
            f"{name} is experimental and the API might change without any "
            "deprecation cycle. To use it, you need to explicitly import "
            "enable_halving_search_cv:\n"
            "from sklearn.experimental import enable_halving_search_cv"
        )
    raise AttributeError(f"module {__name__} has no attribute {name}")
  • 触发时机:当用户在运行时访问模块属性且该属性未在全局字典中时,Python 会调用模块级 __getattr__

  • 实验特性拦截:如果属性名是 HalvingGridSearchCVHalvingRandomSearchCV,函数抛出带有明确使用说明的 ImportError,强制用户通过 sklearn.experimental.enable_halving_search_cv 显式开启实验特性。

  • 其他属性:不在实验名单中的缺失属性会触发普通的 AttributeError

这相当于服务台的保安,只允许持有“邀请函”(即显式启用)的访客进入实验区域。

42.3.5.1 体系结构图(文字示意)

+------------------------+
| sklearn.model_selection|
+----------+-------------+
           |
   +-------v-------+      (导入阶段)
   |  Core APIs    |---+  FixedThresholdClassifier、GridSearchCV、KFold …
   +---------------+   |
           |          |
   +-------v-------+  |   (TYPE_CHECKING)
   |  Experimental |<-+  Halving*  (仅在类型检查时可见)
   +---------------+
           |
   +-------v-------+   (运行时访问)
   | __getattr__   |---+---> 若访问 Halving* → ImportError + 使用指引
   +---------------+

42.4 测试公共设施:OneTimeSplitter —— 单次切分的“一次性试管”

42.4.1 核心类型定义:OneTimeSplitter

class OneTimeSplitter:
    """A wrapper to make KFold single entry cv iterator"""

    def __init__(self, n_splits=4, n_samples=99):
        self.n_splits = n_splits
        self.n_samples = n_samples
        # 预先创建 KFold 并一次性获取分割生成器
        self.indices = iter(KFold(n_splits=n_splits).split(np.ones(n_samples)))
  • 目的:在单元测试中模拟只能遍历一次的交叉验证器,帮助检测高层估计器(如 GridSearchCV)是否会错误地多次调用 split

  • 实现要点:构造函数中立刻调用 KFold.split 并把返回的生成器包装为一次性迭代器 self.indices,确保后续 split 只能消费一次。

42.4.2 split 方法逐行解析

def split(self, X=None, y=None, groups=None):
    """Split can be called only once"""
    for index in self.indices:
        yield index
  • 循环:遍历预先存储的 self.indices。因为 self.indices 是一次性迭代器,第一次遍历会产生全部 (train_index, test_index) 元组。

  • 一次性特性:第二次调用 splitself.indices 已经耗尽,循环直接结束,不会再产生任何划分。

42.4.3 get_n_splits 方法逐行解析

def get_n_splits(self, X=None, y=None, groups=None):
    return self.n_splits
  • 返回值:保持与 BaseCrossValidator 接口一致,返回在初始化时声明的折数。即便实际只能遍历一次,外部仍能通过该方法获知预期的划分数量。

42.4.3.1 流程图(一次性迭代)

sequenceDiagram participant Tester participant Splitter as OneTimeSplitter participant Gen as KFold.split() generator Tester->>Splitter: instantiate (n_splits=3) Splitter->>Gen: KFold(3).split(dummy) Gen-->>Splitter: generator stored in self.indices Tester->>Splitter: first split() Splitter->>Gen: iterate self.indices Gen-->>Splitter: yield (train0, test0) Splitter->>Tester: (train0, test0) ... (repeat for remaining folds) ... Tester->>Splitter: second split() Splitter->>Gen: iterate exhausted generator Gen-->>Splitter: no more items Splitter->>Tester: empty iterator

42.5 设计中的取舍

在实现 公共 API 与实验特性的平衡 时,需要在可见性、使用便利和向后兼容之间做出权衡。

在进入表格前先给出概述:

下表对三种常见的实验特性发布策略进行比较,帮助读者理解为何 Scikit‑Learn 选择了当前的 “默认隐藏 + getattr 报错” 方案。

| 方案 | 优点 | 缺点 |

|------|------|------|

| 默认隐藏 + __getattr__ 报错(当前实现) | - 防止用户在生产环境误用未成熟 API。
- 错误信息明确指向官方启用方式。
- 仅在真正需要时才加载实验代码,保持包体积小。 | - 需要记住额外的导入语句;IDE 自动补全不显示实验类。 |

| __all__ 中直接暴露 | - 开发体验最佳,用户可直接 import 实验类。 | - 实验类一旦改动会直接破坏已有代码;缺少显式的风险提示。 |

| 通过子模块 experimental 完全隔离 | - 完全解耦,实验代码与正式代码分离。 | - 增加学习成本,需要记住不同的导入路径;文档需额外维护两套入口。 |

42.5.1 延迟加载的成本分析

  • 内存占用:只有在用户显式启用实验特性时才会把相关模块加载进内存,降低默认启动时的资源消耗。

  • 启动时间__init__.py 只执行核心导入,保持 sklearn 包的快速导入。

  • 维护便利:实验代码可以独立迭代,后期如果成熟只需在 __init__ 中去掉 __getattr__ 检查即可,实现平滑迁移。


42.6 动手练习

下面的练习旨在让读者亲自验证本章节的关键实现。每一项都以完整的段落形式呈现,便于直接复制运行。

练习 1 – 探索 __all__

通过打印 __all__,观察哪些符号被公开。确认实验特性 Halving* 已列入列表,但在未显式启用前访问会抛出错误。

import sklearn.model_selection as ms

print("公开的符号列表:")
print(ms.__all__)

练习 2 – 触发 __getattr__

直接导入实验类会触发 __getattr__,观察它抛出的 ImportError 提示信息,从而学习正确的启用方式。

try:
    from sklearn.model_selection import HalvingGridSearchCV
except ImportError as e:
    print("捕获到的错误信息:")
    print(e)

练习 3 – 使用 OneTimeSplitter

实例化 OneTimeSplitter,遍历一次划分后再次调用 split,验证第二次不会产生任何输出。

from sklearn.model_selection.tests.common import OneTimeSplitter

splitter = OneTimeSplitter(n_splits=3, n_samples=30)

print("第一次遍历划分:")
for train_idx, test_idx in splitter.split():
    print(f"train: {len(train_idx)}, test: {len(test_idx)}")

print("\n第二次遍历划分(应为空):")
for train_idx, test_idx in splitter.split():
    print("这行不会被执行")

练习 4 – 实现简易延迟加载

创建一个包装包 my_pkg/__init__.py,利用 __getattr__ 在访问属性 pd 时才真正导入 pandas。运行以下代码验证只在首次访问时才加载。

# 第 42 章 —— 假设已经在 my_pkg/__init__.py 中写入如下代码:
# 第 42 章 —— def __getattr__(name):
# 第 42 章 —— if name == "pd":
# 第 42 章 —— import pandas as pd
# 第 42 章 —— globals()["pd"] = pd
# 第 42 章 —— return pd
# 第 42 章 —— raise AttributeError(name)

import importlib, sys, time

start = time.time()
import my_pkg          # 不会导入 pandas
print("导入 my_pkg 用时:", time.time() - start, "秒")

start = time.time()
df = my_pkg.pd.DataFrame({"a": [1, 2, 3]})   # 第一次访问 pd,触发导入
print("首次访问 pd 用时:", time.time() - start, "秒")
print(df)

start = time.time()
_ = my_pkg.pd.DataFrame({"b": [4, 5]})       # 第二次访问已缓存
print("二次访问 pd 用时:", time.time() - start, "秒")

通过对比首次与二次访问的耗时,你可以直观看到延迟加载的效果。


42.7 本章小结

本章系统梳理了 sklearn.model_selection 作为模型选择工具总服务台的内部实现细节:

  1. 统一入口__init__.py 通过显式导入把切分、搜索、阈值调优、可视化等功能聚合为统一命名空间。

  2. 实验特性管控typing.TYPE_CHECKING 为类型检查提供实验类声明,__getattr__ 在运行时拦截访问并抛出指引性的 ImportError,实现“默认隐藏、显式启用”。

  3. 源码细化:对每个核心类(FixedThresholdClassifierLearningCurveDisplayGridSearchCVKFold 等)给出导入路径、职责概述以及关键接口。

  4. __all__ 解析:逐项列出 30+ 公共符号,解释即使是实验特性也被列入以保持文档一致性。

  5. 测试设施OneTimeSplitter 通过一次性 KFold 生成器模拟只能遍历一次的流式数据场景,为上层估计器的稳健性提供专用夹具。

  6. 设计取舍:通过类比阐明实验特性可见性与使用便利之间的平衡,并展示延迟加载带来的资源与维护优势。

类比回顾:正如总服务台既要提供丰富的常规业务,也要对实验性专属房间设立严格的访问规则,sklearn.model_selection 通过上述机制实现了功能完整、风险可控、用户友好的整体设计。

42.8 生活类比

想象 sklearn.model_selection 是一家大型机器学习服务公司的“总服务台”__init__.py = 前台导览图 & 智能分诊台: 清晰列出所有可直接办理的业务(__all__ 中的 30+ 个类/函数) 将业务按类型分区:切分器(_split)、搜索器(_search)、验证器(_validation)、阈值调优(_classification_threshold)、可视化(_plot) __getattr__ 机制 = “实验性业务”专用通道守卫: 普通客户(用户)直接询问 HalvingGridSearchCV 时,守卫拦截并告知:

“该业务为内测阶段,需凭‘邀请函’(from sklearn.experimental import enable_halving_search_cv)方可办理” 类型检查员(mypy)来视察时,守卫出示“规划蓝图”(TYPE_CHECKING 分支),确保蓝图完整但不实际开放业务 OneTimeSplitter = “一次性体检套餐”试管: 标准体检(KFold)可反复预约多次,但某些特殊测试(流式验证)要求只能体检一次 该试管预装好一次性采样流程(预存 indices 生成器),用完即废,精准模拟不可重置的数据流场景

42.9 源码地图:模块结构概览

sklearn/model_selection/__init__.py
├── 导入与重导出(1-42行)
│   ├── from _classification_threshold import FixedThresholdClassifier, TunedThresholdClassifierCV
│   ├── from _plot import LearningCurveDisplay, ValidationCurveDisplay
│   ├── from _search import GridSearchCV, ParameterGrid, ParameterSampler, RandomizedSearchCV
│   ├── from _split import (BaseCrossValidator, KFold, StratifiedKFold, GroupKFold,
│   │   LeaveOneOut, LeavePOut, LeaveOneGroupOut, LeavePGroupsOut,
│   │   ShuffleSplit, StratifiedShuffleSplit, GroupShuffleSplit,
│   │   RepeatedKFold, RepeatedStratifiedKFold, TimeSeriesSplit,
│   │   PredefinedSplit, check_cv, train_test_split)
│   └── from _validation import (cross_validate, cross_val_score, cross_val_predict,
│       learning_curve, validation_curve, permutation_test_score)
├── 类型检查分支(44-50行)
│   └── if typing.TYPE_CHECKING: import HalvingGridSearchCV, HalvingRandomSearchCV
├── 公共 API 声明(52-78行)
│   └── __all__ = [... 30+ 个公共符号 ...]
└── 动态属性访问钩子(80-90行)
    └── __getattr__(name) # 实验性特性访问拦截与报错引导

sklearn/model_selection/tests/common.py
└── OneTimeSplitter 类
    ├── __init__(n_splits=4, n_samples=99) # 预先实例化 KFold 生成器
    ├── split(X=None, y=None, groups=None) # 单次迭代协议,yield 预存索引
    └── get_n_splits(X=None, y=None, groups=None) # 委托内部 n_splits 属性

第 43 章 —— model_selection 公共 API —— 模型选择工具的“总服务台”

43.1 学习目标

  • 难度:★★★☆☆(3/5)

  • 预备知识:Python 基础、面向对象编程与 Markdown/代码阅读基础

  • 理解 model_selection 公共 API 的组织结构与延迟加载机制

  • 掌握测试公共设施 OneTimeSplitter 的实现与用途

  • 理解阈值分类器(FixedThresholdClassifier 与 TunedThresholdClassifierCV)的核心逻辑与元数据路由集成

  • 掌握阈值调优中交叉验证评分、最佳阈值选择与 refit 行为的验证要点

  • 能够阅读并扩展模型选择模块的公共测试用例

  • 掌握 _fit_and_score_over_thresholds 核心执行引擎的单次拟合多阈值评分逻辑

  • 理解 _CurveScorer 在阈值空间扫描评分的实现原理

43.2 生活类比

想象 model_selection 模块是一座 机器学习模型调优的指挥中心

  • __init__.py 就是大门,统一对外暴露切分器、验证器、搜索器、阈值调优器等“作战单元”,并通过 __getattr__ 实现实验性功能的“动态部署”。

  • tests/common.py 充当后勤补给站,OneTimeSplitter 确保交叉验证切分器在测试中只被调用一次,防止重复消耗。

  • test_classification_threshold.py 是阈值调优的专项演练场,内部包括:

    • FixedThresholdClassifier —— 固定阈值的“标准哨所”,直接将连续分数按阈值转为类别;

    • TunedThresholdClassifierCV —— 智能阈值搜索的“雷达站”,利用 _CurveScorer 在阈值空间扫描,结合交叉验证寻找最佳决策边界;

    • _fit_and_score_over_thresholds —— 核心执行引擎,在各折训练/验证并跨阈值评分,支持样本权重、fit_params 等元数据精准路由;

    • _CurveScorer —— 多阈值评分的“扫描仪”,在概率/决策函数输出上按阈值序列计算指标。

  • 元数据路由 充当战场通讯网络:sample_weightgroupsfit_params 等情报通过 set_*_request 精准分派给 estimator/scorer/splitter,enable_metadata_routing=True 时生效。

正如指挥中心需要统一调度各单元协同作战,model_selection 通过统一 API、延迟加载、元数据路由,确保切分、验证、搜索、阈值调优各环节高效联动、配置灵活。

43.3 源码地图

sklearn/model_selection/__init__.py
├── 公共 API 组装与延迟加载
│   ├── __all__ 列表定义
│   ├── 从 _split、_validation、_search、_search_successive_halving、_classification_threshold 导入核心类
│   └── __getattr__ 实现实验性组件延迟加载
sklearn/model_selection/tests/common.py
├── OneTimeSplitter.__init__()              # 单次迭代 CV 包装器初始化
├── OneTimeSplitter.split()                 # 仅产出一次切分
├── OneTimeSplitter.get_n_splits()          # 返回切分数
sklearn/model_selection/tests/test_classification_threshold.py
├── _fit_and_score_over_thresholds 验证
│   ├── test_fit_and_score_over_thresholds_curve_scorers()       # 曲线评分器阈值有序性
│   ├── test_fit_and_score_over_thresholds_prefit()              # 预拟合估计器行为
│   ├── test_fit_and_score_over_thresholds_sample_weight()       # 样本权重路由等价性
│   ├── test_fit_and_score_over_thresholds_fit_params()          # fit_params 透传
├── TunedThresholdClassifierCV 验证
│   ├── test_tuned_threshold_classifier_no_binary()              # 非二分类报错
│   ├── test_tuned_threshold_classifier_conflict_cv_refit()    # cv/refit 冲突检查
│   ├── test_threshold_classifier_estimator_response_methods() # 响应方法暴露
│   ├── test_tuned_threshold_classifier_without_constraint_value() # 优化目标指标
│   ├── test_tuned_threshold_classifier_metric_with_parameter()  # 带参数评分器
│   ├── test_tuned_threshold_classifier_with_string_targets()  # 字符串标签支持
│   ├── test_tuned_threshold_classifier_refit()                # refit 行为
│   ├── test_tuned_threshold_classifier_fit_params()           # fit_params 透传
│   ├── test_tuned_threshold_classifier_cv_zeros_sample_weights_equivalence() # 零权重等价性
│   ├── test_tuned_threshold_classifier_thresholds_array()     # 自定阈值数组
│   ├── test_tuned_threshold_classifier_store_cv_results()      # cv_results_ 存储控制
│   ├── test_tuned_threshold_classifier_cv_float()              # 浮点数 cv 参数
│   ├── test_tuned_threshold_classifier_error_constant_predictor() # 常数预测器报错
├── FixedThresholdClassifier 验证
│   ├── test_fixed_threshold_classifier_equivalence_default()    # 默认阈值等价性
│   ├── test_fixed_threshold_classifier()                        # 自定义阈值与 pos_label
│   ├── test_fixed_threshold_classifier_metadata_routing()       # 元数据路由集成
│   ├── test_fixed_threshold_classifier_fitted_estimator()       # 预拟合估计器
│   ├── test_fixed_threshold_classifier_classes_()               # classes_ 属性
sklearn/model_selection/_classification_threshold.py
├── _fit_and_score_over_thresholds()    # 核心函数:单次 fit 后多阈值评分
├        └── 解释:此函数实现单次模型拟合、多阈值评分,核心是避免重复 fit,提升效率。其流程为:若提供 train_idx 按索引分割数据并切片 fit_params/score_params,调用一次 classifier.fit;若 train_idx 为 None(预拟合),则直接使用已有模型;最后交给 curve_scorer 在验证集上按阈值序列评分,返回 scores 与 thresholds。
├── _CurveScorer                       # 曲线评分器:阈值扫描评分
├        └── 解释:此类封装了在多个决策阈值上评估指标的逻辑。它接受一个 scorer(如 balanced_accuracy_score)、响应方法(predict_proba/decision_function)和阈值序列,返回每个阈值对应的得分。内部通过 _get_response_values_binary 获取估计器的连续输出,再按阈值转换为类别标签,最后调用底层 scorer 计算指标。它是 TunedThresholdClassifierCV 能够在阈值空间“扫描”的关键。
├── FixedThresholdClassifier           # 固定阈值分类器
│   ├── __init__()
│   │   └── 解释:构造函数接受 estimator、threshold(“auto”或 float)、pos_label 和 response_method(“auto”、"predict_proba"、"decision_function")。它将这些参数存储为实例属性,为后续的 predict 决策做准备。
│   ├── fit()
│   │   └── 解释:通过 process_routing 处理元数据路由(如 sample_weight),克隆底层 estimator 并调用 fit 方法。此步骤确保在 enable_metadata_routing=True 时,fit_params 中的 sample_weight 等信息能正确传递给底层 estimator。
│   ├── predict()
│   │   └── 解释:检查是否已拟合,获取 estimator_(或直接使用 estimator),调用 _get_response_values_binary 获取连续得分(概率或决策函数)及实际使用的响应方法。根据 threshold="auto" 自动设定阈值(概率用 0.5,决策函数用 0.0)或使用自定义阈值,最后调用 _threshold_scores_to_class_labels 将得分映射为类别标签。
│   ├── predict_proba()
│   │   └── 解释:代理到底层 estimator_.predict_proba(X),前提是 estimator 已拟合(通过 _check_is_fitted 检查)。
│   ├── decision_function()
│   │   └── 解释:代理到底层 estimator_.decision_function(X),同样需确保 estimator 已拟合。
│   └── classes_ 属性
│       └── 解释:返回 self.estimator_.classes_(若已拟合)或尝试检查底层 estimator 是否已拟合;若未拟合则抛出 AttributeError。此属性确保类别标签的一致性,便于上游下游使用。
├── TunedThresholdClassifierCV         # 交叉验证阈值调优器
│   ├── __init__()
│   │   └── 解释:构造函数接受 estimator、scoring(“balanced_accuracy”等)、response_method(“auto”等)、thresholds(int 或 array-like)、cv(整数、浮点、splitter 或 "prefit")、refit、n_jobs、random_state 和 store_cv_results 等参数,并存傚为实例属性。
│   ├── fit()
│   │   └── 解释:核心逻辑包括:① 解析 cv 参数(浮点触发 StratifiedShuffleSplit,"prefit" 检查 refit 冲突,否则用 check_cv 生成 splitter);② 通过 process_routing 整合 fit 与 score 的元数据需求;③ 得到最终用于预测的 estimator_(全数据 fit 或单折 fit)和用于阈值搜索的 classifier;④ 并行遍历每折,调用 _fit_and_score_over_thresholds 实现“一次拟合、多阈值评分”;⑤ 检查常数预测器(若所有折的阈值范围退化为常数则报错);⑥ 统一阈值网格(线性或自定义)并在每折上做插值;⑦ 计算平均得分并选出 best_threshold_ 和 best_score_;⑧ 可选地存储 cv_results_。
│   ├── predict()
│   │   └── 解释:检查是否已拟合,从 _curve_scorer 获取 pos_label,调用 _get_response_values_binary 获取 estimator_ 的连续得分,最后用 best_threshold_ 将得分转换为类别标签。
│   ├── predict_proba()
│   │   └── 解释:代理到底层 estimator_.predict_proba(X)。
│   ├── decision_function()
│   │   └── 解释:代理到底层 estimator_.decision_function(X)。
│   ├── best_threshold_ 属性
│   │   └── 解释:在 fit 过程中计算得到的最优决策阈值(浮点数)。
│   └── cv_results_ 属性
│       └── 解释:若 store_cv_results=True,则为包含 "thresholds"(阈值网格)和 "scores"(对应平均得分)的字典;否则不存在。

43.4 模型选择公共 API 与延迟加载 —— 统一入口的“总服务台”

模型选择模块的 __init__.py 作为统一入口,承担两个关键职责:

1️⃣ 通过 __all__ 明确对外暴露的公共接口;

2️⃣ 通过 __getattr__ 实现实验性功能的延迟加载与向后兼容。

43.4.1 公共 API 组装

源码路径sklearn/model_selection/__init__.py - __all__(第 17‑65 行)

# 第 43 章 —— 模块文档字符串,说明此文件的用途
"""Tools for model selection, such as cross validation and hyper-parameter tuning."""

# 第 43 章 —— 版权信息和许可证声明
# 第 43 章 —— Authors: The scikit-learn developers
# 第 43 章 —— SPDX-License-Identifier: BSD-3-Clause

import typing  # 导入 typing 模块,用于类型检查

# 第 43 章 —— 从子模块导入核心类和函数,构建公共 API
from sklearn.model_selection._classification_threshold import (
    FixedThresholdClassifier,
    TunedThresholdClassifierCV,
)
# 第 43 章 —— …(省略其他导入)

# 第 43 章 —— 条件导入,仅在类型检查阶段出现
if typing.TYPE_CHECKING:
    # 避免类型检查器(如 mypy)对实验性估计器报错
    from sklearn.model_selection._search_successive_halving import (
        HalvingGridSearchCV,
        HalvingRandomSearchCV,
    )

# 第 43 章 —— 定义 __all__ 列表,明确对外暴露的公共接口
__all__ = [
    "BaseCrossValidator",
    "BaseShuffleSplit",
    "FixedThresholdClassifier",
    "GridSearchCV",
    # …(其余符号)
    "TunedThresholdClassifierCV",
    "check_cv",
    "cross_val_predict",
    # …
]

核心要点:统一从子模块(_classification_threshold_plot_search_split_validation)导入关键类/函数,确保用户只能访问经过审查的符号。typing.TYPE_CHECKING 为实验性类提供类型信息,避免运行时加载。__all__ 列表实现封装与清晰的 API。

43.4.2 延迟加载机制

源码路径sklearn/model_selection/__init__.py - __getattr__(第 68‑78 行)

def __getattr__(name):
    if name in {"HalvingGridSearchCV", "HalvingRandomSearchCV"}:
        raise ImportError(
            f"{name} is experimental and the API might change without any "
            "deprecation cycle. To use it, you need to explicitly import "
            "enable_halving_search_cv:\n"
            "from sklearn.experimental import enable_halving_search_cv"
        )
    raise AttributeError(f"module {__name__} has no attribute {name}")

实现思路:当访问未知属性时,Python 自动调用 __getattr__。若属性属于实验性功能,抛出 ImportError 并指明启用方式;否则抛出标准 AttributeError。该机制避免在模块加载时导入实验性代码,提升启动性能,同时提供明确迁移路径。

43.4.3 架构图(整体流程)

flowchart TD subgraph API入口 A[__init__.py] -->|导入| B[_split、_validation、_search、_classification_threshold] A -->|定义| C[__all__] A -->|实现| D[__getattr__] end subgraph 延迟加载 D -->|实验性| E[HalvingGridSearchCV, HalingRandomSearchCV] end

43.5 测试公共设施 OneTimeSplitter —— 确保 CV 切分器“一次性使用”

在模型选择的测试体系中,需要验证交叉验证函数是否只调用一次切分器。OneTimeSplitter 将标准 KFold 包装为只能产生一次切分的迭代器。

43.5.1 单次迭代包装器

源码路径sklearn/model_selection/tests/common.py - OneTimeSplitter(第 1‑23 行)

class OneTimeSplitter:
    """A wrapper to make KFold single entry cv iterator"""

    def __init__(self, n_splits=4, n_samples=99):
        self.n_splits = n_splits
        self.n_samples = n_samples
        # 预先创建一个 KFold 迭代器,只保留其迭代器对象
        self.indices = iter(KFold(n_splits=n_splits).split(np.ones(n_samples)))

    def split(self, X=None, y=None, groups=None):
        """Split can be called only once"""
        for index in self.indices:
            yield index

    def get_n_splits(self, X=None, y=None, groups=None):
        return self.n_splits

工作原理

__init__ 预先生成一次 KFold.split 的迭代器并保存。split 只遍历该迭代器一次,后续调用返回空迭代器,确保只产生一次真实切分。get_n_splits 返回初始化时的折数,保持与 KFold 接口兼容。

43.5.2 架构图(OneTimeSplitter 工作流)

sequenceDiagram participant Test as 测试代码 participant Splitter as OneTimeSplitter participant KFold as KFold Test->>Splitter: 初始化 Splitter->>KFold: 创建 KFold 并获取 split 迭代器 loop Split 调用 Splitter->>Splitter: 读取一次迭代器 end Note right of Splitter: 之后再调用返回空

43.6 阈值调优核心执行引擎 —— _fit_and_score_over_thresholds 的单次拟合多阈值评分

阈值调优需要在 多个阈值 上评估模型性能。如果为每个阈值重新拟合模型,计算成本会呈指数增长。_fit_and_score_over_thresholds 通过 一次拟合多阈值评分 的策略,大幅提升效率。

43.6.1 核心逻辑

源码路径sklearn/model_selection/_classification_threshold.py - _fit_and_score_over_thresholds(第 140‑150 行)

def _fit_and_score_over_thresholds(
    classifier,
    X,
    y,
    *,
    fit_params,
    train_idx,
    val_idx,
    curve_scorer,
    score_params,
):
    """Fit a classifier and compute the scores for different decision thresholds."""
    if train_idx is not None:
        X_train, X_val = _safe_indexing(X, train_idx), _safe_indexing(X, val_idx)
        y_train, y_val = _safe_indexing(y, train_idx), _safe_indexing(y, val_idx)
        fit_params_train = _check_method_params(X, fit_params, indices=train_idx)
        score_params_val = _check_method_params(X, score_params, indices=val_idx)
        classifier.fit(X_train, y_train, **fit_params_train)
    else:  # prefit estimator, only a validation set is provided
        X_val, y_val, score_params_val = X, y, score_params

    return curve_scorer(classifier, X_val, y_val, **score_params_val)

流程要点:若提供 train_idx,使用 _safe_indexing 按索引分割训练/验证集。通过 _check_method_paramsfit_params / score_params 按索引切片(如 sample_weight)。只调用一次 classifier.fit()(或直接使用预拟合模模型)。将已拟合的分类器、验证数据以及评分参数交给 curve_scorer,一次性返回所有阈值的 scoresthresholds

43.6.2 架构图(_fit_and_score_over_thresholds 流程)

flowchart TD A[输入: classifier, X, y, train_idx, val_idx, fit_params, score_params] A --> B{train_idx 是否为 None?} B -- 否 --> C[划分训练/验证集] C --> D[切片 fit_params / score_params] D --> E[一次拟合 classifier.fit()] B -- 是 --> F[使用预拟合 classifier] E --> G[调用 curve_scorer] F --> G G --> H[返回 scores, thresholds]

43.7 FixedThresholdClassifier —— 固定阈值的“标准哨所”

在二分类任务中,模型常输出概率或决策分数,阈值用于将其映射为离散标签。FixedThresholdClassifier 提供 手动设定 的阈值,且不需要重新训练底层模型。

43.7.1 阈值决策逻辑

源码路径sklearn/model_selection/_classification_threshold.py - FixedThresholdClassifier(第 1‑200 行)

class FixedThresholdClassifier(BaseThresholdClassifier):
    """Binary classifier that manually sets the decision threshold."""
    def __init__(self, estimator, *, threshold="auto", pos_label=None, response_method="auto"):
        super().__init__(estimator=estimator, response_method=response_method)
        self.pos_label = pos_label
        self.threshold = threshold

    def predict(self, X):
        _check_is_fitted(self)
        estimator = getattr(self, "estimator_", self.estimator)
        y_score, _, response_method_used = _get_response_values_binary(
            estimator, X, self._get_response_method(),
            pos_label=self.pos_label, return_response_method_used=True,
        )
        decision_threshold = 0.5 if self.threshold == "auto" and response_method_used == "predict_proba" else \
                             0.0 if self.threshold == "auto" else self.threshold
        return _threshold_scores_to_class_labels(
            y_score, decision_threshold, self.classes_, self.pos_label
        )

    @property
    def classes_(self):
        return self.estimator_.classes_

关键步骤:通过 _get_response_values_binary 获取 连续得分(概率或决策函数),并记录实际使用的响应方法。threshold="auto" 时,根据响应方法自动设为 0.5(概率)或 0.0(决策函数)。调用 _threshold_scores_to_class_labels 把得分与阈值比较,映射到 pos_label 对应的类别。classes_ 直接代理底层 estimator 的 classes_,保持标签一致性。

43.7.2 元数据路由与预拟合支持

测试代码test_fixed_threshold_classifier_metadata_routingtest_fixed_threshold_classifier_fitted_estimator

enable_metadata_routing=True 时,FixedThresholdClassifier 通过 set_fit_request(sample_weight=True)sample_weight 正确路由到底层 LogisticRegression.fit。当底层估计器已经预拟合,FixedThresholdClassifier.fit 不会再次调用 fit,直接使用 estimator_ 进行预测。

43.7.3 架构图(FixedThresholdClassifier 工作流)

flowchart LR subgraph 初始化 A[FixedThresholdClassifier(estimator, threshold, ...)] end subgraph 预测 B[predict(X)] --> C{是否已拟合?} C -- 已拟合 --> D[使用 estimator_] C -- 未拟合 --> E[调用 estimator.fit()] D --> F[_get_response_values_binary()] E --> F F --> G[确定决策阈值] G --> H[_threshold_scores_to_class_labels()] H --> I[返回类别标签] end

43.8 TunedThresholdClassifierCV —— 交叉验证阈值搜索的“智能雷达站”

FixedThresholdClassifier 只能手动调节阈值,实际场景往往需要 自动寻找最优阈值TunedThresholdClassifierCV 通过交叉验证、阈值扫描以及元数据路由,实现了端到端的阈值调优。

43.8.1 交叉验证阈值优化流程

源码路径sklearn/model_selection/_classification_threshold.py - TunedThresholdClassifierCV._fit(第 200‑500 行)

def _fit(self, X, y, **params):
    # 1️⃣ 解析 cv 参数 → 构造实际的交叉验证划分对象
    if isinstance(self.cv, Real) and 0 < self.cv < 1:
        cv = StratifiedShuffleSplit(...)
    elif self.cv == "prefit":
        # 检查 refit 与 prefit 的合法性
        ...
    else:
        cv = check_cv(self.cv, y=y, classifier=True)
        if self.refit is False and cv.get_n_splits() > 1:
            raise ValueError(...)

    routed_params = process_routing(self, "fit", **params)
    self._curve_scorer = self._get_curve_scorer()

    # 2️⃣ 决定最终 estimator_(用于预测)以及用于阈值搜索的 classifier
    if cv == "prefit":
        self.estimator_ = self.estimator
        classifier = self.estimator_
        splits = [(None, range(_num_samples(X)))]
    else:
        self.estimator_ = clone(self.estimator)
        classifier = clone(self.estimator)
        splits = cv.split(X, y, **routed_params.splitter.split)

        if self.refit:
            X_train, y_train, fit_params_train = X, y, routed_params.estimator.fit
        else:
            # 单折场景,只使用第一个训练划分
            train_idx, _ = next(cv.split(X, y, **routed_params.splitter.split))
            X_train = _safe_indexing(X, train_idx)
            y_train = _safe_indexing(y, train_idx)
            fit_params_train = _check_method_params(...)

        self.estimator_.fit(X_train, y_train, **fit_params_train))

    # 3️⃣ 并行遍历每个折,对每折执行一次拟合 + 多阈值评分
    cv_scores, cv_thresholds = zip(
        *Parallel(n_jobs=self.n_jobs)(
            delayed(_fit_and_score_over_thresholds)(
                clone(classifier) if cv != "prefit" else classifier,
                X, y,
                fit_params=routed_params.estimator.fit,
                train_idx=train_idx,
                val_idx=val_idx,
                curve_scorer=self._curve_scorer,
                score_params=routed_params.scorer.score,
            )
            for train_idx, val_idx in splits
        )
    )
    # 4️⃣ 检查常数预测器 → 抛异常
    if any(np.isclose(th[0], th[-1]) for th in cv_thresholds):
        raise ValueError(...)

    # 5️⃣ 统一阈值网格(线性或自定义),并对每折的阈值曲线做插值
    min_threshold = min(...); max_threshold = max(...)
    decision_thresholds = np.linspace(min_threshold, max_threshold, self.thresholds) \
        if isinstance(self.thresholds, Integral) else np.asarray(self.thresholds)

    objective_scores = _mean_interpolated_score(
        decision_thresholds, cv_thresholds, cv_scores
    )
    best_idx = objective_scores.argmax()
    self.best_score_ = objective_scores[best_idx]
    self.best_threshold_ = decision_thresholds[best_idx]

    if self.store_cv_results:
        self.cv_results_ = {"thresholds": decision_thresholds,
                            "scores": objective_scores}
    return self

关键环节

  • CV 解析:支持整数、浮点(单次 ShuffleSplit)、自定义 splitter、"prefit"

  • 元数据路由process_routing 自动处理 sample_weightgroups 等信息。

  • 并行阈值评分Parallel + _fit_and_score_over_thresholds 实现 每折一次拟合、多阈值评分

  • 常数预测器检测:若任一折的阈值范围为常数,则抛出明确错误。

  • 阈值网格:可通过整数生成等间距阈值或直接提供自定义阈值数组。

  • 最佳阈值选择:在统一网格上对每折的阈值曲线做线性插值,取平均得分最高的阈值。

43.8.2 预测逻辑

def predict(self, X):
    check_is_fitted(self, "estimator_")
    pos_label = self._curve_scorer._get_pos_label()
    y_score, _ = _get_response_values_binary(
        self.estimator_, X, self._get_response_method(), pos_label=pos_label,
    )
    return _threshold_scores_to_class_labels(
        y_score, self.best_threshold_, self.classes_, pos_label
    )
  • 使用 best_threshold_ 与底层 estimator 的连续得分进行类别映射。

43.8.3 元数据路由声明

def get_metadata_routing(self):
    router = (
        MetadataRouter(owner=self)
        .add(
            estimator=self.estimator,
            method_mapping=MethodMapping().add(callee="fit", caller="fit"),
        )
        .add(
            splitter=self.cv,
            method_mapping=MethodMapping().add(callee="split", caller="fit"),
        )
        .add(
            scorer=self._get_curve_scorer(),
            method_mapping=MethodMapping().add(callee="score", caller="fit"),
        )
    )
    return router
  • 明确 estimator.fitsplitter.splitscorer.score 三条路由渠道,保证 sample_weightgroups 等元数据在 训练 → 切分 → 评分 全链路传递。

43.8.4 架构图(TunedThresholdClassifierCV 端到端)

flowchart TD A[输入: X, y, estimator, cv, thresholds, scoring] --> B{cv 类型} B -- prefit --> C[使用已有 estimator_] B -- 常规 --> D[clone estimator, 生成 splits] D --> E[refit?] E -- True --> F[在全数据上 fit estimator_] E -- False --> G[在首个训练折上 fit estimator_] F & G --> H[Parallel: 对每折调用 _fit_and_score_over_thresholds] H --> I[收集 cv_scores, cv_thresholds] I --> J[检查常数预测器] J --> K[统一阈值网格 & 插值] K --> L[计算 objective_scores] L --> M[选取 best_threshold_ / best_score_] M --> N[可选存储 cv_results_]

43.9 设计中的取舍

:为何把阈值扫描放进 _fit_and_score_over_thresholds

:将阈值评估与模型拟合解耦虽然更灵活,但在实际调优流程中几乎总是 先拟合模型再评估阈值。若拆开,用户必须自行管理模型状态、数据划分,容易出错。_fit_and_score_over_thresholds 强制 一次拟合、多阈值,把常见模式封装成安全原语,降低出错概率,提升鲁棒性——这正是 scikit‑learn 更倾向提供“正确方式”而非极端低级 API 的取舍。

:元数据路由在阈值调优中的关键性?

:在真实业务场景(如医学诊断)中,sample_weightgroups 直接影响模型学习与评估。若阈值调优过程丢失这些信息,得到的“最佳阈值”将不再代表真实需求。启用 enable_metadata_routing=True 并在 Estimator、Splitter、Scorer 上声明需求,确保 从训练到阈值评分的全链路 都考虑了这些元数据,从而让调优结果在实际部署时保持可信。

43.10 动手练习

  • 阅读模型选择公共 API 与延迟加载机制

  • 分析 OneTimeSplitter 在测试中的作用

  • 探索阈值分类器的元数据路由与交叉验证集成

43.11 本章小结

本章我们深入探讨了 scikit‑learn 的 model_selection 模块,理解了它如何作为 机器学习模型调优的总指挥部,通过统一的公共 API、智能的延迟加载机制和精巧的元数据路由,协调切分器、验证器、搜索器和阈值调优器等各个组件高效协作。我们从模块的总入口 __init__.py 开始,看到它如何通过 __all__ 明确公开接口,并通过 __getattr__ 为实验性功能提供安全的延迟加载路径。随后,分析了测试体系中的 OneTimeSplitter,它确保交叉验证执行引擎不会重复使用切分器,从而保证测试的精准性。进一步,我们深入阈值调优的核心——_fit_and_score_over_thresholds,领悟了它如何通过 单次拟合、多阈值评分 的模式显著提升效率,并正确处理样本权重、fit_params 等元数据的路由。在此基础上,了解了 FixedThresholdClassifier 提供的手动阈值设定,以及 TunedThresholdClassifierCV 如何在交叉验证折上结合阈值扫描自动搜索最优阈值,同时妥善处理 refit 行为、常数预测器检测和元数据路由等关键细节。最后,通过设计取舍的讨论,体会到这些 API 背后的工程智慧:它们不是追求理论上的完全正交,而是专注于在实际工作流中提供安全、高效、易用的抽象。

本章我们一起学习了以下概念:

以下是本章学习的关键概念总结:

  • model_selection/__init__.py:模型选择模块总入口,聚合切分器、验证器、搜索器、可视化类与阈值分类器,__getattr__ 实现实验性组件延迟加载

  • OneTimeSplitter:测试工具类,包装 KFold 使其 split() 仅产出一次切分,验证交叉验证执行引擎不重复调用 CV 切分器

  • _fit_and_score_over_thresholds:阈值调优核心执行函数,单次 fit 后在多阈值上评分,支持预拟合、样本权重路由、fit_params 透传

  • _CurveScorer:曲线评分器,在多个决策阈值上评估指标,支撑 TunedThresholdClassifierCV 等高级功能

  • FixedThresholdClassifier:固定阈值分类器,将估计器连续响应(predict_proba/decision_function)按指定阈值与 pos_label 转为类别预测

  • TunedThresholdClassifierCV:交叉验证阈值调优器,利用 _CurveScorer 扫描阈值空间,优化指定评分指标,支持 refitstore_cv_results、字符串标签、元数据路由

  • 元数据路由集成:enable_metadata_routing=True 时,阈值分类器正确路由 sample_weightgroupsfit_params 等元数据到 estimator/scorer/splitter

第 44 章 —— 数据集模块概览

44.1 学习目标

  • 难度:★★★☆☆(3/5)

  • 预备知识:Python 基础、面向对象编程与 Markdown/代码阅读基础

  • 理解 metrics 与 model_selection 模块测试体系的整体架构与设计哲学

  • 掌握 test_common.py 中对称性、样本权重不变性、格式不变性等通用属性的系统性验证方法

  • 熟悉 Array API 合规性测试如何在 NumPy、CuPy、PyTorch 等后端下保证指标行为一致性

  • 理解回归测试与边界场景测试如何通过已知结果验证与极端值测试构筑精度防线

  • 掌握 datasets 模块测试套件中针对加载器、获取器、缓存行为与错误处理的验证策略

  • 了解模型选择模块中切分器、验证器、搜索器与阈值调优的测试覆盖范围

  • 能够阅读并扩展自定义指标或估计器的测试用例

  • 熟练运用离线测试夹具实现零网络依赖的确定性测试

动手练习

  • 为一个自定义的聚类指标编写 test_common.py 中的通用属性测试,验证其对称性、样本权重不变性和格式不变性。

  • 利用 array_namespace 装饰器在 test_regression.py 中添加一个针对 PyTorch 后端的回归指标一致性测试。

  • 修改 test_datasets.py 中的 fetch_20newsgroups 测试,使用 _fetch_fixture 装饰器替换真实网络请求为本地离线夹具,并验证返回数据的形状与目标标签正确性。

  • 在 test_model_selection.py 中扩展一个自定义的切分器,使其支持分组数据,并编写测试验证其在分层分组下的索引正确性。

44.2 生活类比

想象 scikit-learn 的测试体系是一个多层级的质量检测工厂:test_common.py 是通用质检标准实验室,每个指标(产品)出厂前必须通过对称性、不变性、边界值等硬性指标检测;Array API 合规性测试是跨平台兼容性认证中心,同一产品在 NumPy(标准车间)、CuPy(GPU 车间)、PyTorch(深度学习车间)等不同生产线上,必须产出质量一致的结果;回归测试与边界场景是极限压力测试场,模拟完美预测、全错预测、NaN/Inf 污染、零除等极端工况,确保指标不崩溃、数值精准;datasets/tests/ 是原材料入厂检验站,验证 load_(自产原料)、fetch_(外购原料)、make_(合成原料)的纯度、包装完整性、标签准确性;model_selection/tests/ 是流水线集成测试车间,切分器(分拣工)、验证器(质检流程)、搜索器(调参专家)、阈值调优(决策门控)协同作业,端到端验证;tests/data/openml/id_/ 是标准样品库,预置已知成分的标准品,离线测试时无需联网送检,即刻对比验证。

44.3 源码地图

sklearn/metrics/tests/init.py # 测试包标识

sklearn/metrics/tests/test_common.py

├── 通用属性验证

│ ├── test_symmetry # 对称性测试

│ ├── test_invariance_sample_weight # 样本权重不变性

│ ├── test_invariance_format # 格式不变性 (稠密/稀疏/数据帧)

│ ├── test_invariance_order # 样本顺序不变性

│ └── test_multiclass_multilabel_consistency # 多分类/多标签一致性

├── Array API 合规性

│ ├── _test_array_api_dispatch # 后端分派测试

│ ├── _test_array_api_consistency # 跨后端一致性

│ └── _test_device_support # 设备支持 (CPU/GPU)

├── 边界与回归场景

│ ├── test_perfect_prediction # 完美预测边界

│ ├── test_zero_division_handling # 零除处理

│ ├── test_nan_inf_handling # NaN/Inf 处理

│ └── test_known_results # 已知数值回归验证

sklearn/metrics/tests/test_classification.py # 分类指标测试

sklearn/metrics/tests/test_regression.py # 回归指标测试

sklearn/metrics/tests/test_ranking.py # 排序/曲线指标测试

sklearn/metrics/tests/test_pairwise.py # 成对距离/核测试

sklearn/metrics/tests/test_pairwise_distances_reduction.py # 邻居搜索分派测试

sklearn/metrics/tests/test_score_objects.py # 评分器系统测试

sklearn/metrics/cluster/tests/init.py # 聚类测试包标识

sklearn/metrics/cluster/tests/test_common.py # 聚类通用属性测试

sklearn/metrics/cluster/tests/test_supervised.py # 有监督聚类指标测试

sklearn/metrics/cluster/tests/test_unsupervised.py # 无监督聚类指标测试

sklearn/metrics/cluster/tests/test_bicluster.py # 双聚类测试

sklearn/metrics/_plot/tests/init.py # 可视化测试包标识

sklearn/metrics/_plot/tests/test_common_curve_display.py # 通用曲线显示测试

sklearn/model_selection/tests/init.py # 模型选择测试包标识

sklearn/model_selection/tests/common.py

│ └── OneTimeSplitter # 单次切分测试工具

sklearn/model_selection/tests/test_split.py # 数据切分器测试

sklearn/model_selection/tests/test_validation.py # 交叉验证执行测试

sklearn/model_selection/tests/test_search.py # 超参数搜索测试

sklearn/model_selection/tests/test_successive_halving.py # 逐次减半搜索测试

sklearn/model_selection/tests/test_classification_threshold.py # 阈值调优测试

sklearn/model_selection/tests/test_plot.py # 学习/验证曲线可视化测试

sklearn/datasets/tests/init.py # 数据集测试包标识

sklearn/datasets/tests/test_base.py # 基础加载器/获取器测试

sklearn/datasets/tests/test_common.py # 数据集通用属性测试

sklearn/datasets/tests/test_california_housing.py # 加州房产测试

sklearn/datasets/tests/test_covtype.py # 森林覆盖测试

sklearn/datasets/tests/test_kddcup99.py # KDD99 测试

sklearn/datasets/tests/test_lfw.py # LFW 人脸测试

sklearn/datasets/tests/test_olivetti_faces.py # Olivetti 人脸测试

sklearn/datasets/tests/test_openml.py # OpenML 交互测试

sklearn/datasets/tests/test_rcv1.py # RCV1 测试

sklearn/datasets/tests/test_20news.py # 20 Newsgroups 测试

sklearn/datasets/tests/test_arff_parser.py # ARFF 解析测试

sklearn/datasets/tests/test_samples_generator.py # 合成数据生成器测试

sklearn/datasets/tests/test_svmlight_format.py # SVMLight 格式测试

sklearn/datasets/tests/data/init.py # 测试数据包标识

sklearn/datasets/tests/data/openml/init.py # OpenML 离线测试夹具根目录

sklearn/datasets/tests/data/openml/id_*/init.py # 具体数据集离线快照标识

44.4 指标通用属性的系统性验证 —— 质检标准实验室的“硬性指标体系”

为什么需要 test_common.py?它是所有度量指标的“宪法”——抽离出对称性、不变性、一致性等与具体数学公式无关的通用数学性质,避免在每个 test_*.py 中重复编写样板测试代码。

对称性测试验证 metric(y_true, y_pred) == metric(y_pred, y_true) 对于对称指标(如 accuracy_score、 rand_score)成立,并通过参数化标记自动跳过非对称指标(如 log_loss、 precision_score)。

不变性三件套:样本权重不变性(加权=复制样本)、样本顺序不变性(打乱不变)、容器格式不变性(稠密/稀疏/DataFrame 等价),确保指标不依赖数据呈现形式。

多分类/多标签一致性:验证 average='macro/micro/weighted/samples' 等聚合策略在二分类扩展到多分类、多标签时的数学自洽性。

pytest 参数化驱动:通过 @pytest.mark.parametrize('metric_func', [accuracy_score, f1_score, ...]) 实现‘写一次测试逻辑,跑遍所有指标’的工程效率。

Array API 合规性集成:在通用测试中引入 array_namespace 检测,自动在 NumPy/CuPy/PyTorch/JAX/Dask 后端下重跑核心断言,一键保障跨后端行为一致。

边界与回归守卫:test_perfect_prediction(全对)、test_all_wrong(全错)、test_zero_division(分母为零)、test_nan_inf(脏数据)、test_known_results(经典数值回归)构成多层安全网。

源码路径:sklearn/metrics/tests/test_common.py - test_symmetry(1-800行)

源码路径:sklearn/metrics/tests/test_classification.py - accuracy_score(1-500行)

源码路径:sklearn/metrics/tests/test_regression.py - mean_squared_error(1-400行)

源码路径:sklearn/metrics/tests/test_ranking.py - roc_auc_score(1-500行)

源码路径:sklearn/metrics/tests/test_pairwise.py - euclidean_distances(1-600行)

源码路径:sklearn/utils/_array_api.py - array_namespace(1-200行)

这段代码定义了通用测试框架的核心结构,包括 REGRESSION_METRICS、CLASSIFICATION_METRICS、CONTINUOUS_CLASSIFICATION_METRICS 和 CURVE_METRICS 四个指标字典,以及 SYMMETRIC_METRICS、NOT_SYMMETRIC_METRICS 等属性集合。它通过 pytest 参数化机制实现了对所有指标的通用属性验证,如对称性、不变性和 Array API 一致性测试,是 metrics 测试体系的基础设施。

graph TD A[指标] --> B[对称性测试] A --> C[不变性测试] A --> D[格式不变性测试] A --> E[样本顺序不变性测试] A --> F[多分类/多标签一致性测试] A --> G[Array API 合规性测试] A --> H[边界与回归守卫]

44.5 评分器与可视化系统的测试覆盖 —— 从‘万能遥控器’到‘驾驶舱仪表盘’的质检

评分器系统测试 (test_score_objects.py):验证 make_scorer 包装器正确捕获响应方法 (predict_proba/decision_function/predict)、符号翻转 (greater_is_better)、元数据路由 (sample_weight 透传) 与多指标缓存优化 (_MultimetricScorer 避免重复预测)。

曲线评分器测试:_CurveScorer 在多阈值下评估指标,支撑 TunedThresholdClassifierCV 阈值搜索,测试覆盖阈值数组生成、插值逻辑、边界阈值处理。

check_scoring 规范化测试:字符串别名 ('accuracy')、可调用对象、列表/字典多指标、错误输入拦截,验证全局 _SCORERS 注册表的完整性。

可视化 Display 测试 (_plot/tests/):ConfusionMatrixDisplay/RocCurveDisplay/PrecisionRecallDisplay/DetCurveDisplay/PredictionErrorDisplay 从 from_estimator/from_predictions 构造到 plot() 绘图的端到端验证。

通用曲线显示基类测试 (test_common_curve_display.py):抽离曲线类共享的坐轴变换、图例处理、样式参数验证、多曲线叠加等公共逻辑,避免各 Display 测试重复造轮子。

绘图工具测试 (utils/test_plotting.py):_validate_style_kwargs 样式合法性、_BinaryClassifierCurveDisplayMixin 二分类曲线混入方法的独立单元测试。

源码路径:sklearn/metrics/tests/test_score_objects.py - make_scorer(1-600行)

源码路径:sklearn/metrics/_plot/tests/test_confusion_matrix_display.py - ConfusionMatrixDisplay(1-300行)

源码路径:sklearn/metrics/_plot/tests/test_roc_curve_display.py - RocCurveDisplay(1-300行)

源码路径:sklearn/metrics/_plot/tests/test_precision_recall_display.py - PrecisionRecallDisplay(1-300行)

源码路径:sklearn/metrics/_plot/tests/test_det_curve_display.py - DetCurveDisplay(1-200行)

源码路径:sklearn/metrics/_plot/tests/test_predict_error_display.py - PredictionErrorDisplay(1-200行)

源码路径:sklearn/metrics/_plot/tests/test_common_curve_display.py - _validate_plot_params(1-200行)

源码路径:sklearn/utils/tests/test_plotting.py - _validate_style_kwargs(1-200行)

这段代码实现了评分器系统的核心逻辑,通过 make_scorer 将度量函数包装为统一接口,支持响应方法选择、符号翻转和元数据路由。它还引入了 _MultimetricScorer 实现多指标缓存优化,以及 _CurveScorer 支持阈值扫描,为 TunedThresholdClassifierCV 等高级功能提供基础。

graph TD A[度量函数] --> B[make_scorer] B --> C[_Scorer基类] C --> D[_MultimetricScorer缓存] C --> E[_CurveScorer阈值扫描] E --> F[TunedThresholdClassifierCV] B --> G[check_scoring规范化] G --> H[_SCORERS注册表]

44.6 聚类评估指标的分层测试策略 —— 无监督世界的‘多维体检单’

有监督聚类指标测试 (test_supervised.py):rand_score、adjusted_rand_score、mutual_info_score/normalized_mutual_info_score/adjusted_mutual_info_score、homogeneity_completeness_v_measure 等基于信息论与组合数学的指标,验证对称性、上界归一化、随机基线期望值。

无监督内部指标测试 (test_unsupervised.py):silhouette_score(轮廓系数需邻域计算)、calinski_harabasz_score(方差比)、davies_bouldin_score(簇内散度/簇间距离),覆盖稠密/稀疏输入、预计算距离矩阵、采样近似参数。

双聚类一致性测试 (test_bicluster.py):consensus_score 评估行/列同时聚类的一致性,验证置换不变性、完全匹配得分。

快速互信息 Cython 加速测试:_expected_mutual_info_fast.pyx 中的精确期望互信息计算,对比纯 Python 实验验证数值等价与性能提升。

聚类通用属性测试 (test_common.py):复用 metrics 通用测试框架,验证聚类指标的对称性、格式不变性、Array API 合规性,体现测试基建的复用哲学。

源码路径:sklearn/metrics/cluster/tests/test_supervised.py - adjusted_rand_score(1-500行)

源码路径:sklearn/metrics/cluster/tests/test_unsupervised.py - silhouette_score(1-400行)

源码路径:sklearn/metrics/cluster/tests/test_bicluster.py - consensus_score(1-300行)

源码路径:sklearn/metrics/cluster/tests/test_common.py - rand_score(1-200行)

源码路径:sklearn/metrics/cluster/_expected_mutual_info_fast.pyx - expected_mutual_information(1-300行)

这段代码实现了互信息的快速计算,通过在对数空间使用 gamma 函数避免直接计算大的阶乘,从而防止数值溢出。它是聚类评估中 adjusted_mutual_info_score 的核心组件,通过预计算对数 gamma 值和谨慎的求和范围,实现了在大规模 contingency 矩阵上的高性能和数值稳定计算。

graph TD A[contingency 矩阵] --> B[计算边际和] B --> C[对数空间中的 gamma 函数] C --> D[精确期望互信息] D --> E[adjusted_mutual_info_score]

44.7 成对距离与邻居搜索分派器测试 —— 高性能内核的‘压力测试台’

pairwise 距离/核测试 (test_pairwise.py):pairwise_distances 统一调度、check_pairwise_arrays 输入合法性、欧氏/曼哈顿/余弦/哈弗辛距离的数值正确性、pairwise_kernels 核函数族(RBF/多项式/Sigmoid/拉普拉斯)的内存优化路径。

缺失值与成对分量距离:nan_euclidean_distances 缺失值优雅处理(忽略 NaN 维度重新归一化)、paired_distances 逐行对齐距离计算。

分块计算与 Cython 并行 (_pairwise_fast.pyx):pairwise_distances_chunked 生成器模式缓解内存压力、OpenMP 并行加速验证、大规模数据分块策略正确性。

分派器架构测试 (test_pairwise_distances_reduction.py):BaseDistancesReductionDispatcher 统一接口、is_usable_for 数据适用性判定、ArgKmin/RadiusNeighbors 分派器按 dtype 选择 32/64 位优化实现。

类模式归约与行范数:ArgKminClassMode/RadiusNeighborsClassMode 集成加权投票、sqeuclidean_row_norms 高效并行行范数计算,验证分类/回归下游任务的数值一致性。

源码路径:sklearn/metrics/tests/test_pairwise.py - pairwise_distances(1-600行)

源码路径:sklearn/metrics/tests/test_pairwise_distances_reduction.py - BaseDistancesReductionDispatcher(1-400行)

源码路径:sklearn/metrics/_pairwise_fast.pyx - _chi2_kernel_fast(1-500行)

源码路径:sklearn/metrics/_pairwise_distances_reduction/_dispatcher.py - is_usable_for(1-300行)

源码路径:sklearn/metrics/_pairwise_distances_reduction/_classmode.pxd - WeightingStrategy(1-100行)

这段代码实现了按 dtype 自动选择 32 位或 64 位优化实现的分派器模式,是 scikit-learn 高性能成对距离计算的核心。它通过 is_usable_for 方法判断数据是否适用于特定分派器,然后将控制权交给对应的 dtype 特化实现(如 ArgKmin32/ArgKmin64),从而在保持 API 一致性的同时实现底层计算的性能优化。

graph TD A[输入数据 X, Y] --> B{BaseDistancesReductionDispatcher} B --> C{is_usable_for 检查} C -->|float64| D[ArgKmin64/RadiusNeighbors64] C -->|float32| E[ArgKmin32/RadiusNeighbors32] D --> F[64位优化实现] E --> G[32位优化实现] F --> H[结果] G --> H

44.8 模型选择核心组件测试矩阵 —— 从切分器到搜索器的‘全链路验收’

数据切分器测试 (test_split.py):BaseCrossValidator/BaseShuffleSplit/GroupsConsumerMixin 接口契约,KFold/StratifiedKFold/GroupKFold 逐折索引正确性与分层保真,TimeSeriesSplit 时间有序性,LeaveOneOut/LeavePOut/LeaveOneGroupOut/LeavePGroupsOut 穷举策略,RepeatedKFold 随机重复机制,PredefinedSplit 用户自定义切分,check_cv 规范化入口与 train_test_split 便捷封装。

交叉验证执行引擎测试 (test_validation.py):cross_validate 多指标并行调度、_fit_and_score 拟合/评分/计时整合、cross_val_predict 预测收集与 _enforce_prediction_order 类别顺序保障、learning_curve/validation_curve 渐进式训练与参数扫描、permutation_test_score 目标置换生成经验零分布显著性评估。

超参数搜索测试 (test_search.py + test_successive_halving.py):ParameterGrid 笛卡尔积与 ParameterSampler 分布采样正确性,GridSearchCV/RandomizedSearchCV 穷举/随机 _run_search 实现,BaseSuccessiveHalving 资源递增与候选减半淘汰逻辑,HalvingGridSearchCV/HalvingRandomSearchCV 网格/随机候选生成与逐轮淘汰结合。

阈值调优分类器测试 (test_classification_threshold.py):FixedThresholdClassifier 硬阈值决策转换,TunedThresholdClassifierCV 利用 _CurveScorer 与 CV 在阈值空间搜索最优决策点,业务指标驱动的阈值优化。

可视化曲线测试 (test_plot.py):LearningCurveDisplay/ValidationCurveDisplay 从 from_estimator 构造到绘图的端到端验证,误差带渲染、对数坐户、多指标叠加。

测试公共设施 (common.py):OneTimeSplitter 仅切分一次的测试专用切分器,支撑各测试模块快速构造确定性切分场景。

源码路径:sklearn/model_selection/tests/test_split.py - KFold(1-800行)

源码路径:sklearn/model_selection/tests/test_validation.py - cross_validate(1-800行)

源码路径:sklearn/model_selection/tests/test_search.py - GridSearchCV(1-600行)

源码路径:sklearn/model_selection/tests/test_successive_halving.py - HalvingGridSearchCV(1-400行)

源码路径:sklearn/model_selection/tests/test_classification_threshold.py - TunedThresholdClassifierCV(1-300行)

源码路径:sklearn/model_selection/tests/test_plot.py - LearningCurveDisplay(1-300行)

源码路径:sklearn/model_selection/tests/common.py - OneTimeSplitter(1-100行)

这段代码实现了一个专用的测试切分器 OneTimeSplitter,它封装了 KFold 但只产生一次分割,用于在测试中快速创建确定性的训练/测试分割场景,避免每次测试都需要初始化完整的交叉验证器。

44.9 datasets 模块测试体系 —— 原材料入厂检验站的‘全流程把关’

基础加载器/获取器测试 (test_base.py):load_iris/load_wine/load_breast_cancer/load_digits/load_diabetes/load_linnerud 等经典玩具数据集的形状、返回格式 (Bunch/ndarray/DataFrame)、as_frame/return_X_y 参数组合、描述字段 (DESCR/feature_names/target_names) 完整性验证。

远程真实数据集测试 (test_california_housing.py 等):fetch_california_housing/fetch_covtype/fetch_kddcup99/fetch_lfw_*/fetch_olivetti_faces/fetch_species_distributions 下载缓存行为、SHA256 校验、文件解析预处理(子采样/标签二值化/图像切片/灰度转换)、缓存损坏自检与重试机制。

OpenML 交互与离线测试 (test_openml.py + tests/data/openml/):fetch_openml 按 name/version/id 检索、ARFF 解析双引擎 (LIAC-ARFF/pandas)、重试装饰器与原子缓存、目标列验证、_fetch_fixture 装饰器拦截网络调用重定向至本地离线夹具,实现零网络依赖 CI。

文本与多标签数据测试 (test_20news.py、test_rcv1.py):fetch_20newsgroups 头部/引用剥离、子集筛选、向量化管线,fetch_rcv1 稀疏特征与多标签主题向量、置换还原与子采样。

合成数据生成器测试 (test_samples_generator.py):make_classification/make_regression 信息特征/冗余特征/噪声参数化控制,make_blobs/make_moons/make_circles/make_swiss_roll/make_s_curve 流形嵌入几何性质,make_biclusters/make_checkerboard 结构化模式生成。

格式解析往返测试 (test_svmlight_format.py、test_arff_parser.py):dump_svmlight_file/load_svmlight_file(s) 稀疏/稠密/多标签/查询ID 往返一致性,ARFF 稀疏列提取、类别编码、分块读取解析正确性。

通用属性测试 (test_common.py):数据集返回对象的键集合一致性、Bunch 属性式访问、缓存目录管理 (get_data_home/clear_data_home) 幂等性与环境变量覆盖。

源码路径:sklearn/datasets/tests/test_base.py - load_iris(1-800行)

源码路径:sklearn/datasets/tests/test_california_housing.py - fetch_california_housing(1-200行)

源码路径:sklearn/datasets/tests/test_covtype.py - fetch_covtype(1-200行)

源码路径:sklearn/datasets/tests/test_kddcup99.py - fetch_kddcup99(1-200行)

源码路径:sklearn/datasets/tests/test_lfw.py - fetch_lfw_people(1-200行)

源码路径:sklearn/datasets/tests/test_olivetti_faces.py - fetch_olivetti_faces(1-200行)

源码路径:sklearn/datasets/tests/test_openml.py - fetch_openml(1-500行)

源码路径:sklearn/datasets/tests/test_rcv1.py - fetch_rcv1(1-300行)

源码路径:sklearn/datasets/tests/test_20news.py - fetch_20newsgroups(1-300行)

源码路径:sklearn/datasets/tests/test_arff_parser.py - _liac_arff_parser(1-300行)

源码路径:sklearn/datasets/tests/test_samples_generator.py - make_classification(1-500行)

源码路径:sklearn/datasets/tests/test_svmlight_format.py - dump_svmlight_file(1-400行)

源码路径:sklearn/datasets/tests/test_common.py - check_as_frame(1-200行)

源码路径:sklearn/datasets/tests/data/openml/id_1/__init__.py - _fetch_fixture(1-1行)

这段代码定义了数据集测试的离线夹具机制,通过在 tests/data/openml/id_*/ 目录下放置 init.py 文件(可以为空)来标识该数据集的离线快照可用。当 fetch_openml 在测试环境中被调用时,_fetch_fixture 装饰器会拦截网络请求并重定向到本地对应的数据文件,实现零网络依赖的确定性测试。

44.10 本章小结

本章我们学习了 scikit-learn 的测试体系如何通过多层次验证确保指标和模型选择组件的正确性。我们首先探索了 test_common.py 如何作为所有指标的通用质检标准,抽离出对称性、不变性等通用属性;然后分析了 Array API 合规性测试如何跨后端保证行为一致性;接着回顾了回归测试与边界场景如何通过已知结果和极端值验证构筑精度防线;随后考察了评分器和可视化系统的测试覆盖,了解其如何从单指标扩展到多指标缓存和阈值搜索;然后深入了聚类评估的分层测试策略,从有监督到无监督再到双聚类;之后审视了成对距离与邻居搜索分派器的高性能内核测试,理解其如何通过 dtype 分派实现跨平台优化;然后考察了模型选择核心组件的全链路验收,从切分器到搜索器的接口契约和行为验证;最后掌握了 datasets 模块的全流程把关策略,从基础加载器到远程真实数据集、合成数据生成器以及离线测试夹具的协同作用。

本章我们一起学习了以下概念:

本章小结前的概念表格概括了 scikit-learn 测试体系中的核心验证策略和测试方法,涵盖了从通用属性验证到离线测试夹具的各个方面。

| 概念 | 解释 |

|------|------|

| test_common.py 通用属性验证 | 对称性、样本权重/顺序/格式不变性、多分类一致性,构成指标正确性的‘宪法’ |

| Array API 合规性测试 | 通过 array_namespace 自动识别后端,在 NumPy/CuPy/PyTorch/JAX/Dask 上验证指标行为一致性 |

| 回归测试与边界场景 | 已知数值验证 + 极端值(完美/全错/NaN/零除)压力测试,筑牢数值精度防线 |

| 评分器系统测试 | 验证 _BaseScorer、_MultimetricScorer 缓存优化、_CurveScorer 阈值扫描、check_scoring 规范化逻辑 |

| 聚类指标测试分层 | 有监督(RI/MI/V-measure)、无监督(轮廓/CH/DB)、双聚类、快速互信息 Cython 加速分别验证 |

| 可视化 Display 测试 | 从构造器 (from_estimator/predictions) 到 plot 绘图细节,以及通用曲线显示基类的公共测试设施 |

| 数据切分器测试矩阵 | KFold/Stratified/Group/TimeSeries/LeaveOneOut/Repeated/Predefined 等切分策略的索引正确性与分层保真度验证 |

| 交叉验证执行引擎测试 | cross_validate 多指标并行、cross_val_predict 顺序保障、learning_curve/validation_curve 渐进扫描、置换检验显著性评估 |

| 超参数搜索测试覆盖 | ParameterGrid/ParameterSampler 生成正确性、Grid/Random/Halving 搜索流程、逐轮淘汰与资源递增逻辑 |

| 阈值调优分类器测试 | FixedThresholdClassifier 硬阈值、TunedThresholdClassifierCV 交叉验证阈值搜索与 _CurveScorer 集成 |

| datasets 测试策略 | load_* 形状/返回格式/描述字段、fetch_* 缓存行为/下载重试/校验、make_* 参数化生成可控性、SVMLight/ARFF 格式解析往返一致性 |

| 离线测试夹具 | tests/data/openml/id_*/init.py 标识预置数据快照,消除测试对网络依赖,实现确定性 CI |

| 测试工具集复用 | OneTimeSplitter、assert_allclose、ignore_warnings、TempMemmap、CheckingClassifier 等基建工具支撑全库测试编写 |

设计中的取舍

问:为什么在 test_common.py 中使用参数化来验证所有指标的通用属性,而不是为每个指标编写单独的测试?

答:这种设计避免了样板代码的重复,提高了工程效率。通过参数化,我们可以编写一次测试逻辑,然后在所有指标上运行它,确保一致性并减少维护负担。

问:Array API 合规性测试如何在不牺牲性能的前提下保证跨后端行为一致性?

答:它使用 array_namespace 动态检测后端,并在 NumPy/CuPy/PyTorch 等后端下重新运行核心断言。这种方法仅在必要时进行后端分派,否则回退到 NumPy 兼容模式,从而在保证一致性的同时最小化性能开销。

问:在数据集测试中使用离线夹具(如 tests/data/openml/id_*/init.py)的权衡是什么?

答:离线夹具消除了测试对网络的依赖,使 CI 更快速和确定性。然而,它要求维护本地数据快照,这可能会占用磁盘空间,并且需要定期更新以反映上游数据的更改。

问:在模型选择测试中,OneTimeSplitter 的设计取舍是什么?

答:OneTimeSplitter 是一个专用的测试工具,它封装了 KFold 但仅产生一次分割。这使得测试能够快速构建确定性的训练/测试场景,而无需为每个测试实例化完整的交叉验证器,从而提高了测试速度和可重复性。其取舍在于功能受限(仅一次分割),但换来了测试的简便性和效率。

下一章中,我们将学习 datasets 模块的详细实现,包括数据主目录管理、远程文件抓取与校验、CSV 与压缩数据的本地装载以及 SVMLight 格式的读写引擎。

44.11 设计中的取舍

为什么采用当前方案,而不是更复杂的替代方案? 本章源码优先选择清晰、可维护且与既有 API 兼容的实现;这降低了使用和调试成本,但也意味着部分极端场景需要调用者自行权衡性能、灵活性与实现复杂度。

44.12 动手练习

44.12.1 阅读 metrics 通用属性测试实现

阅读 sklearn/metrics/tests/test_common.pytest_symmetrytest_invariance_sample_weighttest_invariance_format 的实现

回答问题:

  • 这些测试如何通过 pytest.mark.parametrize 自动化覆盖分类、回归、聚类等不同指标?

  • test_invariance_format 如何构造稠密数组、CSR 稀疏矩阵、pandas DataFrame 等不同容器输入?

  • 对称性测试中 y_true, y_pred 互换为何对某些指标(如 log_loss)不适用?如何在参数化中排除?

44.12.2 剖析 Array API 合规性测试机制

阅读 sklearn/metrics/tests/test_common.py 中 Array API 相关测试函数及 sklearn/utils/_array_api.pyarray_namespace 实现

回答问题:

  • _test_array_api_dispatch 如何验证指标函数正确分派到对应后端(NumPy/CuPy/PyTorch)?

  • 跨后端一致性测试中,如何处理不同后端浮点数精度差异导致的断言阈值调整?

  • 如果某后端不支持特定操作(如 PyTorch 无 nanmedian),测试框架如何优雅跳过而非失败?

44.12.3 设计自定义指标的测试用例

参考 sklearn/metrics/tests/test_classification.pytest_precision_recall_fscore_supporttest_log_loss 的模式

要求:

  • 为一个自定义的非对称分类指标 asymmetric_fbeta_score(beta, alpha) 编写完整测试用例

  • 包含:对称性(预期失败)、样本权重不变性、格式不变性、完美/全错预测边界、零除处理、已知数值回归验证

  • 使用 pytest.mark.parametrize 覆盖二分类、多分类 (ovr/ovo)、多标签场景

  • 集成 Array API 合规性测试标记(如 @pytest.mark.array_api_compatible

44.12.4 追踪 datasets 离线测试夹具工作流

阅读 sklearn/datasets/tests/test_openml.pysklearn/datasets/tests/data/openml/id_1/__init__.py 等离线夹具

回答问题:

  • fetch_openml 在测试环境下如何被 _fetch_fixture 装饰器拦截并重定向到本地 tests/data/openml/

  • id_* 目录下的 ARFF 文件与 __init__.py 空文件如何配合实现 importlib.resources 定位?

  • 如果要为新数据集 id_999 添加离线测试,需要在 tests/data/openml/id_999/ 放置哪些文件?命名约定是什么?

第 45 章 —— 数据基座与文件调度 —— 构筑“数据仓库的传送带”

45.1 学习目标

  • 难度:★★★☆☆(3/5)

  • 预备知识:Python 基础、面向对象编程与 Markdown/代码阅读基础

本章节将帮助你:

  • 了解数据主目录的自检与脚本直接运行的行为机制。

  • 掌握 __main__ 块在数据模块中的定位与使用场景。

  • 深入理解模块作为脚本执行时的路径解析与依赖加载原理。

  • 熟悉数据加载器的快速验证与演示方式。

  • 掌握 SVMLight/LibSVM 格式的高性能解析与序列化原理。

  • 理解稀疏与稠密数据在 Cython 层的零拷贝转换技术。

  • 明白多标签、查询 ID 等扩展特性在文件格式中的编码方式。

生活类比:把 sklearn.datasets 想象成一家 智能仓储公司的演示厅。当用户直接运行模块(python -m sklearn.datasets._base)时,就像走进演示厅,自动播放的产品介绍(__main__ 块)会展示各种样本(Iris、Digits 等),帮助用户快速判断数据是否符合需求。SVMLight 解析器则是 高速自动分拣机:文本行像传送带上的包裹,Cython 内核 _load_svml_svmlight_file 是扫描头,瞬间识别标签、索引与数值并装入稀疏矩阵(CSR)货箱。随后的 get_dense_row_string / get_sparse_row_string 如同逆向打包线,将内存中的矩阵重新封装为文本。整个过程无 GIL、零拷贝、流式处理,类似分拣机全天候不间断运转。


45.2 源码地图

sklearn/datasets/_base.py
├── __main__
│   ├── 演示数据加载功能
│   └── 自测数据完整性与格式
├── get_data_home()
├── clear_data_home()
├── load_csv_data()
├── load_gzip_compressed_csv_data()
├── load_descr()
├── load_files()
├── load_iris()
├── load_wine()
├── load_breast_cancer()
├── load_digits()
├── load_diabetes()
├── load_linnerud()
├── load_sample_images()
├── load_sample_image()
├── _convert_data_dataframe()
├── _pkl_filepath()
├── _sha256()
├── _fetch_remote()
├── _filter_filename()
├── _derive_folder_and_filename_from_url()
└── fetch_file()
sklearn/datasets/_svmlight_format_io.py
├── load_svmlight_file()
├── load_svmlight_files()
├── _gen_open()
├── _open_and_load()
├── _dump_svmlight()
└── dump_svmlight_file()
sklearn/datasets/_svmlight_format_fast.pyx
├── _load_svmlight_file()
├── get_dense_row_string()
├── get_sparse_row_string()
└── _dump_svmlight_file()

说明:上述树状图展示了核心文件与函数的层级关系。后续章节会分别展开每个子模块的实现细节。


45.3 数据主目录管理 —— 构建可配置的缓存根目录

45.3.1 代码实行(逐行注释)

# 第 45 章 —— src/sklearn/datasets/_base.py - get_data_home() (20-36 行)
@validate_params(
    {"data_home": [str, os.PathLike, None]},  # 参数校验:接受 str、Path 或 None
    prefer_skip_nested_validation=True,
)
def get_data_home(data_home=None) -> str:
    """
    Return the path of the scikit‑learn data directory.
    """
    # ① 若未显式传入 data_home,则优先读取环境变量 SCIKIT_LEARN_DATA
    #    若环境变量不存在,则回退到用户主目录下的默认路径 ~/scikit_learn_data
    if data_home is None:
        data_home = environ.get("SCIKIT_LEARN_DATA", join("~", "scikit_learn_data"))

    # ② 展开 “~” 为实际用户目录,兼容 Linux/macOS/Windows
    data_home = expanduser(data_home)

    # ③ 按需创建目录,exist_ok=True 防止并发创建导致的竞争
    makedirs(data_home, exist_ok=True)
    return data_home

解释:环境变量优先策略让高级用户可以自行指定缓存位置,保持灵活性。expanduser 确保在所有平台上都能正确解析 ~makedirs(..., exist_ok=True) 在多进程/多线程场景下安全创建,避免 FileExistsError

45.3.2 清理缓存的“一键清空”

# 第 45 章 —— src/sklearn/datasets/_base.py - clear_data_home() (38-50 行)
@validate_params(
    {"data_home": [str, os.PathLike, None]},  # 同上
    prefer_skip_nested_validation=True,
)
def clear_data_home(data_home=None):
    """
    Delete all the content of the data home cache.
    """
    # 统一使用 get_data_home 获取目录路径(保证路径一致性)
    data_home = get_data_home(data_home)

    # 递归删除整个目录,等价于“核按钮”
    shutil.rmtree(data_home)

解释:当本地缓存损坏或需要强制重新下载时,调用此函数即可彻底清空目录,避免手动逐文件删除的繁琐。

45.3.3 流程图(Mermaid)

flowchart TD A[调用 get_data_home()] --> B{环境变量 SCIKIT_LEARN_DATA ?} B -- 是 --> C[使用环境变量的路径] B -- 否 --> D[使用默认路径 ~/scikit_learn_data] C --> E[expanduser() 解析 ~] D --> E E --> F[makedirs(..., exist_ok=True)] F --> G[返回 data_home] style A fill:#f9f,stroke:#333,stroke-width:2px style G fill:#bbf,stroke:#333,stroke-width:2px

类比延伸:这相当于仓库的“入口门禁系统”,先检查是否有专属通道(环境变量),没有则走公共通道(默认路径),随后打开门(创建目录)让后续货物进出。


45.4 远程文件抓取与校验 —— 安全可靠的下载流程

45.4.1 SHA256 校验(逐行注释)

# 第 45 章 —— src/sklearn/datasets/_base.py - _sha256() (120-132 行)
def _sha256(path):
    """Calculate the sha256 hash of the file at path."""
    sha256hash = hashlib.sha256()
    chunk_size = 8192                     # 8 KB,避免一次性读取大文件导致 OOM
    with open(path, "rb") as f:
        while True:
            buffer = f.read(chunk_size)    # 分块读取
            if not buffer:
                break
            sha256hash.update(buffer)      # 累计更新哈希
    return sha256hash.hexdigest()

解释:分块读取既节省内存,又能对任意大小文件进行完整性校验。

45.4.2 原子下载(逐行注释)

# 第 45 章 —— src/sklearn/datasets/_base.py - _fetch_remote() (134-186 行)
def _fetch_remote(remote, dirname=None, n_retries=3, delay=1):
    """
    下载 remote.url 到本地目录,采用原子写入并校验 SHA256。
    """
    folder_path = Path(dirname) if dirname else Path(".")
    file_path = folder_path / remote.filename

    # 若本地已有文件且校验通过,则直接返回,避免重复下载
    if file_path.exists():
        if remote.checksum is None:
            return file_path
        if _sha256(file_path) == remote.checksum:
            return file_path
        else:
            warnings.warn(
                f"SHA256 checksum of existing file {file_path.name} mismatched; re‑downloading."
            )

    # 创建唯一临时文件,防止并发冲突
    temp_file = NamedTemporaryFile(
        prefix=remote.filename + ".part_", dir=folder_path, delete=False
    )
    temp_file.close()                      # 立即释放文件句柄,保证唯一占用
    try:
        temp_path = Path(temp_file.name)
        while True:
            try:
                urlretrieve(remote.url, temp_path)   # 下载到临时文件
                break
            except (URLError, TimeoutError):
                if n_retries == 0:
                    raise
                warnings.warn(f"Retry downloading {remote.url}")
                n_retries -= 1
                time.sleep(delay)

        # 下载完成后校验 SHA256
        if remote.checksum and _sha256(temp_path) != remote.checksum:
            raise OSError("Checksum mismatch after download.")
    except Exception:
        os.unlink(temp_file.name)           # 清理残留的临时文件
        raise

    # 原子移动:若同一文件系统上,move 是原子操作
    shutil.move(temp_path, file_path)
    return file_path

解释原子写入:使用 NamedTemporaryFile + shutil.move,确保下载过程即使被中断也不会留下半截文件。重试机制:在网络不稳时自动重试 n_retries 次,提升成功率。异常清理:捕获异常后手动删除临时文件,防止磁盘产生垃圾。

45.4.3 URL → 本地路径推导(逐行注释)

# 第 45 章 —— src/sklearn/datasets/_base.py - _derive_folder_and_filename_from_url() (198-226 行)
def _derive_folder_and_filename_from_url(url):
    """
    将 URL 拆解为安全的本地文件夹层级与文件名。
    """
    parsed_url = urlparse(url)
    if not parsed_url.hostname:
        raise ValueError(f"Invalid URL: {url}")

    # 主机名作为顶层文件夹(保留点以区分子域)
    folder_components = [_filter_filename(parsed_url.hostname, filter_dots=False)]

    path = parsed_url.path
    if "/" in path:
        base_folder, raw_filename = path.rsplit("/", 1)
        base_folder = _filter_filename(base_folder)   # 去除非法字符
        if base_folder:
            folder_components.append(base_folder)
    else:
        raw_filename = path

    filename = _filter_filename(raw_filename, filter_dots=False) or "downloaded_file"
    return "/".join(folder_components), filename

解释:通过过滤非法字符(_filter_filename),确保生成的路径在所有文件系统上安全。示例:http://example.com/data/foo.csvscikit_learn_data/example/data/foo.csv

45.4.4 统一下载入口 fetch_file

# 第 45 章 —— src/sklearn/datasets/_base.py - fetch_file() (228-258 行)
def fetch_file(url, folder=None, local_filename=None, sha256=None,
               n_retries=3, delay=1):
    """
    根据 URL 下载文件,若本地已存在且 SHA256 匹配则直接返回。
    """
    folder_from_url, filename_from_url = _derive_folder_and_filename_from_url(url)

    # 允许用户自定义文件名或保存目录
    local_filename = local_filename or filename_from_url
    folder = Path(get_data_home()) / folder_from_url if folder is None else Path(folder)
    makedirs(folder, exist_ok=True)

    remote_metadata = RemoteFileMetadata(
        filename=local_filename, url=url, checksum=sha256
    )
    return _fetch_remote(remote_metadata, dirname=folder,
                          n_retries=n_retries, delay=delay)

解释fetch_file 将 URL → 本地路径映射、目录创建、下载调用流程统一封装,提供给上层 API 使用。

45.4.5 流程图(Mermaid)

flowchart LR A[fetch_file(url)] --> B[_derive_folder_and_filename_from_url] B --> C{本地文件是否存在?} C -- 是 --> D[_sha256 校验] D -- 匹配 --> E[返回本地路径] D -- 不匹配 --> F[_fetch_remote] C -- 否 --> F F --> G[下载到临时文件] G --> H[校验 SHA256] H -- 成功 --> I[shutil.move(原子移动)] I --> E style F fill:#ffe,stroke:#333,stroke-width:2px style I fill:#bbf,stroke:#333,stroke-width:2px

类比延伸:这相当于 签收快递 过程:先检查是否已经在仓库(本地)并验收(SHA256),没有则派送快递员(urlretrieve)并在签收后把快递搬进仓库(原子移动)。


45.5 CSV 与压缩数据的本地装载 —— 高效转化包内静态资源

45.5.1 load_csv_data(逐行注释)

# 第 45 章 —— src/sklearn/datasets/_base.py - load_csv_data() (278-320 行)
def load_csv_data(
    data_file_name,
    *,
    data_module=DATA_MODULE,
    descr_file_name=None,
    descr_module=DESCR_MODULE,
    encoding="utf-8",
):
    """
    从包内资源读取 CSV,并返回 data、target、target_names(可选 descr)。
    """
    # 1️⃣ 使用 importlib.resources 统一定位包内文件
    data_path = resources.files(data_module) / data_file_name
    with data_path.open("r", encoding=encoding) as csv_file:
        data_file = csv.reader(csv_file)

        # 2️⃣ 读取首行获取尺寸信息(避免动态扩容)
        temp = next(data_file)
        n_samples = int(temp[0])
        n_features = int(temp[1])
        target_names = np.array(temp[2:])

        # 3️⃣ 预分配内存
        data = np.empty((n_samples, n_features))
        target = np.empty((n_samples,), dtype=int)

        # 4️⃣ 逐行读取并填充数组
        for i, ir in enumerate(data_file):
            data[i] = np.asarray(ir[:-1], dtype=np.float64)
            target[i] = np.asarray(ir[-1], dtype=int)

    # 5️⃣ 若需要描述文件,则使用 load_descr 读取
    if descr_file_name is None:
        return data, target, target_names
    else:
        descr = load_descr(descr_module=descr_module,
                          descr_file_name=descr_file_name)
        return data, target, target_names, descr

解释预分配np.empty((n_samples, n_features)) 防止在循环中不断扩容。资源定位resources.files 通过包的元数据定位,兼容源码安装、wheel、conda。

45.5.2 load_gzip_compressed_csv_data(逐行注释)

# 第 45 章 —— src/sklearn/datasets/_base.py - load_gzip_compressed_csv_data() (322-364 行)
def load_gzip_compressed_csv_data(
    data_file_name,
    *,
    data_module=DATA_MODULE,
    descr_file_name=None,
    descr_module=DESCR_MODULE,
    encoding="utf-8",
    **kwargs,
):
    """
    读取 gzip 压缩的 CSV,采用零拷贝流式解压 → np.loadtxt。
    """
    data_path = resources.files(data_module) / data_file_name
    # ① 以二进制方式打开压缩资源
    with data_path.open("rb") as compressed_file:
        # ② 直接在流上使用 gzip 解压,避免生成临时解压文件
        compressed_file = gzip.open(compressed_file, mode="rt", encoding=encoding)
        # ③ np.loadtxt 读取解压后文本,**kwargs 允许自定义分隔符等
        data = np.loadtxt(compressed_file, **kwargs)

    if descr_file_name is None:
        return data
    else:
        descr = load_descr(descr_module=descr_module,
                          descr_file_name=descr_file_name)
        return data, descr

解释:完全在内存流中完成解压 → 解析,既省磁盘 I/O,也避免峰值内存。

45.5.3 load_descr(逐行注释)

# 第 45 章 —— src/sklearn/datasets/_base.py - load_descr() (366-386 行)
def load_descr(descr_file_name, *, descr_module=DESCR_MODULE, encoding="utf-8"):
    """
    读取 .rst 描述文件的纯文本内容。
    """
    path = resources.files(descr_module) / descr_file_name
    return path.read_text(encoding=encoding)

解释:为数据集提供人类可读的背景说明,便于教学与交互式探索。

45.5.4 流程图(Mermaid)

flowchart TD A[load_csv_data()] --> B[resources.files → data_path] B --> C[打开 CSV (text mode)] C --> D[读取首行获取 n_samples, n_features] D --> E[预分配 NumPy 数组] E --> F[逐行填充 data & target] F --> G{descr_file_name?} G -- 有 --> H[load_descr() 读取 .rst] G -- 无 --> I[返回 data, target, target_names] H --> I style D fill:#ddf,stroke:#333,stroke-width:2px style F fill:#edf,stroke:#333,stroke-width:2px

类比延伸:这相当于 仓库内部的自动装配线:先定位商品(资源文件),再检查尺寸(首行),随后快速分配仓位(预分配),最后逐件装箱(逐行填充)。


45.6 经典玩具数据集加载器 —— 即拿即用的标准样本之间

以下示例展示了 加载器结构统一:读取 CSV → 生成特征名称 → 可选 as_frame 转 Pandas → 返回 Bunch。这里重点展示 load_irisload_wineload_digits 三个加载器的关键实现,其他加载器遵循相同模式。

45.6.1 load_iris(逐行注释)

# 第 45 章 —— src/sklearn/datasets/_base.py - load_iris() (400-446 行)
@validate_params(
    {"return_X_y": ["boolean"], "as_frame": ["boolean"]},
    prefer_skip_nested_validation=True,
)
def load_iris(*, return_X_y=False, as_frame=False):
    """Load and return the iris dataset (classification)."""
    data_file_name = "iris.csv"
    # 读取数据、目标、目标名称以及描述文件
    data, target, target_names, fdescr = load_csv_data(
        data_file_name=data_file_name,
        descr_file_name="iris.rst"
    )

    feature_names = [
        "sepal length (cm)",
        "sepal width (cm)",
        "petal length (cm)",
        "petal width (cm)",
    ]

    frame = None
    target_columns = ["target"]
    if as_frame:
        frame, data, target = _convert_data_dataframe(
            "load_iris", data, target, feature_names, target_columns
        )

    if return_X_y:
        return data, target

    return Bunch(
        data=data,
        target=target,
        frame=frame,
        target_names=target_names,
        DESCR=fdescr,
        feature_names=feature_names,
        filename=data_file_name,
        data_module=DATA_MODULE,
    )

45.6.2 load_wine

# 第 45 章 —— src/sklearn/datasets/_base.py - load_wine() (448-498 行)
@validate_params(
    {"return_X_y": ["boolean"], "as_frame": ["boolean"]},
    prefer_skip_nested_validation=True,
)
def load_wine(*, return_X_y=False, as_frame=False):
    """Load and return the wine dataset (classification)."""
    data, target, target_names, fdescr = load_csv_data(
        data_file_name="wine_data.csv",
        descr_file_name="wine_data.rst"
    )

    feature_names = [
        "alcohol", "malic_acid", "ash", "alcalinity_of_ash",
        "magnesium", "total_phenols", "flavanoids",
        "nonflavanoid_phenols", "proanthocyanins",
        "color_intensity", "hue", "od280/od315_of_diluted_wines",
        "proline",
    ]

    frame = None
    target_columns = ["target"]
    if as_frame:
        frame, data, target = _convert_data_dataframe(
            "load_wine", data, target, feature_names, target_columns
        )

    if return_X_y:
        return data, target

    return Bunch(
        data=data,
        target=target,
        frame=frame,
        target_names=target_names,
        DESCR=fdescr,
        feature_names=feature_names,
    )

45.6.3 load_digits(包含 gzip 解压)

# 第 45 章 —— src/sklearn/datasets/_base.py - load_digits() (500-556 行)
@validate_params(
    {"n_class": [Interval(Integral, 1, 10, closed="both")],
     "return_X_y": ["boolean"], "as_frame": ["boolean"]},
    prefer_skip_nested_validation=True,
)
def load_digits(*, n_class=10, return_X_y=False, as_frame=False):
    """Load and return the digits dataset (classification)."""
    # 读取 gzip 压缩的 CSV,内部已完成流式解压
    data, fdescr = load_gzip_compressed_csv_data(
        data_file_name="digits.csv.gz",
        descr_file_name="digits.rst",
        delimiter=","
    )

    target = data[:, -1].astype(int, copy=False)   # 最后一列为标签
    flat_data = data[:, :-1]                        # 前 64 列为特征
    images = flat_data.view()
    images.shape = (-1, 8, 8)                       # 零拷贝视图

    # 可选只保留前 n_class 类
    if n_class < 10:
        idx = target < n_class
        flat_data, target = flat_data[idx], target[idx]
        images = images[idx]

    feature_names = [
        f"pixel_{r}_{c}" for r in range(8) for c in range(8)
    ]

    frame = None
    target_columns = ["target"]
    if as_frame:
        frame, flat_data, target = _convert_data_dataframe(
            "load_digits", flat_data, target, feature_names, target_columns
        )

    if return_X_y:
        return flat_data, target

    return Bunch(
        data=flat_data,
        target=target,
        frame=frame,
        feature_names=feature_names,
        target_names=np.arange(10),
        images=images,
        DESCR=fdescr,
    )

45.6.4 流程图(Mermaid)

flowchart TD A[load_iris()] --> B[load_csv_data()] B --> C[生成 feature_names] C --> D{as_frame?} D -- 是 --> E[_convert_data_dataframe()] D -- 否 --> F[直接返回 Bunch] E --> F style B fill:#def,stroke:#333,stroke-width:2px style D fill:#f9f,stroke:#333,stroke-width:2px

45.6.5 样本图像加载(新增流程图)

# 第 45 章 —— src/sklearn/datasets/_base.py - load_sample_images() (558-606 行)
def load_sample_images():
    """Load the two sample JPEG images (china.jpg, flower.jpg)."""
    try:
        from PIL import Image
    except ImportError as exc:
        raise ImportError(
            "PIL is required to load JPEG images. Install Pillow."
        ) from exc

    descr = load_descr("README.txt", descr_module=IMAGES_MODULE)

    filenames, images = [], []
    # 只加载 .jpg 文件并保持字母顺序
    jpg_paths = sorted(
        r for r in resources.files(IMAGES_MODULE).iterdir()
        if r.is_file() and r.match("*.jpg")
    )
    for path in jpg_paths:
        filenames.append(str(path))
        with path.open("rb") as f:
            pil_image = Image.open(f)
            images.append(np.asarray(pil_image))

    return Bunch(images=images, filenames=filenames, DESCR=descr)
# 第 45 章 —— src/sklearn/datasets/_base.py - load_sample_image() (608-632 行)
@validate_params(
    {"image_name": [StrOptions({"china.jpg", "flower.jpg"})]},
    prefer_skip_nested_validation=True,
)
def load_sample_image(image_name):
    """Return a single sample image as a NumPy array."""
    images = load_sample_images()
    for idx, fname in enumerate(images.filenames):
        if fname.endswith(image_name):
            return images.images[idx]
    raise AttributeError(f"Cannot find sample image: {image_name}")

45.6.5.1 流程图(Mermaid)(新增)

flowchart TD A[load_sample_images()] --> B[检查 Pillow 是否可用] B --> C[定位 JPEG 资源文件] C --> D[逐文件打开 → PIL 解码 → 转 ndarray] D --> E[返回 Bunch(images, filenames, DESCR)] style C fill:#ddf,stroke:#333,stroke-width:2px

解释load_sample_images 自动检测 PIL 是否可用,若缺失给出明确错误提示。两张图片以 Bunch 返回,方便后续 images.images[0] 直接使用。


45.7 兼容性与转换工具函数(新增流程图)

45.7.1 _convert_data_dataframe

# 第 45 章 —— src/sklearn/datasets/_base.py - _convert_data_dataframe() (634-660 行)
def _convert_data_dataframe(
    caller_name, data, target, feature_names, target_names, sparse_data=False
):
    """
    将 NumPy/稀疏矩阵转换为 pandas DataFrame(as_frame=True 时使用)。
    """
    pd = check_pandas_support(f"{caller_name} with as_frame=True")
    # 稠密 vs 稀疏分支
    if not sparse_data:
        data_df = pd.DataFrame(data, columns=feature_names, copy=False)
    else:
        data_df = pd.DataFrame.sparse.from_spmatrix(data, columns=feature_names)

    target_df = pd.DataFrame(target, columns=target_names)
    combined_df = pd.concat([data_df, target_df], axis=1)

    X = combined_df[feature_names]
    y = combined_df[target_names]

    # 如果目标只有一列,返回 Series 而不是单列 DataFrame
    if y.shape[1] == 1:
        y = y.iloc[:, 0]
    return combined_df, X, y

解释check_pandas_support 动态检测 pandas 是否可用,保持轻量化部署仍能工作。对稀疏矩阵使用 DataFrame.sparse.from_spmatrix,避免稠密化导致内存激增。

45.7.2 _pkl_filepath

# 第 45 章 —— src/sklearn/datasets/_base.py - _pkl_filepath() (662-672 行)
def _pkl_filepath(*args, **kwargs):
    """
    为 Python 3 的 pickle 文件生成兼容路径(在 .pkl 前加 _py3)。
    """
    py3_suffix = kwargs.get("py3_suffix", "_py3")
    basename, ext = splitext(args[-1])            # 提取文件名与扩展名
    basename += py3_suffix                         # 插入后缀
    new_args = args[:-1] + (basename + ext,)
    return join(*new_args)

解释:在同一目录下保留 _py3 与原始 .pkl,实现 Python 2/3 兼容。

45.7.3 流程图(Mermaid)(新增)

flowchart TD A[_convert_data_dataframe()] --> B[检测 pandas 是否可用] B --> C[稠密路径: pandas.DataFrame] B --> D[稀疏路径: DataFrame.sparse.from_spmatrix] C & D --> E[拼接特征与目标 → 返回 combined, X, y] style B fill:#f9f,stroke:#333,stroke-width:2px

45.8 SVMLight 格式高性能解析器 —— 稀疏数据的“极速翻译官”

45.8.1 load_svmlight_file(逐行注释)

# 第 45 章 —— src/sklearn/datasets/_svmlight_format_io.py - load_svmlight_file() (20-56 行)
@validate_params(
    {"f": [str, Interval(Integral, 0, None, closed="left"), os.PathLike,
           HasMethods("read")],
     "n_features": [Interval(Integral, 1, None, closed="left"), None],
     "dtype": "no_validation",
     "multilabel": ["boolean"],
     "zero_based": ["boolean", StrOptions({"auto"})],
     "query_id": ["boolean"],
     "offset": [Interval(Integral, 0, None, closed="left")],
     "length": [Integral],
    },
    prefer_skip_nested_validation=True,
)
def load_svmlight_file(
    f,
    *,
    n_features=None,
    dtype=np.float64,
    multilabel=False,
    zero_based="auto",
    query_id=False,
    offset=0,
    length=-1,
):
    """
    统一入口:调用 load_svmlight_files 处理单文件列表。
    """
    return tuple(
        load_svmlight_files(
            [f],
            n_features=n_features,
            dtype=dtype,
            multilabel=multilabel,
            zero_based=zero_based,
            query_id=query_id,
            offset=offset,
            length=length,
        )
    )

解释:为保持 API 一致性,单文件入口直接转到多文件实现 load_svmlight_files

45.8.2 _gen_open(逐行注释)

# 第 45 章 —— src/sklearn/datasets/_svmlight_format_io.py - _gen_open() (58-78 行)
def _gen_open(f):
    """
    根据文件扩展名返回对应的二进制打开对象。
    """
    if isinstance(f, int):                         # 文件描述符
        return open(f, "rb", closefd=False)
    elif isinstance(f, os.PathLike):
        f = os.fspath(f)
    elif not isinstance(f, str):
        raise TypeError(f"expected {{str, int, path-like, file-like}}, got {type(f)}")

    _, ext = os.path.splitext(f)
    if ext == ".gz":
        import gzip
        return gzip.open(f, "rb")
    elif ext == ".bz2":
        from bz2 import BZ2File
        return BZ2File(f, "rb")
    else:
        return open(f, "rb")

解释:统一处理普通文件、gzip、bz2 以及整数文件描述符,使后续解析函数可以透明地接受多种输入形式。

45.8.3 _open_and_load(逐行注释)

# 第 45 章 —— src/sklearn/datasets/_svmlight_format_io.py - _open_and_load() (80-106 行)
def _open_and_load(f, dtype, multilabel, zero_based, query_id, offset=0, length=-1):
    """
    打开文件(或直接使用已有的文件对象),调用 Cython 内核进行解析。
    """
    if hasattr(f, "read"):      # 已经是 file‑like
        actual_dtype, data, ind, indptr, labels, query = _load_svmlight_file(
            f, dtype, multilabel, zero_based, query_id, offset, length
        )
    else:
        # 使用 _gen_open 根据扩展名打开
        with closing(_gen_open(f)) as f:
            actual_dtype, data, ind, indptr, labels, query = _load_svmlight_file(
                f, dtype, multilabel, zero_based, query_id, offset, length
            )

    # 将 array.array 转换为 NumPy,保持零拷贝
    if not multilabel:
        labels = np.frombuffer(labels, np.float64)
    data = np.frombuffer(data, actual_dtype)
    indices = np.frombuffer(ind, np.longlong)
    indptr = np.frombuffer(indptr, np.longlong)
    query = np.frombuffer(query, np.int64)

    data = np.asarray(data, dtype=dtype)  # 强制目标 dtype
    return data, indices, indptr, labels, query

解释:该函数封装了文件打开、Cython 解析以及 array.array → NumPy 零拷贝的所有步骤,外部只需要关心返回的稀疏结构。

45.8.4 load_svmlight_files(逐行注释)

# 第 45 章 —— src/sklearn/datasets/_svmlight_format_io.py - load_svmlight_files() (108-176 行)
@validate_params(
    {"files": ["array-like", str, os.PathLike, HasMethods("read"),
               Interval(Integral, 0, None, closed="left")],
     "n_features": [Interval(Integral, 1, None, closed="left"), None],
     "dtype": "no_validation",
     "multilabel": ["boolean"],
     "zero_based": ["boolean", StrOptions({"auto"})],
     "query_id": ["boolean"],
     "offset": [Interval(Integral, 0, None, closed="left")],
     "length": [Integral],
    },
    prefer_skip_nested_validation=True,
)
def load_svmlight_files(
    files,
    *,
    n_features=None,
    dtype=np.float64,
    multilabel=False,
    zero_based="auto",
    query_id=False,
    offset=0,
    length=-1,
):
    """
    同时加载多个 SVMLight 文件,统一特征维度。
    """
    # 若使用 offset/length,则关闭 auto heuristic,保证一致性
    if (offset != 0 or length > 0) and zero_based == "auto":
        zero_based = True

    if (offset != 0 or length > 0) and n_features is None:
        raise ValueError("n_features 必须在使用 offset/length 时提供。")

    # 对每个文件调用 _open_and_load,返回 (data, indices, indptr, y, query)
    r = [
        _open_and_load(
            f, dtype, multilabel, bool(zero_based), bool(query_id),
            offset=offset, length=length,
        )
        for f in files
    ]

    # 自动检测 one‑based 并统一转为 zero‑based
    if zero_based is False or (
        zero_based == "auto" and all(len(tmp[1]) and np.min(tmp[1]) > 0 for tmp in r)
    ):
        for _, indices, _, _, _ in r:
            indices -= 1

    # 计算整体特征数(最大的索引 + 1)
    n_f = max(ind[1].max() if len(ind[1]) else 0 for ind in r) + 1

    if n_features is None:
        n_features = n_f
    elif n_features < n_f:
        raise ValueError(
            f"n_features 被设为 {n_features},但实际需要 {n_f} 个特征"
        )

    result = []
    for data, indices, indptr, y, query_vals in r:
        shape = (indptr.shape[0] - 1, n_features)
        X = sp.csr_matrix((data, indices, indptr), shape)
        X.sort_indices()                    # 保证 CSR 索引有序
        result += [X, y]
        if query_id:
            result.append(query_vals)

    return result

解释特征对齐:先解析每个文件的最大特征索引,统一填充零向量,保证后续训练/测试维度一致。Zero‑based 统一:若文件使用 one‑based(常见 libsvm),在这里统一减 1。返回结构[X1, y1, X2, y2, …] 或带 query_id 的三元组,便于解包。

45.8.5 Cython 内核 _load_svmlight_file(逐行注释)

# 第 45 章 —— src/sklearn/datasets/_svmlight_format_fast.pyx - _load_svmlight_file() (20-124 行)
def _load_svmlight_file(f, dtype, bint multilabel, bint zero_based,
                        bint query_id, long long offset, long long length):
    cdef array.array data, indices, indptr
    cdef bytes line
    cdef char *hash_ptr
    cdef char *line_cstr
    cdef int idx, prev_idx
    cdef Py_ssize_t i
    cdef bytes qid_prefix = b'qid'
    cdef Py_ssize_t n_features
    cdef long long offset_max = offset + length if length > 0 else -1

    # ① 根据 dtype 初始化 data 缓冲区(float32 vs float64)
    if dtype == np.float32:
        data = array.array("f")
    else:
        dtype = np.float64
        data = array.array("d")

    indices = array.array("q")               # long long 索引
    indptr = array.array("q", [0])           # 行指针
    query = np.arange(0, dtype=np.int64)     # 初始查询 ID(仅在需要时扩大)

    if multilabel:
        labels = []                         # 多标签存为 Python list
    else:
        labels = array.array("d")            # 单标签使用 array.array

    # ② 若需跳过前 offset 字节,先定位并丢弃残缺行
    if offset > 0:
        f.seek(offset)
        f.readline()

    # ③ 主循环:逐行读取
    for line in f:
        # ---- 跳过注释(#)----
        line_cstr = line
        hash_ptr = strchr(line_cstr, 35)   # ASCII '#'
        if hash_ptr != NULL:
            line = line[:hash_ptr - line_cstr]

        line_parts = line.split()
        if len(line_parts) == 0:
            continue                         # 空行或全注释行

        # ---- 目标与特征分离 ----
        target, features = line_parts[0], line_parts[1:]

        # ---- 处理多标签 -----
        if multilabel:
            if COLON in target:             # 若目标本身携带 ':' → 视作特征
                target, features = [], line_parts[0:]
            else:
                target = [float(y) for y in target.split(COMMA)]
            target.sort()
            labels.append(tuple(target))
        else:
            # 单标签使用 array.resize_smart 预分配,避免频繁 Python list append
            array.resize_smart(labels, len(labels) + 1)
            labels[len(labels) - 1] = float(target)

        # ---- 处理 qid(查询 ID)----
        prev_idx = -1
        n_features = len(features)
        if n_features and features[0].startswith(qid_prefix):
            _, value = features[0].split(COLON, 1)
            if query_id:
                query = np.append(query, np.int64(value))
            features.pop(0)
            n_features -= 1

        # ---- 解析每个 feature:idx:value ----
        for i in range(0, n_features):
            idx_s, value = features[i].split(COLON, 1)
            idx = int(idx_s)
            # 索引合法性检查
            if idx < 0 or (not zero_based and idx == 0):
                raise ValueError(f"Invalid index {idx} in SVMLight file.")
            if idx <= prev_idx:
                raise ValueError("Feature indices must be sorted and unique.")
            # 追加索引和值到数组
            array.resize_smart(indices, len(indices) + 1)
            indices[len(indices) - 1] = idx
            array.resize_smart(data, len(data) + 1)
            data[len(data) - 1] = float(value)
            prev_idx = idx

        # ---- 更新 indptr(每行结束)----
        array.resize_smart(indptr, len(indptr) + 1)
        indptr[len(indptr) - 1] = len(data)

        # ---- 检查是否已超过 offset+length ----
        if offset_max != -1 and f.tell() > offset_max:
            break

    # 返回所有缓冲区以及查询 ID(即使未使用仍返回空数组)
    return (dtype, data, indices, indptr, labels, query)

解释:采用 Cythonarray.array 实现 零拷贝:所有数据在 C 层直接写入缓冲区,随后通过 np.frombuffer 转为 NumPy,避免 Python 列表的频繁扩容。使用 array.resize_smart 动态扩容而不是 Python list.append,因为 array.array 在 C 级别扩容更高效且保持连续内存布局。无 GIL:整个解析过程在 C 层运行,无需获取全局解释器锁,适合并行读取大文件。异常安全:若索引不递增或出现负值,将立即抛出 ValueError,帮助定位损坏的 SVMLight 文件。

45.8.6 行序列化工具 get_dense_row_string / get_sparse_row_string

# 第 45 章 —— src/sklearn/datasets/_svmlight_format_fast.pyx - get_dense_row_string() (126-152 行)
def get_dense_row_string(
    const int_or_float[:, :] X,
    Py_ssize_t[:] x_inds,
    double_or_longlong[:] x_vals,
    Py_ssize_t row,
    str value_pattern,
    bint one_based,
):
    """
    将稠密矩阵的单行转换为 "index:value" 空格分隔字符串。
    """
    cdef Py_ssize_t row_length = X.shape[1]
    cdef Py_ssize_t x_nz_used = 0
    cdef Py_ssize_t k
    cdef int_or_float val

    for k in range(row_length):
        val = X[row, k]
        if val == 0:
            continue
        x_inds[x_nz_used] = k
        x_vals[x_nz_used] = <double_or_longlong> val
        x_nz_used += 1

    # 使用用户提供的格式化模式(比如 "%d:%.16g")
    reprs = [
        value_pattern % (x_inds[i] + one_based, x_vals[i])
        for i in range(x_nz_used)
    ]
    return " ".join(reprs)
# 第 45 章 —— src/sklearn/datasets/_svmlight_format_fast.pyx - get_sparse_row_string() (154-176 行)
def get_sparse_row_string(
    int_or_float[:] X_data,
    int[:] X_indptr,
    int[:] X_indices,
    Py_ssize_t row,
    str value_pattern,
    bint one_based,
):
    """
    将 CSR 稀疏矩阵的单行转换为 "index:value"。
    """
    cdef Py_ssize_t row_start = X_indptr[row]
    cdef Py_ssize_t row_end = X_indptr[row+1]

    reprs = [
        value_pattern % (X_indices[i] + one_based, X_data[i])
        for i in range(row_start, row_end)
    ]
    return " ".join(reprs)

解释:两个函数分别为稠密与稀疏路径提供高效的行字符串化,实现 “逆向打包线”,用于 dump_svmlight_file 的写回。

45.8.7 流程图(Mermaid)

flowchart LR A[load_svmlight_files()] --> B[_open_and_load() 对每个文件] B --> C[_load_svmlight_file (Cython) 解析行] C --> D[返回 data, indices, indptr, labels, query] D --> E[统一特征维度 → 构造 CSR 矩阵] E --> F[返回 X, y (以及可选的 query_id)] style B fill:#ffe,stroke:#333,stroke-width:2px style E fill:#bbf,stroke:#333,stroke-width:2px

45.9 SVMLight 格式高效序列化 —— 从内存矩阵到标准文本的“打包流水线”

45.9.1 Python 层包装 _dump_svmlight

# 第 45 章 —— src/sklearn/datasets/_svmlight_format_io.py - _dump_svmlight() (178-200 行)
def _dump_svmlight(X, y, f, multilabel, one_based, comment, query_id):
    """
    负责写入文件头(注释、版本、索引基数),并调用 Cython 内核完成逐行写入。
    """
    if comment:
        f.write(
            (f"# Generated by dump_svmlight_file from scikit-learn {__version__}\n"
             f"# Column indices are {'zero' if not one_based else 'one'}-based\n"
             "#\n").encode()
        )
        f.writelines(b"# %s\n" % line for line in comment.splitlines())

    X_is_sp = sp.issparse(X)
    y_is_sp = sp.issparse(y)

    # 确保 y 为二维(单标签时强制列向量)
    if not multilabel and not y_is_sp:
        y = y[:, np.newaxis]

    _dump_svmlight_file(
        X, y, f, multilabel, one_based, query_id,
        X_is_sp, y_is_sp,
    )

解释:首先写入统一的文件头,随后判断稀疏/稠密情况并把控制权交给 Cython 实现,以获得最高的写文件性能。

45.9.2 dump_svmlight_file(逐行注释)

# 第 45 章 —— src/sklearn/datasets/_svmlight_format_io.py - dump_svmlight_file() (202-294 行)
@validate_params(
    {"X": ["array-like", "sparse matrix"],
     "y": ["array-like", "sparse matrix"],
     "f": [str, HasMethods(["write"])],
     "zero_based": ["boolean"],
     "comment": [str, bytes, None],
     "query_id": ["array-like", None],
     "multilabel": ["boolean"]},
    prefer_skip_nested_validation=True,
)
def dump_svmlight_file(
    X,
    y,
    f,
    *,
    zero_based=True,
    comment=None,
    query_id=None,
    multilabel=False,
):
    """
    将稠密或 CSR 矩阵写入 SVMLight 文本文件。
    """
    # —— 参数检查与统一 —— #
    if comment is not None:
        if isinstance(comment, bytes):
            comment.decode("ascii")        # 确保是 ASCII,否则报错
        else:
            comment = comment.encode("utf-8")
        if b"\0" in comment:
            raise ValueError("comment 包含 NUL 字符")

    yval = check_array(y, accept_sparse="csr", ensure_2d=False)
    if sp.issparse(yval):
        if yval.shape[1] != 1 and not multilabel:
            raise ValueError(f"expected y shape (n_samples,1), got {yval.shape}")
    else:
        if yval.ndim != 1 and not multilabel:
            raise ValueError(f"expected y shape (n_samples,), got {yval.shape}")

    Xval = check_array(X, accept_sparse="csr")
    if Xval.shape[0] != yval.shape[0]:
        raise ValueError(f"X 与 y 行数不匹配:{Xval.shape[0]} vs {yval.shape[0]}")

    # —— 排序 CSR 索引以防止历史 bug(#1501)—— #
    if hasattr(yval, "sorted_indices"):
        y = yval.sorted_indices()
    else:
        y = yval
        if hasattr(y, "sort_indices"):
            y.sort_indices()

    if hasattr(Xval, "sorted_indices"):
        X = Xval.sorted_indices()
    else:
        X = Xval
        if hasattr(X, "sort_indices"):
            X.sort_indices()

    # —— 处理 query_id —— #
    if query_id is None:
        query_id = np.array([], dtype=np.int32)
    else:
        query_id = np.asarray(query_id)
        if query_id.shape[0] != y.shape[0]:
            raise ValueError(f"query_id 行数需与 y 对齐,got {query_id.shape}")

    one_based = not zero_based

    # —— 实际写入 —— #
    if hasattr(f, "write"):
        _dump_svmlight(X, y, f, multilabel, one_based, comment, query_id)
    else:
        with open(f, "wb") as f:
            _dump_svmlight(X, y, f, multilabel, one_based, comment, query_id)

解释参数统一check_array 将稠密/稀疏输入统一为 CSR,确保后端 Cython 能直接使用。索引排序:防止因 CSR 索引未排序导致写回异常。one_based 计算zero_based=False 时输出 one‑based 索引,保持与 libsvm 兼容。

45.9.3 Cython 写出实现 _dump_svmlight_file(关键段落)

# 第 45 章 —— src/sklearn/datasets/_svmlight_format_fast.pyx - _dump_svmlight_file() (178-238 行)
def _dump_svmlight_file(
    X,
    y,
    f,
    bint multilabel,
    bint one_based,
    int_or_longlong[:] query_id,
    bint X_is_sp,
    bint y_is_sp,
):
    # 1️⃣ 确定数值格式(整数 vs 浮点)
    X_is_integral = X.dtype.kind == "i"
    value_pattern = "%d:%d" if X_is_integral else "%d:%.16g"
    label_pattern = "%d" if y.dtype.kind == "i" else "%.16g"

    # 2️⃣ 行模板,若有 query_id 则加入
    line_pattern = "%s"
    if query_id.size > 0:
        line_pattern = "%s qid:%d %s\n"
    else:
        line_pattern = "%s %s\n"

    # 3️⃣ 为稠密路径预分配缓冲区
    cdef Py_ssize_t row_length = X.shape[1]
    cdef Py_ssize_t[:] x_inds = np.empty(row_length, dtype=np.intp)
    cdef signed long long[:] x_vals_int
    cdef double[:] x_vals_float
    if not X_is_sp:
        if X_is_integral:
            x_vals_int = np.zeros(row_length, dtype=np.longlong)
        else:
            x_vals_float = np.zeros(row_length, dtype=np.float64)

    # 4️⃣ 遍历每一行并写入
    for i in range(X.shape[0]):
        if not X_is_sp:
            # 稠密路径:使用 get_dense_row_string
            s = get_dense_row_string(
                X, x_inds,
                x_vals_int if X_is_integral else x_vals_float,
                i, value_pattern, one_based
            )
        else:
            # 稀疏路径:直接遍历 CSR 切片
            s = get_sparse_row_string(
                X.data, X.indptr, X.indices,
                i, value_pattern, one_based
            )

        # 目标标签转字符串(单标签或多标签)
        if multilabel:
            if y_is_sp:
                col_start = y.indptr[i]
                col_end = y.indptr[i+1]
                labels_str = ','.join(
                    label_pattern % y.indices[j]
                    for j in range(col_start, col_end)
                    if y.data[j] != 0
                )
            else:
                labels_str = ','.join(
                    label_pattern % j
                    for j in range(y.shape[1])
                    if y[i, j] != 0
                )
        else:
            if y_is_sp:
                labels_str = label_pattern % y.data[i]
            else:
                labels_str = label_pattern % y[i, 0]

        # 组装最终行并写入文件
        if query_id.size > 0:
            f.write((line_pattern % (labels_str, query_id[i], s)).encode())
        else:
            f.write((line_pattern % (labels_str, s)).encode())

解释稠密 vs 稀疏分支:稠密路径使用预分配的 x_inds / x_vals,稀疏路径直接利用 CSR 的 indptr / indices,实现 几乎无额外开销 的写回。多标签处理:当 multilabel=True 时,若目标是稀疏矩阵,会遍历其非零索引并使用逗号拼接;若是稠密数组,则遍历列检查非零位置,同样生成逗号分隔的标签串。one_based 决定是否在写出时将列索引加 1,以兼容 libsvm 的 one‑based 约定。

45.9.4 流程图(Mermaid)

flowchart LR A[dump_svmlight_file()] --> B[参数统一检查 & CSR 排序] B --> C[_dump_svmlight() 写头部] C --> D[_dump_svmlight_file (Cython)] D --> E[循环遍历每行:稠密或稀疏路径] E --> F[调用 get_dense_row_string / get_sparse_row_string] F --> G[构造标签字符串(单标签或多标签)] G --> H[写入文件行] style D fill:#bbf,stroke:#333,stroke-width:2px

45.10 数据模块自测与演示 —— 验证加载的即时反馈机制

sklearn/datasets/_base.py 作为脚本直接执行时(python -m sklearn.datasets._base),以下 __main__ 块会被触发:

if __name__ == "__main__":
    # 依次演示经典数据集的加载状态
    for loader in (load_iris, load_wine, load_breast_cancer,
                   load_digits, load_sample_images):
        try:
            data = loader()
            print(f"{loader.__name__}:")
            # 根据返回类型打印关键属性
            if hasattr(data, "data"):
                print(f"  data.shape = {data.data.shape}")
                print(f"  target.shape = {data.target.shape}")
                if hasattr(data, "target_names"):
                    print(f"  target_names = {data.target_names}")
            if hasattr(data, "images"):
                print(f"  images count = {len(data.images)}")
                print(f"  first image shape = {data.images[0].shape}")
        except ImportError as exc:
            # 例如缺少 Pillow 时给出友好提示,继续其它演示
            print(f"Skipping {loader.__name__} due to missing dependency: {exc}")

这相当于 智能仓储演示厅的现场巡检:只要模块能成功跑通,用户即可即时看到每个 “商品” 的基本信息,确认数据完整性。

45.10.1 流程图(Mermaid)

flowchart TD A[模块作为脚本运行] --> B[遍历所有加载函数] B --> C[调用加载函数获取 Bunch] C --> D[打印 data.shape、target.shape、target_names 等信息] D --> E[捕获 ImportError 并打印友好提示] style B fill:#def,stroke:#333,stroke-width:2px

45.11 设计取舍分析

在 scikit‑learn 的设计哲学中,兼容性优先、依赖最小化 是核心价值。SVMLight、CSV 与纯文本是几乎所有机器学习工具链(LibSVM、Vowpal Wabbit、Spark MLlib)都能直接读取的“通用语言”。虽然列式存储(Parquet、Feather 等)在压缩率和列裁剪上更具优势,但它们需要额外的重量级依赖(pyarrowfastparquet),违背了 scikit‑learn “轻量、零配置” 的原则。此外,文本格式天然可读,教学与调试时可以直接 cathead 查看,对新人更加友好。

因此,scikit‑learn 采用了 “兼容性优先、性能适中” 的权衡:通过 Cython 零拷贝、原子下载、流式解压等手段提升性能,同时保持文件格式的极低门槛。这使得本库在科研、教学、快速原型阶段拥有无可替代的易用性。

收益:跨语言兼容、透明文件结构、极低部署门槛。

代价:存储空间相对更大、在极端大规模稀疏数据上仍受 I/O 与 CPU 解析瓶颈限制、缺少列裁剪与自动压缩等企业级特性。

整体来看,这一取舍正好匹配 scikit‑learn 的目标用户群(科研、教学、快速实验),在该场景下提供了最佳的使用体验。


45.12 小结

下面的表格总结了本章涉及的关键概念,帮助你快速回顾。

| 概念 | 解释 |

|------|------|

| __main__ | 当 _base.py 被直接运行时,演示加载若干玩具数据集并打印其形状、目标、描述等信息,提供零配置的即时自检手段。 |

| load_svmlight_file | 解析单个 SVMLight/LibSVM 文本文件,返回稀疏 CSR 矩阵 X、目标向量 y 与可选 query_id,底层依赖 Cython _load_svmlight_file 实现零拷贝、无 GIL 的高速解析。 |

| load_svmlight_files | 批量加载多个 SVMLight 文件,自动统一特征维度并返回 [X1, y1, …](若 query_id=True 则附加 q),确保训练/测试集特征数一致。 |

| _load_svmlight_file (Cython) | 逐行流式读取,跳过注释、解析 qid、检测索引递增、使用 array.resize_smart 动态扩容,最终通过 np.frombuffer 零拷贝构建 CSR。 |

| get_dense_row_string / get_sparse_row_string | 两个 Cython 级行序列化工具:稠密路径预分配缓冲区,稀疏路径直接遍历 CSR 切片,生成 index:value 片段供 _dump_svmlight_file 写入。 |

| _dump_svmlight_file | 根据稠密/稀疏分支、是否多标签、是否含 query_id 生成完整的 SVMLight 行文本,并写入文件,全部在 Cython 中完成,保持极低的 Python 开销。 |

| fetch_file / _derive_folder_and_filename_from_url / _filter_filename | 将 URL → 本地目录映射、创建目录、下载调用统一封装,提供安全、原子、可重试的文件抓取流程。 |

| array.resize_smart | 在 Cython 中动态扩容 array.array,相比 Python list.append 能够保持连续内存布局并避免频繁的内存重新分配,从而提升大文件解析的性能。 |

通过上述细致的源码剖析,你已经掌握了 scikit‑learn 数据模块从 缓存目录远程抓取本地加载稀疏格式 全链路的实现细节。接下来,你可以自行在本地实验 load_svmlight_filedump_svmlight_file 的往返转换,体会零拷贝与流式处理带来的性能优势。


45.13 练习

  1. 阅读并调试 _load_svmlight_file 的 Cython 源码,尝试在不改动逻辑的前提下加入对 # 注释后空格的宽容处理。

  2. 对比 load_svmlight_fileload_svmlight_files 在特征对齐上的实现差异,解释为何多文件加载必须统一特征维度。

  3. 运行 python -m sklearn.datasets._base,观察各数据集的打印信息,体会 __main__ 块的自检作用。

祝你在源码探索的旅程中收获满满!

45.14 生活类比

想象 scikit-learn 的数据模块是一家智能仓储公司的演示厅: 当客户走进演示厅(直接运行模块),会看到自动播放的产品介绍(__main__ 块) 演示厅会展示各类数据“样本间”(如加载 Iris、Digits 等数据集),验证数据是否完整可用 这就像智能仓储公司设有展示区,客户无需下载完整库存,就能现场体验数据格式、结基本特征 演示内容包括数据形状、特征名称、目标分布等,帮助用户快速判断该数据集是否符合需求 SVMLight 格式解析器则像一台高速自动分拣机: 传送带上的包裹(文本行)流经扫描头(Cython 解析内核 _load_svmlight_file) 扫描头瞬间识别标签(target)、特征索引与数值,直接装入稀疏矩阵(CSR)货箱 get_dense_row_stringget_sparse_row_string 是反向打包线,将内存中的矩阵还原为标准格式文本 query_idmultilabel 如同包裹上的特殊标记,支持学习排序与多标签等高级业务场景 整个过程无 GIL、零拷贝、流式处理,如同分拣机全天候不间断高速运转

第 46 章 —— 经典玩具数据集 —— 把玩“机器学习入门标本”

46.1 学习目标

  • 难度:★★★☆☆(3/5)

  • 预备知识:Python 基础、面向对象编程与 Markdown/代码阅读基础

  • 理解 scikit‑learn 数据集模块的缓存目录管理机制与环境变量配置

  • 掌握远程文件下载的原子写入、SHA256 校验与重试机制

  • 熟悉 CSV、GZIP 压缩数据及 RST 说明文件的本地加载流程

  • 了解玩具数据集(分类、回归、图像)的加载逻辑与 Bunch 封装过程

  • 理解 load_files 目录树加载器的文本分类数据组织方式

  • 掌握 fetch_file 等通用工具函数实现自定义数据获取

  • 熟悉实用工具函数:DataFrame 转换、pickle 路径兼容、模块级常量


46.2 生活类比 —— 数据物流配送中心

把 scikit‑learn 的数据集模块想象成一个 智能物流配送中心,它在每一次实验中负责把数据“快递”到你的代码手中:首先是仓库选址与清场,get_data_home() 决定仓库位置并按需创建目录,clear_data_home() 可一次性清空缓存;接着是国际快递签收环节,_fetch_remote()fetch_file() 负责原子写入、SHA256 验单以及指数退避重试,确保数据完整可靠;随后是标准化拆箱作业,load_csv_data() 负责逐行解析 CSV 首行元信息,load_gzip_compressed_csv_data() 直接解压 GZIP 后交给 np.loadtxt 一次性加载;说明书与样品展示则由 load_descr() 读取 RST 文档、load_sample_images() 通过 PIL 解码示例图片完成;标准样品箱对应各玩具数据集加载器,如 load_wine()load_iris() 等,它们复用底层装载器并统一封装为 Bunch 返回;批量入库的分拣员是 load_files(),它把子目录当作类别,支持扩展名过滤、编码指定以及随机打乱;最后是格式转换工位,_convert_data_dataframe() 零拷贝把 NumPy 原料转为 Pandas 成品,实现数据在不同库间的无缝流转。这样一套物流系统让我们在实验室里像在超市挑选商品一样,快速、可靠地获得各种数据标本。


46.3 代码地图

src/sklearn/datasets/_base.py
├── 缓存目录管理
│   ├── get_data_home()                    # 获取/创建数据根目录(支持环境变量 SCIKIT_LEARN_DATA)
│   └── clear_data_home()                  # 递归删除缓存目录
├── 远程文件获取与校验
│   ├── _fetch_remote()                    # 原子下载 + SHA256 校验 + 重试
│   ├── _sha256()                          # 计算文件 SHA256
│   ├── _derive_folder_and_filename_from_url() # 从 URL 推导安全本地路径
│   ├── _filter_filename()                 # 字符串安全化为合法文件名
│   └── fetch_file()                       # 高层 API,封装 URL 解析与下载
├── 本地数据装载器
│   ├── load_csv_data()                    # 解析普通 CSV(首行携带元信息)
│   ├── load_gzip_compressed_csv_data()    # 解析 GZIP 压缩 CSV(np.loadtxt)
│   ├── load_descr()                       # 读取 RST 说明文档
│   ├── load_sample_images()               # 加载 JPG 示例图片(依赖 PIL)
│   └── load_sample_image()                # 按名称加载单张示例图片
├── 经典玩具数据集加载器
│   ├── load_wine()
│   ├── load_iris()
│   ├── load_breast_cancer()
│   ├── load_digits()
│   ├── load_diabetes()
│   └── load_linnerud()
├── 文件夹数据集加载器
│   └── load_files()                       # 从目录树加载文本文件(子目录 = 类别)
├── 实用工具函数
│   ├── _convert_data_dataframe()
│   ├── _pkl_filepath()
│   └── __main__                           # 常量定义:DATA_MODULE / DESCR_MODULE / IMAGES_MODULE / RemoteFileMetadata

46.4 数据主目录管理 —— 仓库选址与清场

46.4.1 核心代码(src/sklearn/datasets/_base.py

def get_data_home(data_home=None) -> str:
    """Return the path of the scikit-learn data directory."""
    if data_home is None:
        # 环境变量优先,其次默认路径
        data_home = environ.get("SCIKIT_LEARN_DATA", join("~", "scikit_learn_data"))
    # ~ 展开为用户主目录
    data_home = expanduser(data_home)
    # 自动创建目录,exist_ok 防止重复创建报错
    makedirs(data_home, exist_ok=True)
    return data_home

这段实现首先检查 SCIKIT_LEARN_DATA 环境变量。如果未定义,则使用默认的 ~/scikit_learn_data。随后使用 expanduser 展开 ~,并通过 makedirs(..., exist_ok=True) 确保目录已经创建好,最后返回绝对路径。整个过程相当于 为物流中心预先划定并准备好仓库

def clear_data_home(data_home=None):
    """Delete all the content of the data home cache."""
    # 统一获取根目录(同 get_data_home 的解析逻辑)
    data_home = get_data_home(data_home)
    # 递归删除目录及其所有子文件
    shutil.rmtree(data_home)

clear_data_home 只需要调用一次 get_data_home 取得根目录,然后使用 shutil.rmtree 递归删除整个目录,实现“一键清空仓库”。当磁盘空间紧张或需要强制重新下载时,这一步骤尤为重要。

46.4.2 流程图

graph TD A[用户调用 get_data_home()] --> B{环境变量 SCIKIT_LEARN_DATA} B -->|有| C[返回环境变量路径] B -->|无| D[使用默认 ~/scikit_learn_data] C --> E[expanduser & makedirs] D --> E E --> F[返回最终目录路径]

46.5 远程文件抓取与校验 —— 签收快递的标准动作

46.5.1 关键概念回顾(类比继续)

  • 原子写入:下载过程先写入唯一的临时文件(.part_XXXXX),成功后使用 shutil.move 替换为目标文件,防止并发写冲突。

  • SHA256 校验:在下载前后分别检查文件的 SHA256 哈希,确保完整性。

  • 指数退避重试:捕获网络异常并在 n_retries 次内以 delay 秒间隔重试,提升鲁棒性。

  • 异常安全:除 KeyboardInterrupt 外的任意异常都会在 except 中手动删除临时文件,避免残留半成品。

46.5.2 核心实现(src/sklearn/datasets/_base.py

46.5.2.1 _fetch_remote()

def _fetch_remote(remote, dirname=None, n_retries=3, delay=1):
    """Helper function to download a remote dataset."""
    if dirname is None:
        folder_path = Path(".")
    else:
        folder_path = Path(dirname)

    file_path = folder_path / remote.filename

    # 1. 本地已有文件且 checksum 匹配 → 直接返回
    if file_path.exists():
        if remote.checksum is None:
            return file_path
        checksum = _sha256(file_path)
        if checksum == remote.checksum:
            return file_path
        else:
            warnings.warn(
                f"SHA256 checksum of existing local file {file_path.name} "
                f"({checksum}) differs from expected ({remote.checksum}): "
                f"re-downloading from {remote.url} ."
            )

    # 2. 创建唯一临时文件(并发安全)
    temp_file = NamedTemporaryFile(
        prefix=remote.filename + ".part_", dir=folder_path, delete=False
    )
    temp_file.close()  # 保持空文件,防止 GC 删除
    try:
        temp_file_path = Path(temp_file.name)
        while True:
            try:
                urlretrieve(remote.url, temp_file_path)
                break
            except (URLError, TimeoutError):
                if n_retries == 0:
                    raise
                warnings.warn(f"Retry downloading from url: {remote.url}")
                n_retries -= 1
                time.sleep(delay)

        # 3. 下载完成后校验 SHA256
        checksum = _sha256(temp_file_path)
        if remote.checksum is not None and remote.checksum != checksum:
            raise OSError(
                f"The SHA256 checksum of {remote.filename} ({checksum}) "
                f"differs from expected ({remote.checksum})."
            )
    except (Exception, KeyboardInterrupt):
        # 4. 任何异常或手动中断 → 删除临时文件防止残留
        os.unlink(temp_file.name)
        raise

    # 5. 原子移动至目标路径
    shutil.move(temp_file_path, file_path)
    return file_path
  • 步骤 1:如果本地已有文件且 SHA256 校验通过,直接返回,避免不必要的网络请求。

  • 步骤 2:创建唯一的临时文件 *.part_,保证并发下载时不会产生冲突。

  • 步骤 3:内部循环实现重试逻辑;每次捕获 URLErrorTimeoutError,递减 n_retriessleep(delay)

  • 步骤 4:下载成功后使用 _sha256 再次校验;若不匹配抛出异常并删除临时文件。

  • 步骤 5:通过 shutil.move 完成原子写入,确保在同一文件系统上是原子的。

46.5.2.2 _sha256()

def _sha256(path):
    """Calculate the sha256 hash of the file at path."""
    sha256hash = hashlib.sha256()
    chunk_size = 8192
    with open(path, "rb") as f:
        while True:
            buffer = f.read(chunk_size)
            if not buffer:
                break
            sha256hash.update(buffer)
    return sha256hash.hexdigest()

函数一次读取 8 KB 块,保持内存常数,适用于大文件的校验。

46.5.2.3 _derive_folder_and_filename_from_url()

def _derive_folder_and_filename_from_url(url):
    parsed_url = urlparse(url)
    if not parsed_url.hostname:
        raise ValueError(f"Invalid URL: {url}")
    folder_components = [_filter_filename(parsed_url.hostname, filter_dots=False)]
    path = parsed_url.path

    if "/" in path:
        base_folder, raw_filename = path.rsplit("/", 1)
        base_folder = _filter_filename(base_folder)
        if base_folder:
            folder_components.append(base_folder)
    else:
        raw_filename = path

    filename = _filter_filename(raw_filename, filter_dots=False)
    if not filename:
        filename = "downloaded_file"
    return "/".join(folder_components), filename

该函数把 URL 分解为安全的本地目录层级和文件名,防止出现非法字符导致的文件系统错误。

46.5.2.4 fetch_file()

def fetch_file(url, folder=None, local_filename=None, sha256=None,
               n_retries=3, delay=1):
    """Fetch a file from the web if not already present in the local folder."""
    folder_from_url, filename_from_url = _derive_folder_and_filename_from_url(url)

    if local_filename is None:
        local_filename = filename_from_url

    if folder is None:
        folder = Path(get_data_home()) / folder_from_url
        makedirs(folder, exist_ok=True)

    remote_metadata = RemoteFileMetadata(
        filename=local_filename, url=url, checksum=sha256
    )
    return _fetch_remote(
        remote_metadata, dirname=folder, n_retries=n_retries, delay=delay
    )

fetch_file 把 URL → 本地目录映射、目录创建和 _fetch_remote 的调用全部封装,向用户提供了一个“一站式”下载入口。

46.5.3 流程图

graph TD A[调用 fetch_file(url)] --> B[_derive_folder_and_filename_from_url] B --> C[生成本地目录 & 文件名] C --> D[_fetch_remote] D -->|文件已存在且 checksum 匹配| E[直接返回本地路径] D -->|需要下载| F[创建临时文件 .part_] F --> G[循环 urlretrieve + 重试] G --> H[_sha256 校验临时文件] H -->|checksum 不匹配| I[抛异常并删除临时文件] H -->|checksum 匹配| J[shutil.move 原子替换] J --> K[返回最终文件路径]

46.6 CSV 与压缩数据的本地装载 —— 拆箱预制件

46.6.1 load_csv_data()

def load_csv_data(
    data_file_name,
    *,
    data_module=DATA_MODULE,
    descr_file_name=None,
    descr_module=DESCR_MODULE,
    encoding="utf-8",
):
    """Loads `data_file_name` from `data_module` with `importlib.resources`."""
    data_path = resources.files(data_module) / data_file_name
    with data_path.open("r", encoding="utf-8") as csv_file:
        data_file = csv.reader(csv_file)
        # 首行约定: n_samples n_features target_name1 target_name2 ...
        temp = next(data_file)
        n_samples = int(temp[0])
        n_features = int(temp[1])
        target_names = np.array(temp[2:])
        data = np.empty((n_samples, n_features))
        target = np.empty((n_samples,), dtype=int)

        for i, ir in enumerate(data_file):
            data[i] = np.asarray(ir[:-1], dtype=np.float64)
            target[i] = np.asarray(ir[-1], dtype=int)

    if descr_file_name is None:
        return data, target, target_names
    else:
        descr = load_descr(descr_module=descr_module,
                           descr_file_name=descr_file_name)
        return data, target, target_names, descr
posted @ 2026-09-04 04:07  绝不原创的飞龙  阅读(2)  评论(0)    收藏  举报