Sklearn-源码解析-书-v1-0-六-

Sklearn 源码解析(书)v1.0(六)

| 二值化 | 无 | 无 | binarize 阈值 | 无(要求整数编码) |

| 关键参数 | alpha, force_alpha | alpha, norm | alpha, binarize | alpha, min_categories |


13.8 测试辅助函数与综合验证 —— 构建“朴素贝叶斯测试工坊”

13.8.1 get_random_normal_x_binary_y:高斯数据的“标准生成器”

# 第 13 章 —— sklearn/tests/test_naive_bayes.py (第28-34行)
def get_random_normal_x_binary_y(global_random_seed):
    rng = np.random.RandomState(global_random_seed)
    X1 = rng.normal(size=(10, 3))
    y1 = (rng.normal(size=10) > 0).astype(int)
    return X1, y1

利用 np.random.RandomState(global_random_seed) 生成可复现的 (10, 3) 正态分布特征矩阵,目标向量通过正态分布采样后取符号得到二元标签 {0, 1}。为 test_gnb_priortest_gnb_sample_weight 提供标准化的二分类测试数据。


13.8.2 get_random_integer_x_three_classes_y:离散数据的“工厂流水线”

# 第 13 章 —— sklearn/tests/test_naive_bayes.py (第37-43行)
def get_random_integer_x_three_classes_y(global_random_seed):
    rng = np.random.RandomState(global_random_seed)
    X2 = rng.randint(5, size=(6, 100))
    y2 = np.array([1, 1, 2, 2, 3, 3])
    return X2, y2

生成 (6, 100) 形状的随机整数特征矩阵,取值范围 [0, 5),固定目标向量 [1,1,2,2,3,3] 恰好覆盖三个类别。被 test_discretenb_priortest_mnnbtest_categoricalnb 等大量离散 NB 测试复用。


13.8.3 测试中的随机种子传递机制

  • global_random_seed 作为 fixture 参数由 pytest 注入,确保每次运行可复现。

  • 内部显式创建 np.random.RandomState 而非直接使用全局随机状态,隔离不同测试的随机流。

  • 这种设计保证了 test_gnb_sample_weight 中“重复样本等价于样本权重”验证的确定性。


13.9 设计中的取舍

13.9.1 为什么不用单一的 _joint_log_likelihood 实现覆盖所有变体?

朴素贝叶斯的“朴素”在于条件独立假设,但不同变体对特征分布的建模截然不同:高斯分布用均值方差、多项式用计数向量、伯努利用二元概率、类别用逐特征分布、补集用反向计数。将这些差异封装在抽象方法中,既保证了预测流程(predict_log_proba 等)的复用,又给予子类最大的实现自由度。这是“模板方法模式”的典型应用。

13.9.2 为什么 GaussianNBvar_smoothing 而不是直接加在方差上?

var_smoothing 乘以最大特征方差,而非固定常数。这种相对缩放策略使得平滑量随数据量纲自适应:若特征方差很大(如 1e6),平滑量也随之变大;若特征方差很小(如 1e-6),平滑量也很小。固定常数会导致在大尺度特征上平滑不足、小尺度特征上过度平滑。

13.9.3 为什么 ComplementNB 要维护 feature_all_

补集计数 \(N_{\neg c, f} = \sum_{k \ne c} N_{k,f} = N_{all,f} - N_{c,f}\)。维护 feature_all_ 避免每次 _update_feature_log_prob 重复求和,体现了“计数即充分统计量”的工程思想:拟合时多存一点,预测时少算一点。

13.9.4 为什么 BernoulliNB 的平滑分母是 class_count_ + 2*alpha

伯努利模型下每个特征每类是二元分布(出现/不出现),等价于两个类别的多项分布。拉普拉斯平滑给每个类别加 \(\alpha\),两个类别共加 \(2\alpha\)。这与 MultinomialNBn_features * alpha 本质不同:前者是每特征两个结果,后者是每样本 n 个特征

13.9.5 为什么 CategoricalNB 要用列表存储 category_count_feature_log_prob_

不同特征的类别数不同(特征 A 有 3 个类别,特征 B 有 100 个),无法用单一矩阵存储。列表结构 [array(n_classes, n_cat_i)] 自然适配这种“锯齿状”数据,同时也使得逐特征独立平滑、高级索引预测成为可能。


13.10 动手练习

13.10.1 练习 1:追踪概率推理全流程

阅读 sklearn/naive_bayes.py 第81-177行,理解 _BaseNB 的四个预测方法:

  1. predict_joint_log_proba() - 第81-108行

  2. predict() - 第110-133行

  3. predict_log_proba() - 第135-160行

  4. predict_proba() - 第162-177行

回答问题:

  • 为什么 predict_log_proba 需要使用 _logsumexp?直接相减会有什么问题?

  • xpx.atleast_nd(log_prob_x, ndim=2).T 这行代码的作用是什么?

  • classes_ 包含字符串标签时,predict() 为何需要 _convert_to_numpy

13.10.2 练习 2:验证在线更新算法的数学正确性

阅读 sklearn/naive_bayes.py 第260-318行的 _update_mean_variance()

  1. 手工推导合并新旧数据的均值公式

  2. 理解 total_ssd 中修正项 (n_new*n_past/n_total)*(mu-new_mu)^2 的数学含义

  3. 设计一个实验:用 np.var 验证在线更新的方差与全量计算的方差是否一致

结合 test_gnb_check_update_with_no_data()test_gnb_sample_weight() 理解空数据和加权场景:

  • n_past==0 时,函数直接返回 new_munew_var,这有什么特殊意义?

  • sample_weight 为 None 时,为什么 new_var 使用 xp.var(X, axis=0)

  • 样本权重为 0 时,函数如何避免除零错误?

13.10.3 练习 3:对比四种离散朴素贝叶斯的平滑差异

阅读以下方法的实现:

  1. MultinomialNB._update_feature_log_prob() - 第759-766行

  2. ComplementNB._update_feature_log_prob() - 第864-876行

  3. BernoulliNB._update_feature_log_prob() - 第1014-1021行

  4. CategoricalNB._update_feature_log_prob() - 第1194-1200行

结合 test_bnb()test_cnb()test_categoricalnb() 验证手工计算结果:

  • 为什么 BernoulliNB 的分母是 class_count_+2αMultinomialNB 是总计数+α*n_features?

  • ComplementNBnorm=Truenorm=False 分别产生什么效果?

  • CategoricalNB 为何对每个特征独立做平滑?与 MultinomialNB 的全局平滑有何本质区别?

13.10.4 练习 4:稀疏矩阵与 Array API 兼容性探索

  1. 阅读 test_mnnb()(第330-390行),理解稀疏与稠密输入如何产生相同结果

  2. 阅读 test_gnb_array_api_compliance()(第598-650行),理解不同 Array API 后端的验证逻辑

  3. 阅读 test_gnb_naive_bayes_scale_invariance(),理解数据缩放对高斯 NB 的影响

回答问题:

  • safe_sparse_dot 在处理 CSR 矩阵时与普通 np.dot 有什么不同?

  • GaussianNB 为何需要显式设置 tags.array_api_support = True

  • 在 Array API 测试中,为何要检查 device()dtype 的一致性?

  • 为什么高斯 NB 具有尺度不变性,而 BernoulliNBCategoricalNB 不具有?

13.10.5 练习 5:增量学习的边界场景处理

阅读以下测试用例,理解增量学习的边界情况:

  1. test_gnb_partial_fit() - 第205-225行

  2. test_discretenb_degenerate_one_class_case() - 第390-440行

  3. test_mnb_prior_unobserved_targets() - 第440-470行

  4. test_discretenb_provide_prior_with_partial_fit() - 第472-490行

回答问题:

  • partial_fity 包含初始 classes 中不存在的标签时,会发生什么?

  • 单类别训练集(退化解)下,离散 NB 如何处理属性形状?

  • 未观测类别的先验平滑如何避免 RuntimeWarning

  • 将 X 分成两半分别 partial_fit 与一次 fit,结果是否完全一致?为什么?


13.11 本章小结

这一章中我们学习了朴素贝叶斯分类器家族的完整源码实现。首先我们剖析了 _BaseNB 抽象基类如何定义统一的预测流水线:联合对数似然 → log-sum-exp 归一化 → 指数变换,这是所有变体共享的“推理骨架”。其次我们深入了 GaussianNB 的在线学习核心——Chan-Golub-LeVeque 算法,理解如何仅用均值、方差、样本数三个充分统计量实现精确的增量更新。接着我们探索了 _BaseDiscreteNB 为四种离散变体提供的共享流水线:标签二值化、计数累积、三类先验策略调度、平滑参数安全阀。然后我们对比分析了四种离散变体的独特设计:MultinomialNB 的拉普拉斯平滑、ComplementNB 的补集纠偏策略、BernoulliNB 的 neg_prob 计算技巧、CategoricalNB 的逐特征独立计数与动态扩展。最后我们了解了测试工坊中的标准化数据生成器与随机种子隔离机制,以及 Array API 兼容性验证的工程实践。

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

| 概念 | 解释 |

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

| _BaseNB.predict_proba() | 统一预测入口:联合对数似然 → log-sum-exp归一化 → 指数变换 |

| GaussianNB._update_mean_variance() | Chan-Golub-LeVeque在线算法,用充分统计量合并新旧数据 |

| GaussianNB._partial_fit() | 首次调用初始化theta_/var_,逐类别更新统计量,支持先验覆盖 |

| GaussianNB.fit() | 通过validate_data验证y后调用_partial_fit完成拟合 |

| GaussianNB.__sklearn_tags__() | 声明array_api_support=True,支持跨后端数组计算 |

| _BaseDiscreteNB.fit() | LabelBinarizer二值化标签,二分类特殊拼接,加权计数 |

| _BaseDiscreteNB.partial_fit() | label_binarize处理增量拟合,支持未观测类别的先验平滑 |

| _BaseDiscreteNB._update_class_log_prior() | 三类先验策略:用户指定、经验估计、均匀分布 |

| _BaseDiscreteNB._check_alpha() | 平滑参数安全阀,过小自动提升至1e-10或由force_alpha强制 |

| _BaseDiscreteNB._init_counters() | 初始化class_count_和feature_count_为零矩阵 |

| _BaseDiscreteNB._check_X_y() | fit方法中的输入验证,接受CSR稀疏矩阵 |

| MultinomialNB.__init__() | 调用基类初始化alpha、force_alpha、fit_prior、class_prior |

| MultinomialNB._count() | safe_sparse_dot高效计算(class, feature)加权计数 |

| MultinomialNB._update_feature_log_prob() | 拉普拉斯平滑:log(feature_count+α) - log(总计数+α*n_features) |

| MultinomialNB._joint_log_likelihood() | safe_sparse_dot(X, feature_log_prob_.T) + class_log_prior_ |

| MultinomialNB.__sklearn_tags__() | 标记positive_only=True,声明输入非负约束 |

| ComplementNB.__init__() | 额外norm参数控制权重是否L1归一化 |

| ComplementNB.__sklearn_tags__() | 标记positive_only=True |

| ComplementNB._count() | safe_sparse_dot计算feature_count_,额外维护feature_all_ |

| ComplementNB._update_feature_log_prob() | 补集权重计算,用其他类别计数惩罚多数类,适合不平衡数据 |

| ComplementNB._joint_log_likelihood() | 单类时仅加class_log_prior_,其余情况返回补集评分 |

| BernoulliNB.__init__() | 额外binarize参数控制二值化阈值 |

| BernoulliNB._check_X() | 预测时对输入执行binarize二值化处理 |

| BernoulliNB._check_X_y() | 拟合时对X和y执行输入验证与二值化 |

| BernoulliNB._count() | safe_sparse_dot计算二元特征计数 |

| BernoulliNB._update_feature_log_prob() | 二元平滑公式:分母为class_count+2α |

| BernoulliNB._joint_log_likelihood() | neg_prob技巧避免重复计算log(1-p),二元特征专用 |

| CategoricalNB.__init__() | 额外min_categories参数控制最小类别数 |

| CategoricalNB.fit() | 继承_BaseDiscreteNB.fit,文档说明类别编码要求 |

| CategoricalNB.partial_fit() | 继承_BaseDiscreteNB.partial_fit,文档说明类别编码要求 |

| CategoricalNB.__sklearn_tags__() | 标记categorical=True,禁用sparse,positive_only=True |

| CategoricalNB._check_X() | 预测时强制整数dtype,禁用稀疏,检查非负 |

| CategoricalNB._check_X_y() | 拟合时强制整数dtype与非负检查 |

| CategoricalNB._init_counters() | 初始化category_count_列表,每特征持有(n_classes, 0)矩阵 |

| CategoricalNB._validate_n_categories() | max(X)+1与min_categories取最大值,确定每特征类别数 |

| CategoricalNB._count() | np.bincount逐特征计数,np.pad动态扩展类别维度 |

| CategoricalNB._update_feature_log_prob() | 逐特征独立拉普拉斯平滑,返回数组列表 |

| CategoricalNB._joint_log_likelihood() | 高级索引feature_log_prob_[i][:, indices].T高效累加 |

| get_random_normal_x_binary_y() | 测试辅助:生成服从正态分布的二分类特征矩阵与目标向量 |

| get_random_integer_x_three_classes_y() | 测试辅助:生成100维整数特征的三分类数据,固定类别分布 |

| test_predict_joint_proba() | 验证predict_joint_log_proba与logsumexp归一化后的一致性 |

| test_mnb_prior_unobserved_targets() | 验证未观测类别先验平滑避免RuntimeWarning,新增类后预测正确 |

| test_discretenb_predict_proba() | 验证伯努利与多项NB在二分类/多分类下的概率形状与求和为1 |

| test_gnb_naive_bayes_scale_invariance() | 验证数据缩放对高斯NB预测结果无影响 |

| test_mnnb() | 验证多项NB在稠密/稀疏输入下的正确性与增量拟合一致性 |

| test_gnb_prior_large_bias() | 验证严重偏置先验下高斯NB的预测行为 |

| test_gnb_check_update_with_no_data() | 验证空数据调用_update_mean_variance返回原统计量 |

| test_gnb_priors_sum_isclose() | 验证10类先验和接近1时的高斯NB拟合 |

| test_gnb_neg_priors() | 验证负数先验引发ValueError |

| test_gnb_priors() | 验证priors参数覆盖经验先验且预测概率正确 |

| test_gnb_wrong_nb_priors() | 验证先验数量与类别数不匹配的错误处理 |

| test_gnb_prior_greater_one() | 验证先验和大于1的错误处理 |

| test_gnb_prior() | 验证经验先验正确性与先验和为1 |

| test_gnb_sample_weight() | 验证样本权重与重复样本在拟合中的等价性 |

| test_gnb() | 验证高斯NB基本拟合、预测与概率一致性 |

| test_mnb_prior_unobserved_targets() | 验证未观测类别先验平滑避免RuntimeWarning |

| test_discretenb_uniform_prior() | 验证fit_prior=False时均匀先验 |

| test_bnb() | 伯努利NB教科书例子,验证特征概率手工计算 |

| test_bnb_feature_log_prob() | 手工验证BernoulliNB特征对数概率公式 |

| test_cnb() | 补集NB权重手工验证,含norm选项和负输入检查 |

| test_categoricalnb() | 验证CategoricalNB计数、预测、样本权重与负输入检查 |

| test_categoricalnb_with_min_categories() | 验证min_categories的int和list输入及类别数扩展 |

| test_categoricalnb_min_categories_errors() | 验证min_categories形状错误引发ValueError |

| test_check_alpha() | 验证_ALPHA_MIN阈值与force_alpha行为 |

| test_categorical_input_tag() | 验证CategoricalNB的categorical标签设置 |

| test_discretenb_provide_prior() | 验证用户指定先验在离散NB中的使用 |

| test_discretenb_provide_prior_with_partial_fit() | 验证partial_fit中用户先验与全量fit一致性 |

| test_discretenb_prior() | 验证离散NB经验先验对数概率 |

| test_discretenb_degenerate_one_class_case() | 验证单类别训练集下各属性的第一轴长度 |

| test_NB_partial_fit_no_first_classes() | 验证首次partial_fit缺少classes参数的错误 |

| test_discretenb_sample_weight_multiclass() | 验证离散NB多样本权重下的fit与partial_fit一致性 |

| test_discretenb_predict_proba() | 验证离散NB概率和为1与二分类/多分类场景 |

| test_discretenb_partial_fit() | 验证离散NB增量拟合与全量fit的计数一致性 |

| test_alpha() | 验证α=0时自动提升至1e-10,含稀疏矩阵场景 |

| test_alpha_vector() | 验证向量alpha支持、负值检查与维度匹配 |

| test_gnb_array_api_compliance() | 验证GaussianNB在不同Array API后端下的一致性 |

下一章中,我们将学习多类别与多输出策略,理解 OvR、OvO、ECOC 与输出链等分解策略如何把复杂任务拆解为多个二分类或单输出子问题。

第 14 章 —— 多类别与多输出策略 —— 组合“二分类器的合纵连横”

14.1 学习目标

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

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

  • 理解 OvR、OvO、ECOC 三种多分类策略的分解原理与适用场景

  • 掌握元估计器如何通过克隆、并行训练与条件方法暴露来组合二分类器

  • 深入 LabelBinarizer 在 OvR 中的标签矩阵转换机制,以及稀疏格式的性能优势

  • 理解多输出估计器(MultiOutput)与链式建模(ClassifierChain/RegressorChain)的设计差异

  • 能阅读并调试元估计器的预测路径,包括平局破解、概率归一化与逆序还原

14.2 生活类比

想象多分类策略是一家快递分拣中心

  • OvR(一对多) = 为每个目的地设立一条独立传送带,包裹经过每条传送带时被标记为“是/否发往该地”

  • OvO(一对一) = 举办两两对决的锦标赛,每个包裹在每场比赛中被分给两位选手之一,最终得票最多者胜出

  • ECOC(纠错输出码) = 为每个目的地分配唯一的二进制“邮政编码”,即使部分编码位在传输中出错,仍能通过最近码字纠错找回正确目的地

  • MultiOutput = 为包裹的每个属性(重量、体积、易碎度)分别安排独立的测量员

  • ClassifierChain = 接力赛跑,前一位测量员的结果直接传递给下一位作为额外线索

就像分拣中心经理需要根据包裹量和目的地数量选择最经济高效的分拣方案,元估计器也需要在训练成本、预测精度和模型可解释性之间权衡。

14.3 源码地图

sklearn/multiclass.py
├── 模块级辅助函数
│   ├── _fit_binary()                  # 训练单个二分类器,处理单类退化为 _ConstantPredictor
│   ├── _partial_fit_binary()         # 增量训练单个二分类器
│   ├── _predict_binary()             # 获取二分类器的连续置信度(decision_function 或 predict_proba)
│   ├── _threshold_for_binary_predict()  # 判定阈值:0.0(decision_function)或 0.5(predict_proba)
│   ├── _ConstantPredictor            # 单类别退化时的常量预测器
│   └── _estimators_has()             # 检查子估计器是否具备某方法(配合 available_if)
├── OneVsRestClassifier
│   ├── __init__()                    # 初始化估计器、并行度与日志级别
│   ├── fit()                         # LabelBinarizer 转换 + Parallel 并行训练 n_classes 个二分类器
│   ├── partial_fit()                 # 增量学习,首次调用锁定 classes,后续批次校验新类别
│   ├── predict()                     # 多分类走 argmax 平局破解,多标签走 CSC 稀疏构建
│   ├── predict_proba()               # 单标签行归一化,多标签返回边际概率
│   ├── decision_function()           # 堆叠各二分类器的决策函数输出
│   ├── multilabel_ / n_classes_      # 属性透传 label_binarizer_.y_type_
│   ├── __sklearn_tags__()            # 透传基估计器的 pairwise 与 sparse 标签
│   └── get_metadata_routing()        # 配置 fit/partial_fit 的元数据路由
├── OneVsOneClassifier
│   ├── __init__()                    # 初始化估计器与并行度
│   ├── _fit_ovo_binary()             # 样本筛选(logical_or)+ 标签重编码为 0/1 + indcond 记录
│   ├── _partial_fit_ovo_binary()     # 增量训练单个 OvO 二分类器
│   ├── fit()                         # 双重循环生成 C(n,2) 组合,Parallel 并行训练
│   ├── partial_fit()                 # itertools.combinations 与已有估计器 zip 配对增量更新
│   ├── predict()                     # 二分类特例走阈值判定,多分类走 argmax
│   ├── decision_function()           # 投票 + 置信度融合(_ovr_decision_function)
│   ├── n_classes_                    # 类别数属性
│   ├── __sklearn_tags__()            # 透传基估计器的 pairwise 与 sparse 标签
│   └── get_metadata_routing()        # 配置 fit/partial_fit 的元数据路由
├── OutputCodeClassifier
│   ├── __init__()                    # 初始化估计器、码本大小、随机状态与并行度
│   ├── fit()                         # 随机码本生成(uniform + 0.5 阈值)+ 按列训练
│   ├── predict()                     # F-contiguous 构造 + pairwise_distances_argmin 最近码字搜索
│   ├── get_metadata_routing()        # 配置 fit 的元数据路由
│   └── __sklearn_tags__()            # 透传基估计器的 sparse 标签
├── sklearn/multioutput.py
├── _fit_estimator()                  # 克隆单个估计器并拟合(含样本权重)
├── _partial_fit_estimator()          # 增量训练单个估计器
├── _available_if_estimator_has()      # 为 MultiOutput 提供条件方法暴露
├── _available_if_base_estimator_has()  # 为 Chain 提供条件方法暴露
├── _MultiOutputEstimator
│   ├── __init__()                    # 初始化估计器与并行度
│   ├── fit()                         # 逐列克隆 + Parallel 并行拟合 + y 维度校验
│   ├── partial_fit()                 # 首次调用克隆,后续批次持续累积
│   ├── predict()                     # 并行预测后 np.asarray(y).T 转置
│   ├── get_metadata_routing()        # 配置 fit/partial_fit 的元数据路由
│   └── __sklearn_tags__()            # 强制 multi_output=True, single_output=False
├── MultiOutputRegressor
│   ├── __init__()                    # 调用基类初始化
│   └── partial_fit()                 # 委托给基类实现
├── MultiOutputClassifier
│   ├── __init__()                    # 调用基类初始化
│   ├── fit()                         # 调用基类 fit 并收集各类别的 classes_
│   ├── _check_predict_proba()         # 双态探测:已拟合查 estimators_,未拟合查 estimator
│   ├── predict_proba()               # 返回每个输出估计器的概率估计列表
│   ├── score()                       # 严格全匹配:np.all(y == y_pred, axis=1)
│   └── __sklearn_tags__()            # 跳过某些通用测试
├── _BaseChain
│   ├── __init__()                    # 初始化链参数(含 base_estimator 弃用处理)
│   ├── _get_estimator()               # 校验并获取 estimator/base_estimator
│   ├── _log_message()                # verbose 模式下打印链进度
│   ├── fit()                         # order 解析(None/数组/'random')+ cv 模式 OOF 预测填充
│   ├── _get_predictions()            # 链路逐级预测 + inv_order 逆序还原
│   ├── predict()                     # 使用 _get_predictions 输出硬预测
│   └── __sklearn_tags__()            # 稀疏支持从基估计器继承
├── ClassifierChain
│   ├── __init__()                    # 初始化链参数与 chain_method
│   ├── fit()                         # 存储 chain_method_ 响应方法名并收集 classes_
│   ├── predict_proba()               # 经 _get_predictions 输出概率
│   ├── predict_log_proba()           # 对概率取对数
│   ├── decision_function()           # 经 _get_predictions 输出决策函数
│   ├── get_metadata_routing()        # 配置 fit 的元数据路由
│   └── __sklearn_tags__()            # 标记 multi_output=True
└── RegressorChain
    ├── __init__()                    # 调用基类初始化
    ├── fit()                         # 固定使用 predict 作为链路方法
    ├── get_metadata_routing()        # 配置 fit 的元数据路由
    └── __sklearn_tags__()            # 标记 multi_output=True

14.4 OvR 核心拆解 —— 一对多策略的“流水线分拣机”

14.4.1 为什么需要 _fit_binary 辅助函数?

在 One-vs-Rest 策略中,每个二分类子估计器只负责区分“当前类别 vs 其余所有类别”。这种分解方式天然要求每个子估计器独立训练、互不干扰。_fit_binary 函数正是这一隔离逻辑的封装:它接收基估计器、特征矩阵 X、单列二值标签 y 以及拟合参数,返回训练好的二分类器。关键点在于单类退化处理:当某列标签全部相同时(例如所有样本都属于“非当前类”),无法训练普通分类器,此时函数会退化为 _ConstantPredictor 并发出警告。通过 clone 保证子估计器之间状态隔离,互不干扰。

14.4.2 _ConstantPredictor 的常量预测语义

当某个类别在训练数据中完全缺失或全为正样本时,OvR 需要一个不会报错的“占位符”预测器。_ConstantPredictor 正是这样一个极简实现:

  • predictdecision_function 输出常量标签

  • predict_proba 返回 one-hot 概率(如 [[1, 0]][[0, 1]]

  • 使用 ensure_all_finite=False 跳过有限性检查,兼容稀疏输入

这段代码定义了单类退化场景下的“兜底预测器”,保证 OvR 在极端标签分布下的鲁棒性。

14.4.3 LabelBinarizer 在 OvR 中的关键角色

LabelBinarizer(sparse_output=True) 是 OvR 将多分类标签转为二分类目标的核心工具。它把形状为 (n_samples,) 的多分类标签 y 转换为 CSC 稀疏指示矩阵 Y,每列对应一个类别的二分类目标(1 表示属于该类,0 表示不属于)。sparse_output=True 在多类别场景下比稠密表示更省内存且性能更优。label_binarizer_.y_type_ 决定后续走多分类还是多标签预测路径。

14.4.4 并行训练 n_classes 个二分类器

OvR 的 fit 方法使用 Parallel 将各列分发给 _fit_binary,支持 n_jobs 控制并行度。代码注释揭示了一个工程细节:当基估计器训练极快时,n_jobs>1 可能因线程开销反而变慢(joblib issue #112)。这就是为什么默认 n_jobs=None 而非 -1 的原因。

# 第 14 章 —— sklearn/multiclass.py - OneVsRestClassifier.fit() (第327-365行)
@_fit_context(
    # OneVsRestClassifier.estimator is not validated yet
    prefer_skip_nested_validation=False
)
def fit(self, X, y, **fit_params):
    """Fit underlying estimators.

    Parameters
    ----------
    X : {array-like, sparse matrix} of shape (n_samples, n_features)
        Data.

    y : {array-like, sparse matrix} of shape (n_samples,) or (n_samples, n_classes)
        Multi-class targets. An indicator matrix turns on multilabel
        classification.

    **fit_params : dict
        Parameters passed to the ``estimator.fit`` method of each
        sub-estimator.

        .. versionadded:: 1.4
            Only available if `enable_metadata_routing=True`. See
            :ref:`Metadata Routing User Guide <metadata_routing>` for more
            details.

    Returns
    -------
    self : object
        Instance of fitted estimator.
    """
    _raise_for_params(fit_params, self, "fit")

    routed_params = process_routing(
        self,
        "fit",
        **fit_params,
    )
    # A sparse LabelBinarizer, with sparse_output=True, has been shown to
    # outperform or match a dense label binarizer in all cases and has also
    # resulted in less or equal memory consumption in the fit_ovr function
    # overall.
    self.label_binarizer_ = LabelBinarizer(sparse_output=True)
    Y = self.label_binarizer_.fit_transform(y)
    Y = Y.tocsc()
    self.classes_ = self.label_binarizer_.classes_
    columns = (col.toarray().ravel() for col in Y.T)
    # In cases where individual estimators are very fast to train setting
    # n_jobs > 1 in can results in slower performance due to the overhead
    # of spawning threads.  See joblib issue #112.
    self.estimators_ = Parallel(n_jobs=self.n_jobs, verbose=self.verbose)(
        delayed(_fit_binary)(
            self.estimator,
            X,
            column,
            fit_params=routed_params.estimator.fit,
            classes=[
                "not %s" % self.label_binarizer_.classes_[i],
                self.label_binarizer_.classes_[i],
            ],
        )
        for i, column in enumerate(columns)
    )

    if hasattr(self.estimators_[0], "n_features_in_"):
        self.n_features_in_ = self.estimators_[0].n_features_in_
    if hasattr(self.estimators_[0], "feature_names_in_"):
        self.feature_names_in_ = self.estimators_[0].feature_names_in_

    return self

这段代码实现了 OvR 的核心训练流程:标签二值化、逐列并行训练、属性透传。

14.4.5 数据流图:OvR 训练与预测

graph TD A[原始标签 y] --> B[LabelBinarizer] B --> C[CSC 稀疏指示矩阵 Y] C --> D[逐列提取] D --> E[_fit_binary] E --> F[并行训练 n_classes 个二分类器] F --> G[estimators_ 列表] G --> H[预测阶段] H --> I{多分类/多标签?} I -->|多分类| J[反向遍历 argmax 平局破解] I -->|多标签| K[阈值判定 + CSC 稀疏构建] J --> L[输出类别标签] K --> L

14.5 OvR 预测路径 —— 从“投票”到“阈值”的双模式切换

14.5.1 多分类模式下的 np.argmax 平局破解技巧

OvR 的 predict 方法在多分类模式下采用反向遍历配合 np.maximum 来精确复现 np.argmax 的平局行为。这是一个极其精妙的实现细节:

# 第 14 章 —— sklearn/multiclass.py - OneVsRestClassifier.predict() (第440-472行)
def predict(self, X):
    """Predict multi-class targets using underlying estimators.

    Parameters
    ----------
    X : {array-like, sparse matrix} of shape (n_samples, n_features)
        Data.

    Returns
    -------
    y : {array-like, sparse matrix} of shape (n_samples,) or (n_samples, n_classes)
        Predicted multi-class targets.
    """
    check_is_fitted(self)

    n_samples = _num_samples(X)
    if self.label_binarizer_.y_type_ == "multiclass":
        maxima = np.empty(n_samples, dtype=float)
        maxima.fill(-np.inf)
        argmaxima = np.zeros(n_samples, dtype=int)
        n_classes = len(self.estimators_)
        # Iterate in reverse order to match np.argmax tie-breaking behavior
        for i, e in enumerate(reversed(self.estimators_)):
            pred = _predict_binary(e, X)
            np.maximum(maxima, pred, out=maxima)
            argmaxima[maxima == pred] = n_classes - i - 1
        return self.classes_[argmaxima]
    else:
        thresh = _threshold_for_binary_predict(self.estimators_[0])
        indices = array.array("i")
        indptr = array.array("i", [0])
        for e in self.estimators_:
            indices.extend(np.where(_predict_binary(e, X) > thresh)[0])
            indptr.append(len(indices))
        data = np.ones(len(indices), dtype=int)
        indicator = sp.csc_matrix(
            (data, indices, indptr), shape=(n_samples, len(self.estimators_))
        )
        return self.label_binarizer_.inverse_transform(indicator)

为什么要反向遍历? np.argmax 在遇到相等最大值时返回第一个出现的索引。反向遍历配合 argmaxima[maxima == pred] = n_classes - i - 1 赋值,保证当后面的类别(原顺序中靠前)与当前最大值相等时,索引会被覆盖为靠前的类别,从而与 np.argmax 行为完全一致。对应测试 test_ovr_ties 验证了这一非平凡的细节。

14.5.2 多标签模式下的稀疏 CSC 矩阵构建

多标签预测时,OvR 使用 array.array('i') 手动组装 indicesindptr,避免逐列修改稀疏矩阵的性能开销。阈值的确定来自 _threshold_for_binary_predict:有 decision_function 用 0.0,否则用 0.5。inverse_transform 将稀疏指示矩阵还原为多标签输出格式。

14.5.3 predict_proba 的单标签归一化逻辑

  • 多标签场景返回边际概率,不强制行和为 1

  • 单标签场景通过 row_sums != 0where 掩码安全归一化,避免除零警告

  • 全零概率的边界情况在 test_ovr_single_label_predict_proba_zero 中验证

# 第 14 章 —— sklearn/multiclass.py - OneVsRestClassifier.predict_proba() (第496-522行)
@available_if(_estimators_has("predict_proba"))
def predict_proba(self, X):
    """Probability estimates.

    The returned estimates for all classes are ordered by label of classes.

    Note that in the multilabel case, each sample can have any number of
    labels. This returns the marginal probability that the given sample has
    the label in question. For example, it is entirely consistent that two
    labels both have a 90% probability of applying to a given sample.

    In the single label multiclass case, the rows of the returned matrix
    sum to 1.

    Parameters
    ----------
    X : {array-like, sparse matrix} of shape (n_samples, n_features)
        Input data.

    Returns
    -------
    T : array-like of shape (n_samples, n_classes)
        Returns the probability of the sample for each class in the model,
        where classes are ordered as they are in `self.classes_`.
    """
    check_is_fitted(self)
    # Y[i, j] gives the probability that sample i has the label j.
    # In the multi-label case, these are not disjoint.
    Y = np.array([e.predict_proba(X)[:, 1] for e in self.estimators_]).T

    if len(self.estimators_) == 1:
        # Only one estimator, but we still want to return probabilities
        # for two classes.
        Y = np.concatenate(((1 - Y), Y), axis=1)

    if not self.multilabel_:
        # Then, (nonzero) sample probability distributions should be normalized.
        row_sums = np.sum(Y, axis=1)[:, np.newaxis]
        np.divide(Y, row_sums, out=Y, where=row_sums != 0)

    return Y

这段代码展示了 OvR 概率输出的核心逻辑:收集各二分类器的正类概率、单估计器特例处理、单标签安全归一化。

14.5.4 decision_function 的形状约定

  • 二分类时返回形状 (n_samples,),多分类时返回 (n_samples, n_classes)

  • 所有子估计器的输出通过 np.vstack 转置后堆叠

14.6 OvR 增量学习 —— partial_fit 的“分批投喂机制”

14.6.1 首次调用判定与标签锁定

OvR 的 partial_fit 通过 _check_partial_fit_first_call 检查 classes 参数是否缺失,首次调用必须提供。后续批次中若出现新类别,通过 np.setdiff1d 检测并抛 ValueError。每个子估计器在第一轮被克隆后持续累积训练,支持 Out-of-Core 学习。

14.6.2 _partial_fit_binary 的固定类别约束

# 第 14 章 —— sklearn/multiclass.py - _partial_fit_binary() (第71-75行)
def _partial_fit_binary(estimator, X, y, partial_fit_params):
    """Partially fit a single binary estimator."""
    estimator.partial_fit(X, y, classes=np.array((0, 1)), **partial_fit_params)
    return estimator

始终传递 classes=np.array((0, 1)),确保二分类器的类别顺序一致。

14.6.3 available_if 的条件方法暴露

_estimators_has("partial_fit") 动态检查基估计器能力。若基估计器不支持 partial_fit,OvR 实例上该属性直接不存在(hasattr 为 False)。测试 test_ovr_partial_fit 验证 SVC 基估计器不会暴露 partial_fit

14.6.4 test_multiclass_estimator_attribute_error 的错误链设计

访问缺失的 partial_fit 时抛出带有因果链的 AttributeError。外层消息定位元估计器,内层 __cause__ 揭示基估计器的真实原因。

14.7 OvR 多标签与稀疏数据处理 —— 处理“多标签世界的十字路口”

14.7.1 多标签指示矩阵的输入语义

y 是二维二元矩阵时,LabelBinarizer 识别为 multilabel-indicatormultilabel_ 属性透传 label_binarizer_.y_type_.startswith("multilabel")。每个标签列独立训练一个二分类器,互不干扰。

14.7.2 稀疏标签矩阵的预测等价性

测试 test_ovr_fit_predict_sparse 验证 CSR/CSC/COO/DOK/LIL 所有稀疏容器。稀疏标签训练与稠密标签训练产生相同的预测输出。predict 返回的稀疏矩阵通过 inverse_transform 还原为多标签格式。

14.7.3 回归器作为基估计器的兼容性

_predict_binary 检测到回归器时直接使用 predict 输出。test_ovr_ovo_regressor 验证 DecisionTreeRegressor 在 OvR/OvO 中正常工作。

14.7.4 Pipeline 与 GridSearchCV 的嵌套支持

test_ovr_pipeline 验证 Pipeline 作为基估计器时方法探测的正确性。test_ovr_gridsearch 验证 estimator__C 双下划线参数透传。

14.8 OvR 边界场景 —— 二分类、字符串标签与常量目标

14.8.1 二分类特例的形状降维

n_classes_ == 2decision_function 返回 (n_samples,) 而非 (n_samples, 2)predict_proba 输出 2 列概率,行和为 1。test_ovr_binary 验证 LinearSVC、LinearRegression、Ridge、ElasticNet、MultinomialNB、LogisticRegression。

14.8.2 字符串标签与标签指示矩阵输入

LabelBinarizer 自动处理字符串标签的编码与解码。test_ovr_multiclass 验证 set(clf.classes_) == {"ham", "eggs", "spam"}。标签指示矩阵输入与原始多分类标签产生一致预测。

14.8.3 常量目标与缺失类别

当某列标签全为 0 或全为 1 时,_fit_binary 发出 UserWarning。_ConstantPredictor 输出常量决策与 one-hot 概率。test_constant_int_target 验证常量标签不触发异常。

14.9 OvO 训练矩阵 —— 两两配对的“锦标赛赛程表”

14.9.1 _fit_ovo_binary 的样本筛选与重编码

# 第 14 章 —— sklearn/multiclass.py - _fit_ovo_binary() (第571-593行)
def _fit_ovo_binary(estimator, X, y, i, j, fit_params):
    """Fit a single binary estimator (one-vs-one)."""
    cond = np.logical_or(y == i, y == j)
    y = y[cond]
    y_binary = np.empty(y.shape, int)
    y_binary[y == i] = 0
    y_binary[y == j] = 1
    indcond = np.arange(_num_samples(X))[cond]

    fit_params_subset = _check_method_params(X, params=fit_params, indices=indcond)
    return (
        _fit_binary(
            estimator,
            _safe_split(estimator, X, None, indices=indcond)[0],
            y_binary,
            fit_params=fit_params_subset,
            classes=[i, j],
        ),
        indcond,
    )

np.logical_or(y == i, y == j) 选出仅属于类别 i 或 j 的样本子集。将标签重编码为 0/1 二元值,丢弃其余类别的样本。返回 (estimator, indcond) 二元组,indcond 记录样本索引供预计算核矩阵使用。_check_method_params 对 fit_params 做样本子集筛选,_safe_split 安全切片。

14.9.2 组合生成与并行调度

双重循环 for i in range(n_classes) for j in range(i + 1, n_classes) 生成 C(n,2) 个二分类任务。Parallel 并行执行后,用 zip(*...) 将估计器列表与索引列表分离。pairwise_indices_ 仅在基估计器 pairwise 标签为 True 时保留,避免不必要的内存开销。

14.9.3 预计算核矩阵的列裁剪语义

当使用 kernel='precomputed' 时,每个 OvO 子估计器只接收对应样本子集的行与列切片。test_pairwise_n_features_in 验证了裁剪后 n_features_in_ 从 149 降为 99 或 100 的细粒度行为。test_pairwise_indices 验证每个索引的长度与核矩阵行数的比例关系。

14.10 OvO 增量学习 —— 组合遍历与二分类特例的双重路径

14.10.1 _partial_fit_ovo_binary 的空批次处理

当某对类别在当前批次中无样本时,直接返回未更新的估计器。使用 np.zeros_like 默认 0 标签,避免为减少的样本子集重新分配。

14.10.2 itertools.combinations 与已有估计器的 zip 配对

combinations(range(self.n_classes_), 2) 生成 (i,j) 对的迭代器。与 self.estimators_ 列表 zip 后逐对增量更新。test_ovo_partial_fit_predict 验证少量类别批次与全量类别批次的一致性。

14.10.3 二分类特例的化简

n_classes_ == 2decision_function 返回 (n_samples,) 形状。predict 中同样以 _threshold_for_binary_predict 判定,与 sklearn 二分类约定对齐。test_ovo_consistent_binary_classification 验证 OvO 二分类与直接训练完全一致。

14.11 OvO 投票与决策融合 —— 矛盾裁决的“计票与加权会议”

14.11.1 _ovr_decision_function 的“投票 + 置信度”双重融合

# 第 14 章 —— sklearn/multiclass.py - OneVsOneClassifier.decision_function() (第762-800行)
def decision_function(self, X):
    """Decision function for the OneVsOneClassifier.

    The decision values for the samples are computed by adding the
    normalized sum of pair-wise classification confidence levels to the
    votes in order to disambiguate between the decision values when the
    votes for all the classes are equal leading to a tie.

    Parameters
    ----------
    X : array-like of shape (n_samples, n_features)
        Input data.

    Returns
    -------
    Y : array-like of shape (n_samples, n_classes) or (n_samples,)
        Result of calling `decision_function` on the final estimator.

        .. versionchanged:: 0.19
            output shape changed to ``(n_samples,)`` to conform to
            scikit-learn conventions for binary classification.
    """
    check_is_fitted(self)
    X = validate_data(
        self,
        X,
        accept_sparse=True,
        ensure_all_finite=False,
        reset=False,
    )

    indices = self.pairwise_indices_
    if indices is None:
        Xs = [X] * len(self.estimators_)
    else:
        Xs = [X[:, idx] for idx in indices]

    predictions = np.vstack(
        [est.predict(Xi) for est, Xi in zip(self.estimators_, Xs)]
    ).T
    confidences = np.vstack(
        [_predict_binary(est, Xi) for est, Xi in zip(self.estimators_, Xs)]
    ).T
    Y = _ovr_decision_function(predictions, confidences, len(self.classes_))
    if self.n_classes_ == 2:
        return Y[:, 1]
    return Y

收集所有二分类器的 predict 硬投票与 _predict_binary 置信度。predictions 矩阵记录了每个样本在每个二分类器上的 0/1 判定。_ovr_decision_function 将投票计数与归一化置信度相加,既保留多数意见又化解平局。

14.11.2 二分类特例的降维输出

n_classes_ == 2 时仅含一个二分类器,decision_function 返回形状 (n_samples,) 而非 (n_samples, 2)predict 中同样以 _threshold_for_binary_predict 判定,与 sklearn 二分类约定对齐。

14.11.3 test_ovo_ties 揭示的平局破解原则

当三个二分类器各投一票(1:1:1)时,使用 decision_function 的连续置信度区分胜者。votes = np.round(decision)normalized_confidences = decision - votes 分离整数票数与分数置信度。test_ovo_ties2 证明平局胜者不限于前两个类别。

14.12 OvO 边界与异常 —— 从列表输入到单类别报错

14.12.1 输入容器的兼容性

test_ovo_fit_on_list 验证列表输入与数组输入产生相同预测。test_ovo_string_y 验证字符串标签不会被错误编码。

14.12.2 异常场景的防御式设计

len(self.classes_) == 1 时抛出 ValueError,因为 C(1,2)=0 无二分类任务。test_ovo_one_class 验证单类别时的报错消息。test_ovo_float_y 验证连续标签触发 check_classification_targets 报错。

14.12.3 GridSearchCV 的参数透传

test_ovo_gridsearch 验证 estimator__C 能深入子估计器参数。

14.13 ECOC 码本生成 —— 随机码本里的“纠错编码艺术”

14.13.1 码本形状与 code_size 弹性控制

n_estimators = int(n_classes * code_size)code_size>1 增加冗余位提升容错,<1 压缩模型规模。码本由 random_state.uniform 生成均匀随机值,再以 0.5 为阈值二值化。有 decision_function 时编码为 {-1, +1},仅 predict_proba 时编码为 {0, 1}。

14.13.2 classes_index 映射与目标矩阵构建

通过字典将类别标签映射到码本行索引,O(1) 查找。Y 矩阵每列对应一个二分类器的目标,每行对应一个样本的码字。

14.13.3 预测阶段的最近码字搜索

pairwise_distances_argmin 在类别码本中寻找欧氏距离最近的码字。注意 order="F" 的 F-contiguous 数组构造,转置后变为 C-contiguous 以满足 ArgKmin 的输入要求。

14.13.4 __sklearn_tags__ 的 sparse 标签透传

从基估计器继承 sparse 输入支持标签。test_ecoc_delegate_sparse_base_estimator 验证稀疏输入的委托与拦截。

14.14 MultiOutput 公共基座 —— 多输出拟合的“工厂车间”

14.14.1 _MultiOutputEstimator 的抽象基类设计

通过 ABCMeta 强制子类实现 __init__,统一 estimatorn_jobs 参数入口。fit 中对 ymulti_output=True 验证,确保目标至少二维。对分类器额外调用 check_classification_targets 拒绝连续标签。

14.14.2 逐列克隆与并行拟合

# 第 14 章 —— sklearn/multioutput.py - _MultiOutputEstimator.fit() (第155-213行)
@_fit_context(
    # MultiOutput*.estimator is not validated yet
    prefer_skip_nested_validation=False
)
def fit(self, X, y, sample_weight=None, **fit_params):
    """Fit the model to data, separately for each output variable.

    Parameters
    ----------
    X : {array-like, sparse matrix} of shape (n_samples, n_features)
        The input data.

    y : {array-like, sparse matrix} of shape (n_samples, n_outputs)
        Multi-output targets. An indicator matrix turns on multilabel
        estimation.

    sample_weight : array-like of shape (n_samples,), default=None
        Sample weights. If `None`, then samples are equally weighted.
        Only supported if the underlying regressor supports sample
        weights.

    **fit_params : dict of string -> object
        Parameters passed to the ``estimator.fit`` method of each step.

        .. versionadded:: 0.23

    Returns
    -------
    self : object
        Returns a fitted instance.
    """
    if not hasattr(self.estimator, "fit"):
        raise ValueError("The base estimator should implement a fit method")

    y = validate_data(self, X="no_validation", y=y, multi_output=True)

    if is_classifier(self):
        check_classification_targets(y)

    if y.ndim == 1:
        raise ValueError(
            "y must have at least two dimensions for "
            "multi-output regression but has only one."
        )

    if _routing_enabled():
        if sample_weight is not None:
            fit_params["sample_weight"] = sample_weight
        routed_params = process_routing(
            self,
            "fit",
            **fit_params,
        )
    else:
        if sample_weight is not None and not has_fit_parameter(
            self.estimator, "sample_weight"
        ):
            raise ValueError(
                "Underlying estimator does not support sample weights."
            )

        fit_params_validated = _check_method_params(X, params=fit_params)
        routed_params = Bunch(estimator=Bunch(fit=fit_params_validated))
        if sample_weight is not None:
            routed_params.estimator.fit["sample_weight"] = sample_weight

    self.estimators_ = Parallel(n_jobs=self.n_jobs)(
        delayed(_fit_estimator)(
            self.estimator, X, y[:, i], **routed_params.estimator.fit
        )
        for i in range(y.shape[1])
    )

    if hasattr(self.estimators_[0], "n_features_in_"):
        self.n_features_in_ = self.estimators_[0].n_features_in_
    if hasattr(self.estimators_[0], "feature_names_in_"):
        self.feature_names_in_ = self.estimators_[0].feature_names_in_

    return self

_fit_estimator 每次克隆一个新估计器,按列 y[:, i] 独立训练。并行结果收集后通过 np.asarray(y).T 将列表转置为 (n_samples, n_outputs) 矩阵。sample_weight 支持采用两套路径:_routing_enabled() 时走元数据路由,否则走兼容逻辑。

14.14.3 _partial_fit_estimator 的增量语义

首次调用时克隆估计器,后续批次持续累积。classes 参数可选传入,供分类器增量训练使用。

14.14.4 标签透传:__sklearn_tags__ 的继承与覆写

稀疏输入支持从基估计器标签继承。目标标签强制设置 multi_output=Truesingle_output=False,确保下游元估计器正确感知。

14.14.5 get_metadata_routing 的元数据路由配置

将 fit 和 partial_fit 调用路由到子估计器的同名方法。

14.15 MultiOutputClassifier 概率接口 —— 条件暴露的“能力探测器”

14.15.1 MultiOutputRegressor 的继承与委托

__init__ 调用基类初始化,保持参数约束一致。partial_fit 直接委托给基类实现。test_multioutput_regressor_has_partial_fit 验证线性回归不具备 partial_fit 时属性不存在。

14.15.2 MultiOutputClassifier.fit 的 classes_ 收集

调用基类 fit 后收集每个子估计器的 classes_ 到列表。test_multi_output_classes_ 验证列表长度与各列类别正确性。

14.15.3 _check_predict_proba 的双态探测逻辑

已拟合时逐一检查每个子估计器是否拥有 predict_proba。未拟合时检查 self.estimator 是否拥有该方法。与 available_if 结合实现条件方法暴露,未支持时 hasattr 直接为 False。

14.15.4 score 的严格全匹配语义

np.all(y == y_pred, axis=1) 要求该样本所有输出列都预测正确才计 1 分。对 y.ndim == 1 抛错,对输出维度不匹配抛错。

14.15.5 test_multi_output_predict_proba 的嵌套 AttributeError 验证

SGDClassifier 的 hinge 损失不暴露 predict_proba,错误链含三层信息。最内层 __cause__.__cause__ 揭示真正原因:概率估计不适用于 hinge 损失。

14.16 MultiOutput 训练与验证全景 —— 从样本权重到嵌套元估计器

14.16.1 样本权重的两套路径

  • test_multi_target_sample_weights_api 验证不支持样本权重的基估计器报错

  • test_multi_target_sample_weights 验证加权拟合等价于重复样本

  • test_multi_output_classification_sample_weights 覆盖分类场景

  • test_multi_target_sample_weight_partial_fit 验证不同权重导致不同预测结果

14.16.2 稀疏输入与回归的跨容器等价性

test_multi_target_sparse_regression 验证 CSR/CSC/COO/LIL/DOK/BSR 容器。稀疏训练与稠密训练的预测输出完全一致。

14.16.3 元估计器嵌套与 Duck Typing

test_multiclass_multioutput_estimator 验证 OneVsRestClassifier 嵌套 MultiOutputClassifier。test_multiclass_multioutput_estimator_predict_proba 验证嵌套概率输出的数值精度。test_multioutputregressor_ducktypes_fitted_estimator 验证 StackingRegressor 作为基估计器。

14.16.4 fit 参数传递与拦截

DummyRegressorWithFitParamsDummyClassifierWithFitParams 记录收到的 fit 参数。test_multioutput_estimator_with_fit_params 验证参数透传。test_fit_params_no_routing 验证未路由参数报错。

14.16.5 增量学习的边界与并行性

test_multi_output_classification_partial_fit 验证逐列增量与单列一致性。test_multi_output_classification_partial_fit_no_first_classes_exception 验证首次调用缺 classes 报错。test_multi_output_classification_partial_fit_parallelism 验证并行时估计器对象不共享。

14.16.6 模块级全局数据构造

__main__ 中加载 iris 数据并生成多输出数据集,供测试函数复用。

14.17 链式基座 _BaseChain —— 前序预测的“接力棒传递”

14.17.1 _BaseChain.__init___get_estimator 的弃用处理

estimatorbase_estimator 互斥,两者同时提供触发 ValueError。base_estimator != "deprecated" 时发出 FutureWarning 并使用该值。test_base_estimator_deprecation 验证弃用警告与互斥条件。

14.17.2 order 参数的三种形态解析

  • None:按列自然顺序 0,1,...,n_outputs-1

  • 数组:自定义链路顺序,sorted(order_) == list(range(n_outputs)) 校验合法性

  • 'random':用 check_random_state 做排列

14.17.3 cv 模式下的交叉验证预测替代真值

cv=None 时直接将真实标签列拼接为增广特征。cv 指定时用 cross_val_predict 生成前序估计器的 OOF 预测填入增广矩阵。预测概率的方法取第二列 [:, 1] 作为前序特征,避免数据泄露。

14.17.4 稀疏矩阵的 hstack 兼容处理

对 dok 格式先转 coo 再 hstack(规避 SciPy 性能问题),最终输出统一为 CSR。稠密输入走 np.hstack,预测零值初始化为 np.zeros

14.17.5 _log_message 的进度打印

verbose 模式下打印 (chain_idx of n_estimators) Processing order Xtest_classifier_chain_verbosetest_regressor_chain_verbose 验证输出格式。

graph TD A[输入特征 X] --> B[第1个估计器] B --> C[预测结果作为新特征] C --> D[拼接到 X 形成 X_aug] D --> E[第2个估计器] E --> F[预测结果作为新特征] F --> G[...] G --> H[第n个估计器] H --> I[收集所有输出] I --> J[inv_order 逆序还原] J --> K[最终预测 Y]

14.18 ClassifierChain 与 RegressorChain —— 链式建模的“预测接力赛”

14.18.1 chain_method 的响应方法选择

支持 predictpredict_probapredict_log_probadecision_function,可用列表按优先级选择。_check_response_method 返回方法的 __name__ 存储于 chain_method_。RegressorChain 无此参数,固定使用 predict

14.18.2 _get_predictions 中的逆序还原

预测按 order_ 顺序生成后,通过 inv_order 将输出列映射回原始标签顺序。Y_feature_chain 累积链路中继特征,Y_output_chain 记录最终输出。

14.18.3 ClassifierChain 的多输出方法

  • predict_proba_get_predictions 输出概率

  • predict_log_proba 对概率取对数

  • decision_function_get_predictions 输出决策函数

14.18.4 test_classifier_chain_vs_independent_models 验证链路增益

与 OneVsRest 独立模型对比,链路捕捉标签间相关性后 Jaccard 分数更高。回归链路同样通过 mean_squared_error 验证 cv 模式的有效性。

14.18.5 fit_params 的链路传递

test_regressor_chain_w_fit_params 验证样本权重传递到每个链估计器。

14.18.6 链式分类器的完整验证矩阵

  • test_classifier_chain_fit_and_predict 验证四种链方法与两种响应方法的组合

  • test_classifier_chain_fit_and_predict_with_linear_svc 验证决策函数与硬预测一致

  • test_classifier_chain_fit_and_predict_with_sparse_data 验证稀疏输入输出与稠密一致

  • test_regressor_chain_fit_and_predict 验证回归链系数维度随链路递增

14.19 设计中的取舍

14.19.1 为什么 OvR 不直接用 np.argmax 而是手写反向遍历?

因为 np.argmax 的平局行为是“返回第一个最大值索引”。OvR 的估计器列表顺序与类别顺序一致,若正向遍历 np.maximum,平局时会保留最后一个最大值(即索引最大的类别),与 np.argmax 行为相反。反向遍历配合索引映射 n_classes - i - 1 精确复现了 np.argmax 的平局语义,保证与单模型多分类器(如 LogisticRegression multinomial)的一致性。

14.19.2 OvO 为何用 _ovr_decision_function 而非简单计票?

简单计票在类别数为偶数时极易出现平局(如 3 类产生 3 个二分类器,每个类参与 2 场对决,极易 1:1 平票)。引入归一化置信度作为“分数票”,将决策函数的连续输出加到整数票数上,既保留了多数表决的鲁棒性,又利用置信度区分度化解平局。这是“硬投票+软置信度”的混合决策策略。

14.19.3 ECOC 为何用欧氏距离而非汉明距离?

码本生成时 decision_function 对应 {-1,+1},predict_proba 对应 {0,1}。预测阶段聚合的是连续置信度(实数向量),而非离散二值码字。欧氏距离在实数空间自然衡量“置信度向量”与“码字向量”的接近度,比汉明距离更利用连续信息量。

14.19.4 MultiOutput 为何强制 multi_output=True 标签?

元估计器嵌套时(如 MultiOutputClassifier(OneVsRestClassifier(...))),内层元估计器需要知道自己正在处理多输出任务,从而正确设置自身标签(如 target_tags.multi_output)。强制覆写确保标签传递链路不断裂,避免下游工具(如 cross_val_predict、元数据路由)误判任务类型。

14.19.5 ClassifierChain 为何要 inv_order 逆序还原?

链路训练顺序 order_ 可能打乱原始列顺序(如 [2, 0, 1])。预测时按链路顺序生成 Y_output_chain,第 0 列对应原始第 2 个标签。inv_order[0,1,2] 映射为 [1,2,0],使 Y_output_chain[:, inv_order] 恢复为原始标签顺序,对用户透明。

14.20 动手练习

14.20.1 练习 1:阅读 OvR 的预测路径与平局破解

阅读 sklearn/multiclass.py 第 440-472 行 OneVsRestClassifier.predict(),理解多分类模式下的 argmax 平局破解技巧。

重点关注:

  1. 为什么 maxima 初始化为 -np.inf

  2. 反向遍历 reversed(self.estimators_) 如何保证与 np.argmax 的平局行为一致?

  3. argmaxima[maxima == pred] 这行代码在什么条件下触发?

回答问题:

  • 如果两个类别的决策分数完全相同,最终会选择哪一个类别?为什么?

  • 参考 test_ovr_ties 测试(第 80-96 行),说明 Dummy 估计器如何构造平局场景。

14.20.2 练习 2:追踪 MultiOutput 的逐列并行训练

阅读 sklearn/multioutput.py 第 155-213 行 _MultiOutputEstimator.fit(),理解多输出估计器的训练流程。

重点关注:

  1. validate_datamulti_output=True 参数如何校验目标维度?

  2. _routing_enabled() 与兼容逻辑两套路径分别如何传递 sample_weight

  3. Parallel(n_jobs=self.n_jobs) 如何对每一列 y[:, i] 并行执行 _fit_estimator

回答问题:

  • 为什么 np.asarray(y).T 可以将并行预测结果还原为 (n_samples, n_outputs) 形状?

  • 如果基估计器不支持 sample_weight 但用户传入了该参数,会发生什么?在源码的哪一行抛出异常?

14.20.3 练习 3:解析 ClassifierChain 的交叉验证预测机制

阅读 sklearn/multioutput.py 第 431-511 行 _BaseChain.fit(),重点理解 cv 模式下的 OOF 预测填充。

重点关注:

  1. order 参数的三种形态(None、数组、'random')分别如何解析?

  2. cross_val_predict 生成的 OOF 预测为何能避免数据泄露?

  3. cv_result.ndim > 1 时为何取 [:, 1]

  4. inv_order 在预测阶段如何将链路输出还原为原始标签顺序?

回答问题:

  • 如果 cv=None,前序估计器的预测特征从哪里来?

  • 稀疏矩阵的 dok 格式为何要先转 coo 再 hstack?参考源码注释说明性能原因。

14.20.4 练习 4:对比三大多分类策略的适用场景

阅读 sklearn/multiclass.py 中 OneVsRestClassifier、OneVsOneClassifier 和 OutputCodeClassifier 的类文档字符串(docstring),总结各自的优缺点。

回答问题:

  • OvR 为何被称为“最常用的多分类策略”和“公平的默认选择”?

  • OvO 在什么情况下比 OvR 更有优势?文档中提到了哪类算法?

  • ECOC 的 code_size 参数如何控制模型规模与容错能力?当 code_size=0.5code_size=2.0 时各有什么效果?

动手实践:

sklearn.datasets.load_iris() 数据分别训练三种策略(基估计器用 LinearSVC),对比训练时间、len(estimators_) 和预测准确率。

14.21 本章小结

这一章中我们学习了 scikit-learn 如何将复杂的多分类与多输出任务分解为多个二分类或单输出子问题。首先我们剖析了 OneVsRestClassifier 的流水线分拣机制:LabelBinarizer 稀疏指示矩阵、并行训练、反向遍历平局破解、多标签 CSC 构建、增量学习锁定类别与条件方法暴露。其次我们深入 OneVsOneClassifier 的锦标赛模式:两两配对样本筛选、预计算核裁剪、投票+置信度融合决策函数、二分类特例降维。接着我们探索了 OutputCodeClassifier 的纠错编码艺术:随机码本生成、最近码字搜索、code_size 控制模型规模。然后我们转入 multioutput 模块,理解 _MultiOutputEstimator 的工厂车间模式:逐列克隆并行拟合、样本权重双路径路由、标签强制 multi_output=True、条件方法暴露机制。最后我们解析了链式建模的接力赛逻辑:BaseChain 的 order 解析与 cv 模式 OOF 预测防泄露、_get_predictions 的逆序还原、ClassifierChain 的 chain_method 灵活选择。

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

| 概念 | 解释 |

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

| _fit_binary | 训练单个二分类器,单类时退化为常量预测器并发出警告 |

| _partial_fit_binary | 增量训练单个二分类器,固定使用 classes=(0,1) |

| _threshold_for_binary_predict | 根据估计器类型选择判定阈值:decision_function=0.0,predict_proba=0.5 |

| _ConstantPredictor | 单类别退化时的常量预测器,实现 predict/decision_function/predict_proba |

| LabelBinarizer(sparse_output=True) | 将多分类标签转为 CSC 稀疏指示矩阵,性能优于稠密表示 |

| OvR.predict 反向遍历 | 倒序遍历估计器精确复现 np.argmax 的平局破解行为 |

| OvR.predict_proba 归一化 | 单标签通过 where 掩码安全归一化,多标签返回边际概率不求和为 1 |

| OvR.multilabel_ | 根据 label_binarizer_.y_type_ 判断是否为多标签模式 |

| available_if | 动态检查基估计器能力,未支持时属性直接不存在(hasattr 为 False) |

| _fit_ovo_binary | 用 logical_or 筛选样本子集并重编码为 0/1,记录 indcond 供预计算核裁剪 |

| _partial_fit_ovo_binary | 增量训练单个 OvO 二分类器,空批次时直接返回未更新估计器 |

| _ovr_decision_function | 将硬投票计数与归一化置信度相加,既保留多数意见又化解平局 |

| ECOC 码本生成 | 随机 uniform 生成后以 0.5 阈值二值化,decision_function 用 {-1,+1},predict_proba 用 {0,1} |

| pairwise_distances_argmin | 在码本中搜索欧氏距离最近的码字,实现纠错解码 |

| _fit_estimator | 克隆单个估计器并按列拟合,支持样本权重 |

| _partial_fit_estimator | 增量训练单个估计器,首次调用时克隆 |

| _MultiOutputEstimator.fit | 逐列克隆估计器并行训练,np.asarray(y).T 转置还原形状 |

| MultiOutputClassifier.score | 严格全匹配语义,所有输出列均正确才计 1 分 |

| MultiOutputClassifier._check_predict_proba | 双态探测 predict_proba 可用性,配合 available_if 实现条件方法暴露 |

| _BaseChain.fit | 解析 order 参数(None/数组/random),cv 模式用 cross_val_predict 防数据泄露 |

| _BaseChain._get_estimator | 校验 estimator 与 base_estimator 的互斥关系并发出弃用警告 |

| _get_predictions 逆序还原 | inv_order 将链路输出映射回原始标签顺序,支持 predict_proba/decision_function 等多种输出 |

| _BaseChain._log_message | verbose 模式下打印链路进度消息 |

| ClassifierChain.predict_log_proba | 对 predict_proba 的结果取对数 |

| test_multi_target_regression_partial_fit | 验证 MultiOutputRegressor 增量学习与逐列 SGDRegressor 增量结果一致 |

| OneVsOneClassifier.__sklearn_tags__ | 从基估计器透传 pairwise 与 sparse 输入支持标签 |

| test_multioutput.py.__main__ | 模块级加载 iris 数据并构造多输出数据集,供后续测试复用 |

下一章中,我们将学习管道与特征联合 —— 构架“可复用的机器学习流水线”,解析 Pipeline 的顺序组合与 FeatureUnion 的并行拼接如何作为元估计器管理多个子步骤,并实现参数共享、缓存与元数据路由。

第 15 章 —— 管道与特征联合 —— 构架“可复用的机器学习流水线”

15.1 学习目标

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

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

  • 理解 Pipeline 如何通过 _BaseComposition 实现嵌套参数管理与顺序执行编排

  • 掌握元数据路由的双模式机制(参数前缀 s__p 与智能分拣 process_routing)

  • 深入 transform_input 元数据变换的设计动机与缓存策略,理解多验证集场景的解决方案

  • 理解 FeatureUnion 的并行执行架构、特征拼接策略与特征名前缀/冲突检测机制

  • 能阅读并修改 Pipeline/FeatureUnion 的切片、缓存、拟合状态判定等核心源码逻辑

  • 熟悉元数据路由测试基础设施(ConsumingTransformer、_Registry 等)的构建方式

  • 掌握 make_pipeline/make_union 自动命名机制与 Pipeline 分数、预测等公共方法实现

  • 理解测试辅助类(Mult、NoTrans、Transf、FitParamT 等)在验证 Pipeline 行为中的作用

15.2 生活类比

想象 Pipeline 是一条智能生产流水线,FeatureUnion 是并行分流装配台中间变换器 = 流水线上的加工工位(清洗→切割→喷漆);最终估计器 = 出厂前的质检包装工位(分类/回归/聚类);参数前缀 s__p = 传统传票制度(每个工位单独发一张传票,格式必须写对);元数据路由 = 智能分拣系统(每个工位提前声明需要什么物料,中心系统自动配送);transform_input = 二次加工服务(半成品在送达某工位前,先经过前序工位的预处理);memory 缓存 = 工位旁的标准化零件库(相同输入直接取成品,不重新加工);FeatureUnion 的 n_jobs = 装配台的多臂机器人(并行处理多路特征流);特征名前缀 verbose_feature_names_out = 为每个零件贴上工位标签(pca__x1 表示 PCA 工位产出);passthrough = 传送带上的直通货箱(不做任何加工);drop = 直接丢弃的废料箱(跳过该工位)。就像工厂经理需要设计合理的工位布局与物料流转路径,Pipeline 的 _fit 方法也需要精确编排每个变换器的拟合顺序与数据流动;而元数据路由的智能分拣系统,则确保了 sample_weight、X_val 等特殊物料能准确送达真正需要的工位。

15.3 源码地图

sklearn/pipeline.py

├── 模块级工具函数

│ ├── _final_estimator_has(attr) # 动态方法可用性检查装饰器

│ ├── _cached_transform() # 元数据变换缓存机制

│ ├── _name_estimators() # make_pipeline/make_union 自动命名

│ ├── _transform_one() # FeatureUnion 单变换器 transform + 加权

│ ├── _fit_transform_one() # Pipeline 单变换器 fit_transform + 加权

│ └── _fit_one() # FeatureUnion 单变换器 fit

├── class Pipeline(_BaseComposition)

│ ├── init() # 属性赋值,不克隆估计器

│ ├── set_output() # 递归设置所有步骤的输出容器

│ ├── get_params()/set_params() # 双下划线嵌套参数管理

│ ├── _validate_steps() # 中间/最终步骤能力校验

│ ├── _iter() # 惰性迭代器,with_final/filter_passthrough 控制

│ ├── len() # 返回步骤数量

│ ├── getitem() # 整数索引/切片返回子 Pipeline

│ ├── named_steps # Bunch 属性式访问

│ ├── _final_estimator # 保护性获取最终估计器

│ ├── _log_message() # verbose 日志消息生成

│ ├── _check_method_params() # 新旧路由模式分支

│ ├── _get_metadata_for_step() # transform_input 元数据变换

│ ├── _fit() # 顺序拟合中间步骤,内存缓存

│ ├── fit() # 拟合中间步骤 + 最终估计器

│ ├── _can_fit_transform() # 判断是否支持 fit_transform

│ ├── fit_transform() # 拟合并用最终估计器变换

│ ├── predict()/predict_proba()/decision_function() # 双模式预测分支

│ ├── fit_predict() # 拟合并用最终估计器聚类

│ ├── score_samples()/predict_log_proba() # 无路由/路由双模式变换

│ ├── transform()/inverse_transform() # 全步骤变换/逆序反变换

│ ├── score() # 最终估计器评分

│ ├── classes_ # 最终分类器的类别属性

│ ├── sklearn_tags() # 标签传播

│ ├── get_feature_names_out() # 链式特征名变换

│ ├── n_features_in_/feature_names_in_ # 首步属性委派

│ ├── sklearn_is_fitted() # 拟合状态检查

│ ├── sk_visual_block() # VisualBlock 序列可视化块

│ └── get_metadata_routing() # 中间/最终步骤方法映射

├── make_pipeline() # 自动命名便捷构造器

├── class FeatureUnion(TransformerMixin, _BaseComposition)

│ ├── init() # 并行变换器列表初始化

│ ├── set_output() # 递归设置所有变换器的输出

│ ├── named_transformers # Bunch 属性式访问

│ ├── get_params()/set_params() # 双下划线嵌套参数管理

│ ├── _validate_transformers() # 变换器实例校验

│ ├── _validate_transformer_weights() # 权重键名校验

│ ├── _iter() # drop 跳过、passthrough 转换

│ ├── get_feature_names_out() # 收集各变换器的输出名

│ ├── _add_prefix_for_feature_names_out() # 前缀/冲突检测

│ ├── fit() # 并行拟合入口

│ ├── fit_transform() # 并行拟合并拼接

│ ├── _log_message() # verbose 日志消息

│ ├── _parallel_func() # joblib 并行调度

│ ├── transform() # 并行变换并拼接

│ ├── _hstack() # 稀疏/容器适配/Array API 三策略拼接

│ ├── _update_transformer_list() # 替换已拟合的变换器

│ ├── n_features_in_/feature_names_in_ # 首步属性委派

│ ├── sklearn_is_fitted() # 全步骤拟合检查

│ ├── sk_visual_block() # VisualBlock 并行可视化块

│ ├── getitem() # 按键名获取变换器

│ └── get_metadata_routing() # fit/fit_transform/transform 映射

└── make_union() # 自动命名便捷构造器

sklearn/tests/test_pipeline.py

├── 测试辅助类

│ ├── NoFit/NoTrans/NoInvTransf # 最小实现校验

│ ├── Transf/TransfFitParams # 可逆变换器

│ ├── Mult # 乘法变换器

│ ├── FitParamT # 带参数的拟合分类器

│ ├── DummyTransf # 带时间戳的缓存测试变换器

│ ├── DummyEstimatorParams # 预测参数测试估计器

│ ├── SimpleEstimator # 完整方法路由测试估计器

│ └── FeatureNameSaver # 特征名保存测试变换器

├── 测试辅助函数

│ └── create_mock_transformer() # 构造带自定义特征名的 mock 变换器

├── 基础流程测试

│ ├── test_pipeline_invalid_parameters() # 参数校验

│ ├── test_meta_estimator_raises_class_not_instance_error() # 实例而非类校验

│ ├── test_empty_pipeline() # 空流水线报错

│ ├── test_pipeline_init_tuple() # steps 元组形式

│ ├── test_pipeline_methods_anova() # 完整方法扫描

│ ├── test_pipeline_fit_params() # fit 参数传递

│ ├── test_pipeline_sample_weight_supported()/unsupported() # 样本权重路由

│ ├── test_pipeline_raise_set_params_error() # 参数错误信息

│ ├── test_pipeline_methods_pca_classifier() # PCA+分类器形状

│ ├── test_pipeline_score_samples_pca_lof() # score_samples 语义

│ ├── test_fit_predict_on_pipeline() # fit_predict 语义

│ ├── test_fit_predict_on_pipeline_without_fit_predict() # 方法缺失回退

│ ├── test_fit_predict_with_intermediate_fit_params() # fit 参数传递到中间步骤

│ ├── test_predict_methods_with_predict_params() # 预测参数透传

│ ├── test_pipeline_transform() # transform 与 inverse_transform

│ ├── test_pipeline_fit_transform() # 缺少 fit_transform 时回退

│ ├── test_set_pipeline_steps() # 步骤替换

│ ├── test_pipeline_named_steps() # 命名步骤访问

│ ├── test_pipeline_ducktyping() # 动态方法可用性

│ ├── test_pipeline_estimator_type() # 估计器类型判定

│ └── test_pipeline_with_no_last_step() # 无末步流水线

├── 参数管理与切片测试

│ ├── test_step_name_validation() # 步骤名规范

│ ├── test_set_params_nested_pipeline() # 嵌套参数设置

│ ├── test_pipeline_slice() # 切片语义验证

│ └── test_pipeline_index() # 索引语义验证

├── passthrough/None 测试

│ ├── test_pipeline_correctly_adjusts_steps() # 步骤调整

│ ├── test_set_pipeline_step_passthrough() # passthrough 步骤

│ └── test_pipeline_get_tags_none() # passthrough 标签

├── 缓存与内存测试

│ ├── test_pipeline_memory() # joblib 缓存命中验证

│ └── test_make_pipeline_memory() # make_pipeline 参数传递

├── 元数据路由测试

│ ├── test_metadata_routing_for_pipeline() # 全方法元数据路由

│ ├── test_metadata_routing_error_for_pipeline() # 未声明元数据报错

│ ├── test_routing_passed_metadata_not_supported() # 路由禁用时报错

│ ├── test_feature_union_metadata_routing() # FeatureUnion 元数据路由

│ ├── test_feature_union_metadata_routing_error() # 元数据未声明报错

│ ├── test_feature_union_get_metadata_routing_without_fit() # 获取路由信息

│ ├── test_feature_union_fit_params() # fit 参数透传

│ └── test_feature_union_fit_params_without_fit_transform() # fit+transform 路由

├── FeatureUnion 测试

│ ├── test_feature_union() # 基本功能

│ ├── test_feature_union_named_transformers() # 属性式访问

│ ├── test_feature_union_weights() # 加权拼接

│ ├── test_feature_union_parallel() # 并行一致性

│ ├── test_feature_union_feature_names() # 特征名前缀

│ ├── test_set_feature_union_steps() # 步骤替换

│ ├── test_set_feature_union_step_drop() # drop 步骤

│ ├── test_set_feature_union_passthrough() # passthrough 步骤

│ ├── test_feature_union_warns_unknown_transformer_weight() # 权重校验

│ ├── test_feature_union_1d_output() # 输出维度校验

│ ├── test_feature_union_array_api_compliance() # Array API 合规

│ ├── test_feature_union_array_api_support_tag() # Array API 标签

│ ├── test_feature_union_getitem() # 按键获取

│ ├── test_feature_union_getitem_error() # 非字符串键报错

│ ├── test_feature_union_feature_names_in_() # DataFrame 列名

│ └── test_feature_union_set_output() # set_output 集成

├── 特征名相关测试

│ ├── test_features_names_passthrough() # 特征名透传

│ ├── test_feature_names_count_vectorizer() # 向量化器特征名

│ ├── test_pipeline_feature_names_out_error_without_definition() # 缺少特征名方法

│ ├── test_pipeline_get_feature_names_out_passes_names_through() # 特征名链

│ ├── test_make_union_passes_verbose_feature_names_out() # 递归参数传递

│ ├── test_feature_union_passthrough_get_feature_names_out_true() # 前缀模式

│ ├── test_feature_union_passthrough_get_feature_names_out_false() # 无前缀模式

│ └── test_feature_union_passthrough_get_feature_names_out_false_errors() # 冲突检测

├── transform_input 测试

│ ├── test_transform_input_pipeline() # 元数据变换全流程

│ ├── test_transform_input_explicit_value_check() # 精确值断言

│ ├── test_transform_input_no_slep6() # 路由禁用时报错

│ └── test_transform_tuple_input() # 多验证集 tuple 输入

├── 拟合状态与属性测试

│ ├── test_pipeline_check_if_fitted() # pipeline 拟合检查

│ ├── test_feature_union_check_if_fitted() # FeatureUnion 拟合检查

│ ├── test_n_features_in_pipeline() # n_features_in_ 委派

│ ├── test_n_features_in_feature_union() # FeatureUnion 特征数

│ ├── test_classes_property() # classes_ 属性

│ ├── test_sklearn_tags_with_empty_pipeline() # 空流水线标签

│ └── test_pipeline_warns_not_fitted() # 未拟合警告

└── 其他

├── test_pipeline_param_error() # 路由禁用参数错误

├── test_pipeline_missing_values_leniency() # 缺失值宽容

├── test_pipeline_with_estimator_with_len() # len 兼容

├── test_pipeline_set_output_integration() # set_output 集成

├── test_score_samples_on_pipeline_without_score_samples() # 方法缺失报错

├── test_search_cv_using_minimal_compatible_estimator() # 极简估计器兼容

├── test_verbose() # verbose 日志格式

└── test_make_pipeline()/test_make_union()/test_make_union_kwargs() # 构造器测试

sklearn/tests/metadata_routing_common.py

├── record_metadata()/record_metadata_not_default() # 元数据记录

├── check_recorded_metadata() # 元数据断言

├── assert_request_is_empty() # 元数据请求空检查

├── assert_request_equal() # 元数据请求相等检查

├── _Registry # 深拷贝共享列表

├── ConsumingRegressor/Classifier # 消费元数据的回归/分类器

├── ConsumingClassifierWithoutPredictProba # 无 predict_proba 的消费分类器

├── ConsumingClassifierWithoutPredictLogProba # 无 predict_log_proba 的消费分类器

├── ConsumingClassifierWithOnlyPredict # 仅 predict 的消费分类器

├── ConsumingTransformer # 消费元数据的变换器

├── ConsumingNoFitTransformTransformer # 无 fit_transform 的变换器

├── NonConsumingClassifier/Regressor # 不消费元数据的分类/回归器

├── ConsumingScorer # 消费元数据的评分器

├── ConsumingSplitter # 消费元数据的切分器

├── ConsumingSplitterInheritingFromGroupKFold # 继承 GroupKFold 的消费切分器

├── MetaRegressor/WeightedMetaRegressor # 元回归器

├── MetaTransformer # 元变换器

└── WeightedMetaClassifier # 加权元分类器

sklearn/tests/test_metadata_routing.py

├── SimplePipeline # 极简元估计器测试类

├── SimpleEstimator # 完整方法路由测试估计器

├── test_assert_request_is_empty() # 空请求检查

├── test_estimator_puts_self_in_registry() # 注册表注入

├── test_request_type_is_alias()/test_request_type_is_valid() # 请求类型判断

├── test_default_requests()/test_default_request_override() # 默认请求

├── test_process_routing_invalid_method()/invalid_object() # 非法方法路由

├── test_process_routing_empty_params_get_with_default() # 空参数路由

├── test_simple_metadata_routing() # 简单路由

├── test_nested_routing() # 嵌套路由

├── test_nested_routing_conflict() # 路由冲突

├── test_invalid_metadata() # 非法元数据

├── test_get_metadata_routing()/test_get_routing_for_object() # 路由信息获取

├── test_setting_default_requests() # 默认请求设置

├── test_removing_non_existing_param_raises() # 移除不存在参数报错

├── test_method_metadata_request() # 方法级请求

├── test_metadata_request_consumes_method() # consumes 方法

├── test_metadata_router_consumes_method() # 路由器 consumes

├── test_metaestimator_warnings()/test_estimator_warnings() # 警告测试

├── test_string_representations() # 字符串表示

├── test_validations() # 参数验证

├── test_methodmapping() # 方法映射

├── test_metadatarouter_add_self_request() # 自请求添加

├── test_metadata_routing_add() # 路由添加

├── test_metadata_routing_get_param_names() # 路由参数名提取

├── test_method_generation() # 方法生成

├── test_composite_methods() # 复合方法

├── test_no_feature_flag_raises_error() # 功能开关关闭报错

├── test_none_metadata_passed() # None 元数据

├── test_no_metadata_always_works() # 无元数据兼容

├── test_unsetmetadatapassederror_correct() # 未设置请求错误

├── test_unsetmetadatapassederror_correct_for_composite_methods() # 复合方法错误

└── test_unbound_set_methods_work() # 未绑定 set 方法

15.4 Pipeline 类骨架与构造 —— 搭建“流水线车间的传送带骨架”

Pipeline 继承自 _BaseComposition,这为它提供了嵌套参数管理(get_params/set_params 的双下划线语法)和命名校验能力。构造时只做属性赋值,不克隆估计器,保持构建时的对象引用,fit 时才按需克隆。

类型定义详解

源码路径:sklearn/pipeline.py - Pipeline.__init__()(第221-226行)

def __init__(self, steps, *, transform_input=None, memory=None, verbose=False):
    self.steps = steps
    self.transform_input = transform_input
    self.memory = memory
    self.verbose = verbose

这段代码定义了 Pipeline 的构造函数,接收步骤列表、元数据变换列表、缓存对象和详细日志开关,仅做属性赋值,不做克隆与校验,校验推迟到 fit 阶段。

核心类型定义

源码路径:sklearn/pipeline.py - Pipeline 类属性(第150-158行)

_parameter_constraints: dict = {
    "steps": [list, Hidden(tuple)],
    "transform_input": [list, None],
    "memory": [None, str, HasMethods(["cache"])],
    "verbose": ["boolean"],
}

使用声明式参数约束:steps 必须为列表或元组,transform_input 是 1.6 新增的字符串列表,用于指定哪些元数据参数需要经过前置变换器处理,memory 接受 joblib.Memory 兼容对象,verbose 为布尔值。

自动命名机制

源码路径:sklearn/pipeline.py - _name_estimators()(第1188-1204行)

def _name_estimators(estimators):
    """Generate names for estimators."""
    names = [
        estimator if isinstance(estimator, str) else type(estimator).__name__.lower()
        for estimator in estimators
    ]
    namecount = defaultdict(int)
    for est, name in zip(estimators, names):
        namecount[name] += 1
    for k, v in list(namecount.items()):
        if v == 1:
            del namecount[k]
    for i in reversed(range(len(estimators))):
        name = names[i]
        if name in namecount:
            names[i] += "-%d" % namecount[name]
            namecount[name] -= 1
    return list(zip(names, estimators))

这段代码实现了 make_pipeline/make_union 的自动命名逻辑:将类名转为小写作为步骤名,重名时从后向前添加 -N 后缀保证唯一性。测试 test_make_pipeline 验证了 make_pipeline(Transf(), Transf()) 产生 transf-1transf-2

便捷构造器

源码路径:sklearn/pipeline.py - make_pipeline()(第1207-1258行)

def make_pipeline(*steps, memory=None, transform_input=None, verbose=False):
    return Pipeline(
        _name_estimators(steps),
        transform_input=transform_input,
        memory=memory,
        verbose=verbose,
    )

make_pipeline 接收变长估计器参数,调用 _name_estimators 自动生成步骤名,再传递给 Pipeline 构造函数。同理 make_unionFeatureUnion 提供自动命名便捷构造。

输出容器递归设置

源码路径:sklearn/pipeline.py - Pipeline.set_output()(第228-257行)

def set_output(self, *, transform=None):
    for _, _, step in self._iter():
        _safe_set_output(step, transform=transform)
    return self

遍历 _iter() 产出的所有有效步骤(不含最终步骤),调用 _safe_set_output 递归设置输出容器,支持 default/pandas/polars 三种模式,为 DataFrame 输出提供统一入口。

15.5 步骤校验与迭代器 —— 流水线的“质量检查站与传送滚轮”

步骤能力校验

源码路径:sklearn/pipeline.py - Pipeline._validate_steps()(第278-310行)

def _validate_steps(self):
    if not self.steps:
        raise ValueError("The pipeline is empty. Please add steps.")
    names, estimators = zip(*self.steps)
    self._validate_names(names)
    self._check_estimators_are_instances(estimators)
    transformers = estimators[:-1]
    estimator = estimators[-1]
    for t in transformers:
        if t is None or t == "passthrough":
            continue
        if not (hasattr(t, "fit") or hasattr(t, "fit_transform")) or not hasattr(t, "transform"):
            raise TypeError(
                "All intermediate steps should be "
                "transformers and implement fit and transform "
                "or be the string 'passthrough' "
                "'%s' (type %s) doesn't" % (t, type(t))
            )
    if (
        estimator is not None
        and estimator != "passthrough"
        and not hasattr(estimator, "fit")
    ):
        raise TypeError(
            "Last step of Pipeline should implement fit "
            "or be the string 'passthrough'. "
            "'%s' (type %s) doesn't" % (estimator, type(estimator))
        )

这段代码实现了分层校验:中间步骤必须实现 fit/fit_transformtransform,或为 passthrough/None;最终步骤只需实现 fit 或为 passthrough/Nonetest_pipeline_invalid_parameters 验证了 NoTrans 作为中间步骤报错。

惰性迭代器

源码路径:sklearn/pipeline.py - Pipeline._iter()(第312-327行)

def _iter(self, with_final=True, filter_passthrough=True):
    stop = len(self.steps)
    if not with_final:
        stop -= 1
    for idx, (name, trans) in enumerate(islice(self.steps, 0, stop)):
        if not filter_passthrough:
            yield idx, name, trans
        elif trans is not None and trans != "passthrough":
            yield idx, name, trans

_iter 是核心遍历工具:with_final=False 排除最终步骤,供 fit 阶段遍历中间变换器;filter_passthrough=False 可遍历 passthrough 步骤,用于 verbose 日志完整展示。使用 islice 实现惰性切片,避免构建中间列表。

最终估计器保护性访问

源码路径:sklearn/pipeline.py - Pipeline._final_estimator(第360-369行)

@property
def _final_estimator(self):
    try:
        estimator = self.steps[-1][1]
        return "passthrough" if estimator is None else estimator
    except (ValueError, AttributeError, TypeError):
        return None

提供保护性获取,处理 steps 格式异常或未验证时的情况,返回 None 让后续 _available_if 抛出更清晰的错误。

15.6 切片与参数访问 —— Pipeline 的“乐高积木拆装术”

切片语义

源码路径:sklearn/pipeline.py - Pipeline.__getitem__()(第329-351行)

def __getitem__(self, ind):
    if isinstance(ind, slice):
        if ind.step not in (1, None):
            raise ValueError("Pipeline slicing only supports a step of 1")
        return self.__class__(
            self.steps[ind], memory=self.memory, verbose=self.verbose
        )
    try:
        name, est = self.steps[ind]
    except TypeError:
        return self.named_steps[ind]
    return est

整数索引返回对应步骤的估计器对象;切片返回新的 Pipeline,共享原步骤对象的引用(浅拷贝)。切片步长必须为 1,防止跳步导致语义不清。test_pipeline_slice 验证了切片后的步骤、参数、named_steps 与原管道一致。

属性式访问

源码路径:sklearn/pipeline.py - Pipeline.named_steps(第353-359行)

@property
def named_steps(self):
    return Bunch(**dict(self.steps))

使用 Bunch 而非 dict,支持属性式访问:pipe.named_steps.scalerpipe.named_steps["scaler"] 更简洁。测试验证了与 dict 的 keys/items 兼容性,以及 values 等保留属性冲突。

嵌套参数管理

源码路径:sklearn/pipeline.py - Pipeline.get_params()/set_params()(第259-285行)

def get_params(self, deep=True):
    return self._get_params("steps", deep=deep)

def set_params(self, **kwargs):
    self._set_params("steps", **kwargs)
    return self

继承自 _BaseComposition_get_params/_set_params,支持 svc__C=10svc=NewEstimator() 两种粒度。test_pipeline_raise_set_params_error 验证了错误参数名会给出精确的合法参数列表。

属性委派

源码路径:sklearn/pipeline.py - Pipeline.__len__()/n_features_in_/feature_names_in_/classes_(第321-324行、第936-944行、第932-934行)

def __len__(self):
    return len(self.steps)

@property
def n_features_in_(self):
    return self.steps[0][1].n_features_in_

@property
def feature_names_in_(self):
    return self.steps[0][1].feature_names_in_

@property
def classes_(self):
    return self.steps[-1][1].classes_

__len__ 返回步骤数量;n_features_in_/feature_names_in_ 委派给第一步;classes_ 委派给最终分类器。

15.7 fit 的执行引擎 —— 流水线的“顺序驱动马达”

核心拟合流程

源码路径:sklearn/pipeline.py - Pipeline._fit()(第507-552行)

def _fit(self, X, y=None, routed_params=None, raw_params=None):
    self.steps = list(self.steps)
    self._validate_steps()
    memory = check_memory(self.memory)
    fit_transform_one_cached = memory.cache(_fit_transform_one)

    for step_idx, name, transformer in self._iter(
        with_final=False, filter_passthrough=False
    ):
        if transformer is None or transformer == "passthrough":
            with _print_elapsed_time("Pipeline", self._log_message(step_idx)):
                continue

        if hasattr(memory, "location") and memory.location is None:
            cloned_transformer = transformer
        else:
            cloned_transformer = clone(transformer)
        step_params = self._get_metadata_for_step(
            step_idx=step_idx,
            step_params=routed_params[name],
            all_params=raw_params,
        )

        X, fitted_transformer = fit_transform_one_cached(
            cloned_transformer,
            X,
            y,
            weight=None,
            message_clsname="Pipeline",
            message=self._log_message(step_idx),
            params=step_params,
        )
        self.steps[step_idx] = (name, fitted_transformer)
    return X

这段代码实现了顺序拟合编排:先浅拷贝 steps 再验证,确保 fit 过程对 steps 的修改不影响原始参数;memory.cache 包裹 _fit_transform_one 实现变换器结果缓存;遍历中间步骤时,启用缓存时 clone 变换器(防止多流水线共享状态),不启用缓存时保留原对象维持向后兼容;每步处理后更新 self.steps[step_idx] 为已拟合的变换器。test_pipeline_memory 验证了缓存命中时 timestamp 不变。

最终估计器拟合

源码路径:sklearn/pipeline.py - Pipeline.fit()(第554-614行)

def fit(self, X, y=None, **params):
    if not _routing_enabled() and self.transform_input is not None:
        raise ValueError(...)
    routed_params = self._check_method_params(method="fit", props=params)
    Xt = self._fit(X, y, routed_params, raw_params=params)
    with _print_elapsed_time("Pipeline", self._log_message(len(self.steps) - 1)):
        if self._final_estimator != "passthrough":
            last_step_params = self._get_metadata_for_step(
                step_idx=len(self) - 1,
                step_params=routed_params[self.steps[-1][0]],
                all_params=params,
            )
            self._final_estimator.fit(Xt, y, **last_step_params["fit"])
    return self

fit 调用 _fit 处理中间步骤,再拟合最终估计器。_check_method_params 处理新旧路由模式分支(下节详解)。最终估计器为 passthrough 时跳过 fit。

fit_transform 与 fit_predict

源码路径:sklearn/pipeline.py - Pipeline.fit_transform()(第624-681行)、Pipeline.fit_predict()(第731-775行)

def fit_transform(self, X, y=None, **params):
    routed_params = self._check_method_params(method="fit_transform", props=params)
    Xt = self._fit(X, y, routed_params)
    last_step = self._final_estimator
    with _print_elapsed_time("Pipeline", self._log_message(len(self.steps) - 1)):
        if last_step == "passthrough":
            return Xt
        last_step_params = self._get_metadata_for_step(...)
        if hasattr(last_step, "fit_transform"):
            return last_step.fit_transform(Xt, y, **last_step_params["fit_transform"])
        else:
            return last_step.fit(Xt, y, **last_step_params["fit"]).transform(
                Xt, **last_step_params["transform"]
            )

def fit_predict(self, X, y=None, **params):
    routed_params = self._check_method_params(method="fit_predict", props=params)
    Xt = self._fit(X, y, routed_params)
    params_last_step = routed_params[self.steps[-1][0]]
    with _print_elapsed_time("Pipeline", self._log_message(len(self.steps) - 1)):
        y_pred = self.steps[-1][1].fit_predict(
            Xt, y, **params_last_step.get("fit_predict", {})
        )
    return y_pred

fit_transform 优先调用最终估计器的 fit_transform,否则回退 fit+transformfit_predict 调用 _fit 后执行最终估计器的 fit_predict_can_fit_transform 检查最终估计器是否支持 transformfit_transform 控制方法可见性。

单变换器拟合变换

源码路径:sklearn/pipeline.py - _fit_transform_one()(第1277-1299行)

def _fit_transform_one(
    transformer, X, y, weight, message_clsname="", message=None, params=None
):
    params = params or {}
    with _print_elapsed_time(message_clsname, message):
        if hasattr(transformer, "fit_transform"):
            res = transformer.fit_transform(X, y, **params.get("fit_transform", {}))
        else:
            res = transformer.fit(X, y, **params.get("fit", {})).transform(
                X, **params.get("transform", {})
            )
    if weight is None:
        return res, transformer
    return res * weight, transformer

优先使用 fit_transform,否则回退 fit+transform。返回变换结果和已拟合变换器,供 _fit 更新 steps。支持权重乘法(FeatureUnion 用)。

15.8 元数据路由的双模式 —— 从“参数前缀时代”到“智能分拣时代”

路由模式分支

源码路径:sklearn/pipeline.py - Pipeline._check_method_params()(第378-402行)

def _check_method_params(self, method, props, **kwargs):
    if _routing_enabled():
        routed_params = process_routing(self, method, **props, **kwargs)
        return routed_params
    else:
        fit_params_steps = Bunch(
            **{
                name: Bunch(**{method: {} for method in METHODS})
                for name, step in self.steps
                if step is not None
            }
        )
        for pname, pval in props.items():
            if "__" not in pname:
                raise ValueError(
                    "Pipeline.fit does not accept the {} parameter. "
                    "You can pass parameters to specific steps of your "
                    "pipeline using the stepname__parameter format, e.g. "
                    "`Pipeline.fit(X, y, logisticregression__sample_weight"
                    "=sample_weight)`.".format(pname)
                )
            step, param = pname.split("__", 1)
            fit_params_steps[step]["fit"][param] = pval
            fit_params_steps[step]["fit_transform"][param] = pval
            fit_params_steps[step]["fit_predict"][param] = pval
        return fit_params_steps

这是新旧路由的分水岭:_routing_enabled() 为 True 时调用 process_routing 返回结构化 Bunch;旧模式下将 s__p 格式参数拆分到对应步骤的 fit/fit_transform/fit_predict 桶中,旧模式不支持裸参数(如 sample_weight),会抛出带格式提示的 ValueError。test_pipeline_param_error 验证了路由禁用时传递 sample_weight 报错。

预测方法的双模式分支

源码路径:sklearn/pipeline.py - Pipeline.predict()(第684-723行)

def predict(self, X, **params):
    check_is_fitted(self)
    Xt = X
    if not _routing_enabled():
        for _, name, transform in self._iter(with_final=False):
            Xt = transform.transform(Xt)
        return self.steps[-1][1].predict(Xt, **params)

    routed_params = process_routing(self, "predict", **params)
    for _, name, transform in self._iter(with_final=False):
        Xt = transform.transform(Xt, **routed_params[name].transform)
    return self.steps[-1][1].predict(Xt, **routed_params[self.steps[-1][0]].predict)

旧模式直接遍历中间步骤调 transform,不带额外参数;新模式向每个步骤的 transform 传递 routed_params[name].transformpredict_proba/predict_log_proba/score 同理。decision_function 仅支持新路由模式(1.4 新增)。

score 方法的双模式逻辑

源码路径:sklearn/pipeline.py - Pipeline.score()(第892-930行)

def score(self, X, y=None, sample_weight=None, **params):
    check_is_fitted(self)
    Xt = X
    if not _routing_enabled():
        for _, name, transform in self._iter(with_final=False):
            Xt = transform.transform(Xt)
        score_params = {}
        if sample_weight is not None:
            score_params["sample_weight"] = sample_weight
        return self.steps[-1][1].score(Xt, y, **score_params)

    routed_params = process_routing(
        self, "score", sample_weight=sample_weight, **params
    )
    Xt = X
    for _, name, transform in self._iter(with_final=False):
        Xt = transform.transform(Xt, **routed_params[name].transform)
    return self.steps[-1][1].score(Xt, y, **routed_params[self.steps[-1][0]].score)

旧模式下 sample_weight 单独处理;新模式下通过 process_routing 统一路由。test_pipeline_sample_weight_supported/unsupported 验证了 sample_weight 的兼容行为。

路由结构示意

graph TD A[用户调用 fit/predict/score] --> B{_routing_enabled?} B -- False --> C[解析 s__p 前缀参数] C --> D[构建 fit_params_steps Bunch] B -- True --> E[调用 process_routing] E --> F[返回 routed_params Bunch] F --> G[各步骤按键名获取参数] D --> H[_fit/_predict 使用对应参数] G --> H

15.9 transform_input 元数据变换 —— 给“元数据也过一次流水线”

元数据变换入口

源码路径:sklearn/pipeline.py - Pipeline._get_metadata_for_step()(第404-467行)

def _get_metadata_for_step(self, *, step_idx, step_params, all_params):
    if (
        self.transform_input is None
        or not all_params
        or not step_params
        or step_idx == 0
    ):
        return step_params

    sub_pipeline = self[:step_idx]
    sub_metadata_routing = get_routing_for_object(sub_pipeline)
    transform_params = {
        key: value
        for key, value in all_params.items()
        if key
        in sub_metadata_routing.consumes(
            method="transform", params=all_params.keys()
        )
    }
    transformed_params = dict()
    transformed_cache = dict()
    for method, method_params in step_params.items():
        transformed_params[method] = Bunch()
        for param_name, param_value in method_params.items():
            if param_name in self.transform_input:
                transformed_params[method][param_name] = _cached_transform(
                    sub_pipeline,
                    cache=transformed_cache,
                    param_name=param_name,
                    param_value=param_value,
                    transform_params=transform_params,
                )
            else:
                transformed_params[method][param_name] = param_value
    return transformed_params

这段代码实现了 transform_input 机制:仅在 transform_input 非空、有参数传入、且 step_idx > 0 时处理。获取子流水线的元数据路由信息,从 all_params 中筛选子流水线 transform 会消费的参数作为 transform_params。对 step_params 中每个方法的每个参数,若参数名在 transform_input 列表中,则用 _cached_transform 变换。

缓存变换实现

源码路径:sklearn/pipeline.py - _cached_transform()(第40-70行)

def _cached_transform(
    sub_pipeline, *, cache, param_name, param_value, transform_params
):
    if param_name not in cache:
        if isinstance(param_value, tuple):
            cache[param_name] = tuple(
                sub_pipeline.transform(element, **transform_params)
                for element in param_value
            )
        else:
            cache[param_name] = sub_pipeline.transform(param_value, **transform_params)
    return cache[param_name]

缓存避免同一元数据被重复变换。若参数值为 tuple,逐元素变换后重组,支持多验证集模式(如 lightgbm/xgboost 的 eval_set)。测试 test_transform_tuple_input 验证了 tuple 中每个数组都被独立变换。

transform_input 测试验证

测试 test_transform_input_explicit_value_check 定义了 Transformer.transform = X + 1,最终估计器请求 X_val,Pipeline transform_input=["X_val"]。输入 X=[[0,1]]X_val=[[1,2]],断言最终估计器收到的 X_val=[[2,3]](经过 +1 变换)。这证明了元数据也经过了前置变换器处理。

15.10 FeatureUnion 的并行拼接 —— 把“多路特征流汇成一条大河”

迭代器与权重

源码路径:sklearn/pipeline.py - FeatureUnion._iter()(第1332-1342行)

def _iter(self):
    get_weight = (self.transformer_weights or {}).get
    for name, trans in self.transformer_list:
        if trans == "drop":
            continue
        if trans == "passthrough":
            trans = FunctionTransformer(feature_names_out="one-to-one")
        yield (name, trans, get_weight(name))

drop 直接跳过;passthrough 替换为 FunctionTransformer(feature_names_out="one-to-one");返回 (name, trans, weight) 三元组,权重来自 transformer_weights

并行调度

源码路径:sklearn/pipeline.py - FeatureUnion._parallel_func()(第1396-1415行)

def _parallel_func(self, X, y, func, routed_params):
    self.transformer_list = list(self.transformer_list)
    self._validate_transformers()
    self._validate_transformer_weights()
    transformers = list(self._iter())
    return Parallel(n_jobs=self.n_jobs)(
        delayed(func)(
            transformer,
            X,
            y,
            weight,
            message_clsname="FeatureUnion",
            message=self._log_message(name, idx, len(transformers)),
            params=routed_params[name],
        )
        for idx, (name, transformer, weight) in enumerate(transformers, 1)
    )

使用 joblib 的 Paralleldelayed 对每个变换器并行执行同一 func_fit_one/_fit_transform_one/_transform_one)。每个任务接收独立的 routed_params[name]test_feature_union_parallel 验证了并行一致性。

三策略拼接

源码路径:sklearn/pipeline.py - FeatureUnion._hstack()(第1432-1450行)

def _hstack(self, Xs):
    xp, _ = get_namespace(*Xs)
    for X, (name, _) in zip(Xs, self.transformer_list):
        if hasattr(X, "shape") and len(X.shape) != 2:
            raise ValueError(...)
    adapter = _get_container_adapter("transform", self)
    if adapter and all(adapter.is_supported_container(X) for X in Xs):
        return adapter.hstack(Xs, self.get_feature_names_out())
    if any(sparse.issparse(f) for f in Xs):
        return sparse.hstack(Xs).tocsr()
    return xp.concat(Xs, axis=1)

拼接策略优先级:

  1. 容器适配器优先:配置了 pandas/polars 输出且所有输出都支持容器适配器,使用 adapter.hstack 保留 DataFrame 类型

  2. 稀疏矩阵:有稀疏矩阵输出,使用 scipy.sparse.hstack 返回 CSR

  3. Array API:否则使用 xp.concat 沿 axis=1 拼接,支持 NumPy/CuPy/PyTorch 等后端

更新已拟合变换器

源码路径:sklearn/pipeline.py - FeatureUnion._update_transformer_list()(第1452-1456行)

def _update_transformer_list(self, transformers):
    transformers = iter(transformers)
    self.transformer_list[:] = [
        (name, old if old == "drop" else next(transformers))
        for name, old in self.transformer_list
    ]

将并行任务返回的已拟合变换器放回 transformer_listdrop 步骤保持不变(不消费迭代器元素)。fit/fit_transform/transform 入口分别调用 _fit_one/_fit_transform_one/_transform_one 并行执行,fit_transform 后 hstack 返回拼接结果。

15.11 特征名生成与冲突检测 —— 给“拼接后的特征贴上专属标签”

收集特征名

源码路径:sklearn/pipeline.py - FeatureUnion.get_feature_names_out()(第1344-1366行)

def get_feature_names_out(self, input_features=None):
    transformer_with_feature_names_out = []
    for name, trans, _ in self._iter():
        if not hasattr(trans, "get_feature_names_out"):
            raise AttributeError(...)
        feature_names_out = trans.get_feature_names_out(input_features)
        transformer_with_feature_names_out.append((name, feature_names_out))
    return self._add_prefix_for_feature_names_out(transformer_with_feature_names_out)

遍历 _iter 结果,要求每个变换器都有 get_feature_names_out 方法,收集 (name, feature_names_out) 列表后交给 _add_prefix_for_feature_names_out

前缀模式与冲突检测

源码路径:sklearn/pipeline.py - FeatureUnion._add_prefix_for_feature_names_out()(第1368-1393行)

def _add_prefix_for_feature_names_out(self, transformer_with_feature_names_out):
    if self.verbose_feature_names_out:
        names = list(
            chain.from_iterable(
                (f"{name}__{feat}" for feat in feature_names_out)
                for name, feature_names_out in transformer_with_feature_names_out
            )
        )
        return np.asarray(names, dtype=object)

    feature_names_count = Counter(
        chain.from_iterable(s for _, s in transformer_with_feature_names_out)
    )
    top_6_overlap = [
        name for name, count in feature_names_count.most_common(6) if count > 1
    ]
    top_6_overlap.sort()
    if top_6_overlap:
        if len(top_6_overlap) == 6:
            names_repr = str(top_6_overlap[:5])[:-1] + ", ...]"
        else:
            names_repr = str(top_6_overlap)
        raise ValueError(
            f"Output feature names: {names_repr} are not unique. Please set "
            "verbose_feature_names_out=True to add prefixes to feature names"
        )
    return np.concatenate(
        [name for _, name in transformer_with_feature_names_out],
    )

verbose_feature_names_out=True(默认)时,每个特征名加前缀 f"{name}__",如 pca__pca0;为 False 时,用 Counter 统计重名特征,取 top 6 展示,超过 5 个重名显示前 5 个加省略号,抛出 ValueError 提示用户开启前缀模式。test_feature_union_passthrough_get_feature_names_out_false_errors_overlap_over_5 验证了 10 个重名特征时错误消息展示前 5 个。

特征名生成与冲突检测流程

flowchart TD A[get_feature_names_out 调用] --> B[遍历 _iter 收集<br/> (name, feature_names_out)] B --> C{verbose_feature_names_out?} C -- True --> D[加前缀模式<br/> f'{name}__{feat}'] D --> E[返回带前缀的特征名数组] C -- False --> F[无前缀模式<br/> 统计所有特征名频次] F --> G{有重名特征?} G -- 否 --> H[直接拼接返回] G -- 是 --> I[取 top 6 重名] I --> J{重名数 > 5?} J -- 是 --> K[截断显示前 5 个 + 省略号] J -- 否 --> L[显示全部重名] K --> M[抛出 ValueError<br/> 提示开启前缀模式] L --> M

15.12 元数据路由的组装 —— Pipeline 的“物流分拣系统说明书”

Pipeline 路由映射

源码路径:sklearn/pipeline.py - Pipeline.get_metadata_routing()(第1016-1073行)

def get_metadata_routing(self):
    router = MetadataRouter(owner=self)
    for _, name, trans in self._iter(with_final=False, filter_passthrough=True):
        method_mapping = MethodMapping()
        if hasattr(trans, "fit_transform"):
            (
                method_mapping.add(caller="fit", callee="fit_transform")
                .add(caller="fit_transform", callee="fit_transform")
                .add(caller="fit_predict", callee="fit_transform")
            )
        else:
            (
                method_mapping.add(caller="fit", callee="fit")
                .add(caller="fit", callee="transform")
                .add(caller="fit_transform", callee="fit")
                .add(caller="fit_transform", callee="transform")
                .add(caller="fit_predict", callee="fit")
                .add(caller="fit_predict", callee="transform")
            )
        (
            method_mapping.add(caller="predict", callee="transform")
            .add(caller="predict", callee="transform")
            .add(caller="predict_proba", callee="transform")
            .add(caller="decision_function", callee="transform")
            .add(caller="predict_log_proba", callee="transform")
            .add(caller="transform", callee="transform")
            .add(caller="inverse_transform", callee="inverse_transform")
            .add(caller="score", callee="transform")
        )
        router.add(method_mapping=method_mapping, **{name: trans})

    final_name, final_est = self.steps[-1]
    if final_est is None or final_est == "passthrough":
        return router

    method_mapping = MethodMapping()
    if hasattr(final_est, "fit_transform"):
        method_mapping.add(caller="fit_transform", callee="fit_transform")
    else:
        method_mapping.add(caller="fit", callee="fit").add(
            caller="fit", callee="transform"
        )
    (
        method_mapping.add(caller="fit", callee="fit")
        .add(caller="predict", callee="predict")
        .add(caller="fit_predict", callee="fit_predict")
        .add(caller="predict_proba", callee="predict_proba")
        .add(caller="decision_function", callee="decision_function")
        .add(caller="predict_log_proba", callee="predict_log_proba")
        .add(caller="transform", callee="transform")
        .add(caller="inverse_transform", callee="inverse_transform")
        .add(caller="score", callee="score")
    )
    router.add(method_mapping=method_mapping, **{final_name: final_est})
    return router

中间步骤映射:若有 fit_transform,则 fit/fit_transform/fit_predict 都映射到 fit_transform;否则展开为 fit+transform 组合。所有预测类方法(predict/predict_proba/decision_function/predict_log_proba/transform/score)均映射到 transform。最终步骤映射:fit_transform 优先映射到自身,其余方法与自己同名映射。

FeatureUnion 路由映射

源码路径:sklearn/pipeline.py - FeatureUnion.get_metadata_routing()(第1494-1514行)

def get_metadata_routing(self):
    router = MetadataRouter(owner=self)
    for name, transformer in self.transformer_list:
        router.add(
            **{name: transformer},
            method_mapping=MethodMapping()
            .add(caller="fit", callee="fit")
            .add(caller="fit_transform", callee="fit_transform")
            .add(caller="fit_transform", callee="fit")
            .add(caller="fit_transform", callee="transform")
            .add(caller="transform", callee="transform"),
        )
    return router

FeatureUnion 只关心 fitfit_transformtransform 三个入口。fit_transform 可能映射到 fit_transformfit+transform 组合,取决于变换器是否实现 fit_transform

Pipeline 与 FeatureUnion 路由映射构建差异

flowchart TD subgraph Pipeline_routing["Pipeline.get_metadata_routing()"] P1[创建 MetadataRouter] P2[遍历中间步骤 _iter(with_final=False)] P3{变换器有 fit_transform?} P3 -- 是 --> P4[fit/fit_transform/fit_predict<br/>均映射到 fit_transform] P3 -- 否 --> P5[fit 映射到 fit+transform<br/>fit_transform 映射到 fit+transform<br/>fit_predict 映射到 fit+transform] P4 --> P6[预测类方法均映射到 transform] P5 --> P6 P6 --> P7[添加到 router] P7 --> P8[处理最终步骤] P8 --> P9{最终估计器有 fit_transform?} P9 -- 是 --> P10[fit_transform 映射到 fit_transform] P9 -- 否 --> P11[fit 映射到 fit+transform] P10 --> P12[其余方法同名映射] P11 --> P12 P12 --> P13[返回 router] end subgraph FeatureUnion_routing["FeatureUnion.get_metadata_routing()"] F1[创建 MetadataRouter] F2[遍历 transformer_list] F3{变换器有 fit_transform?} F3 -- 是 --> F4[fit_transform 映射到 fit_transform] F3 -- 否 --> F5[fit_transform 映射到 fit+transform] F4 --> F6[fit 映射到 fit<br/>transform 映射到 transform] F5 --> F6 F6 --> F7[添加到 router] F7 --> F8[返回 router] end

15.13 标签系统与拟合状态判定 —— 流水线的“能力体检报告”

标签传播

源码路径:sklearn/pipeline.py - Pipeline.__sklearn_tags__()(第876-910行)

def __sklearn_tags__(self):
    tags = super().__sklearn_tags__()
    if not self.steps:
        return tags
    try:
        if self.steps[0][1] is not None and self.steps[0][1] != "passthrough":
            tags.input_tags.pairwise = get_tags(self.steps[0][1]).input_tags.pairwise
        tags.input_tags.sparse = all(
            get_tags(step).input_tags.sparse
            for name, step in self.steps
            if step is not None and step != "passthrough"
        )
    except (ValueError, AttributeError, TypeError):
        pass
    try:
        if self.steps[-1][1] is not None and self.steps[-1][1] != "passthrough":
            last_step_tags = get_tags(self.steps[-1][1])
            tags.estimator_type = last_step_tags.estimator_type
            tags.target_tags.multi_output = last_step_tags.target_tags.multi_output
            tags.classifier_tags = deepcopy(last_step_tags.classifier_tags)
            tags.regressor_tags = deepcopy(last_step_tags.regressor_tags)
            tags.transformer_tags = deepcopy(last_step_tags.transformer_tags)
    except (ValueError, AttributeError, TypeError):
        pass
    return tags

从第一个非 passthrough 步骤继承 pairwise 标签;sparse 标签取所有非 passthrough 步骤的 AND 结果(有警告注释:可能不准确,如 PCA 输出稠密);从最终估计器继承 estimator_typemulti_output、分类/回归/变换器标签。用 try-except 捕获 steps 格式异常,避免未验证时崩溃。

拟合状态判定

源码路径:sklearn/pipeline.py - Pipeline.__sklearn_is_fitted__()(第946-972行)

def __sklearn_is_fitted__(self):
    last_step = None
    for _, estimator in reversed(self.steps):
        if estimator != "passthrough":
            last_step = estimator
            break
    if last_step is None:
        return True
    try:
        check_is_fitted(last_step)
        return True
    except NotFittedError:
        return False

从后向前找第一个非 passthrough 步骤,只对最后一个非 passthrough 步骤做 check_is_fitted,省略前置检查以提速。全是 passthrough 视为已拟合。FeatureUnion 则遍历所有有效变换器逐个检查。

标签传播路径与拟合状态判定流程

flowchart TD subgraph Tags["__sklearn_tags__ 标签传播"] T1[super().__sklearn_tags__ 基础标签] T2{steps 为空?} -- 是 --> T3[直接返回] T2 -- 否 --> T4[首步非 passthrough?] T4 -- 是 --> T5[继承 pairwise 标签] T4 -- 否 --> T6[跳过 pairwise] T6 --> T7[所有非 passthrough 步骤<br/>sparse 标签 AND 运算] T7 --> T8[末步非 passthrough?] T8 -- 是 --> T9[继承 estimator_type<br/>multi_output<br/>classifier/regressor/transformer 标签] T8 -- 否 --> T10[跳过末步标签] T3 & T9 & T10 --> T11[返回 tags] end subgraph Fitted["__sklearn_is_fitted__ 拟合状态判定"] F1[从后向前遍历 steps] F2{找到非 passthrough 步骤?} F2 -- 否 --> F3[全是 passthrough<br/>返回 True] F2 -- 是 --> F4[仅对该步骤调用 check_is_fitted] F4 --> F5{已拟合?} F5 -- 是 --> F6[返回 True] F5 -- 否 --> F7[返回 False] end

15.14 预测与变换方法全景 —— Pipeline 的“方法调用总调度台”

预测方法的双模式分支

前文已分析 predict/predict_proba/predict_log_proba/decision_function:旧模式直接遍历中间步骤 transform,最终步骤调用对应方法;新模式通过 process_routing 向中间步骤传递 transform 参数。

transform 与 inverse_transform

源码路径:sklearn/pipeline.py - Pipeline.transform()(第845-873行)、Pipeline.inverse_transform()(第878-910行)

def transform(self, X, **params):
    check_is_fitted(self)
    _raise_for_params(params, self, "transform")
    routed_params = process_routing(self, "transform", **params)
    Xt = X
    for _, name, transform in self._iter():
        Xt = transform.transform(Xt, **routed_params[name].transform)
    return Xt

def inverse_transform(self, X, **params):
    check_is_fitted(self)
    _raise_for_params(params, self, "inverse_transform")
    routed_params = process_routing(self, "inverse_transform", **params)
    reverse_iter = reversed(list(self._iter()))
    for _, name, transform in reverse_iter:
        X = transform.inverse_transform(X, **routed_params[name].inverse_transform)
    return X

transform 遍历全部步骤(含最终步骤)执行 transforminverse_transform 逆序遍历全部步骤执行 inverse_transform_can_transform/_can_inverse_transformavailable_if 控制方法可见性。

transform/inverse_transform 正向/逆向遍历差异

sequenceDiagram participant User as 用户调用 participant Pipeline as Pipeline participant Steps as 步骤列表 Note over User,Steps: transform 正向遍历 User->>Pipeline: transform(X) Pipeline->>Pipeline: check_is_fitted Pipeline->>Pipeline: process_routing(transform) loop 正序遍历 _iter() Pipeline->>Steps: transform(Xt, **routed_params[name].transform) Steps-->>Pipeline: 返回变换后数据 end Pipeline-->>User: 返回最终结果 Note over User,Steps: inverse_transform 逆向遍历 User->>Pipeline: inverse_transform(X) Pipeline->>Pipeline: check_is_fitted Pipeline->>Pipeline: process_routing(inverse_transform) loop 逆序遍历 reversed(_iter()) Pipeline->>Steps: inverse_transform(X, **routed_params[name].inverse_transform) Steps-->>Pipeline: 返回逆变换数据 end Pipeline-->>User: 返回原始空间数据

score_samples 与 get_feature_names_out

源码路径:sklearn/pipeline.py - Pipeline.score_samples()(第837-865行)、Pipeline.get_feature_names_out()(第912-934行)

def score_samples(self, X):
    check_is_fitted(self)
    Xt = X
    for _, _, transformer in self._iter(with_final=False):
        Xt = transformer.transform(Xt)
    return self.steps[-1][1].score_samples(Xt)

def get_feature_names_out(self, input_features=None):
    feature_names_out = input_features
    for _, name, transform in self._iter():
        if not hasattr(transform, "get_feature_names_out"):
            raise AttributeError(
                "Estimator {} does not provide get_feature_names_out. "
                "Did you mean to call pipeline[:-1].get_feature_names_out"
                "()?".format(name)
            )
        feature_names_out = transform.get_feature_names_out(feature_names_out)
    return feature_names_out

score_samples 无路由分支,直接变换后调用最终步骤;get_feature_names_out 链式调用每个步骤的 get_feature_names_out,缺少方法时给出带步骤名的 AttributeError。

可视化块

源码路径:sklearn/pipeline.py - Pipeline._sk_visual_block_()(第974-994行)

def _sk_visual_block_(self):
    def _get_name(name, est):
        if est is None or est == "passthrough":
            return f"{name}: passthrough"
        return f"{name}: {est.__class__.__name__}"
    names, estimators = zip(
        *[(_get_name(name, est), est) for name, est in self.steps]
    )
    name_details = [str(est) for est in estimators]
    return _VisualBlock(
        "serial",
        estimators,
        names=names,
        name_details=name_details,
        dash_wrapped=False,
    )

构造 serial 类型的 VisualBlock,用于 Jupyter HTML 可视化展示。

15.15 元数据路由测试基础设施 —— 构建“分拣系统的测试车间”

注册表与深拷贝

源码路径:sklearn/tests/metadata_routing_common.py - _Registry(第135-145行)

class _Registry(list):
    def __deepcopy__(self, memo):
        return self
    def __copy__(self):
        return self

重写 __deepcopy____copy__ 返回自身,深拷贝时保持同一列表引用,用于追踪克隆后的子估计器。ConsumingTransformer 构造时接收 registryfitself.registry.append(self) 记录自身引用。

元数据记录与断言

源码路径:sklearn/tests/metadata_routing_common.py - record_metadata()/check_recorded_metadata()(第35-84行)

def record_metadata(obj, record_default=True, **kwargs):
    stack = inspect.stack()
    callee = stack[1].function
    caller = stack[2].function
    if not hasattr(obj, "_records"):
        obj._records = defaultdict(lambda: defaultdict(list))
    if not record_default:
        kwargs = {k: v for k, v in kwargs.items() if not isinstance(v, str) or v != "default"}
    obj._records[callee][caller].append(kwargs)

def check_recorded_metadata(obj, method, parent, split_params=tuple(), **kwargs):
    all_records = getattr(obj, "_records", dict()).get(method, dict()).get(parent, list())
    for record in all_records:
        assert set(kwargs.keys()) == set(record.keys())
        for key, value in kwargs.items():
            recorded_value = record[key]
            if key in split_params and recorded_value is not None:
                assert np.isin(recorded_value, value).all()
            else:
                if isinstance(recorded_value, np.ndarray):
                    assert_array_equal(recorded_value, value)
                else:
                    assert recorded_value is value

record_metadata 使用 inspect.stack 记录调用层级(callee/caller);check_recorded_metadata 验证参数名和值完全匹配,支持 split_params 检查子集关系。

测试辅助估计器

源码路径:sklearn/tests/metadata_routing_common.py - ConsumingTransformer(第196-226行)

class ConsumingTransformer(TransformerMixin, BaseEstimator):
    def __init__(self, registry=None):
        self.registry = registry
    def fit(self, X, y=None, sample_weight="default", metadata="default"):
        if self.registry is not None:
            self.registry.append(self)
        record_metadata_not_default(self, sample_weight=sample_weight, metadata=metadata)
        self.fitted_ = True
        return self
    def transform(self, X, sample_weight="default", metadata="default"):
        record_metadata_not_default(self, sample_weight=sample_weight, metadata=metadata)
        return X + 1
    def fit_transform(self, X, y, sample_weight="default", metadata="default"):
        record_metadata_not_default(self, sample_weight=sample_weight, metadata=metadata)
        return self.fit(X, y, sample_weight=sample_weight, metadata=metadata).transform(
            X, sample_weight=sample_weight, metadata=metadata
        )

手动实现 fit_transform 是必要的,因为 TransformerMixin.fit_transform 不路由元数据到 transform,而这里需要 transform 也收到 sample_weightmetadataConsumingNoFitTransformTransformer 不继承 TransformerMixin,无 fit_transform,用于测试无 fit_transform 时的路由行为。

极简元估计器

源码路径:sklearn/tests/test_metadata_routing.py - SimplePipeline(第33-82行)

class SimplePipeline(BaseEstimator):
    def __init__(self, steps):
        self.steps = steps
    def fit(self, X, y, **fit_params):
        self.steps_ = []
        params = process_routing(self, "fit", **fit_params)
        X_transformed = X
        for i, step in enumerate(self.steps[:-1]):
            transformer = clone(step).fit(
                X_transformed, y, **params.get(f"step_15").fit
            )
            self.steps_.append(transformer)
            X_transformed = transformer.transform(
                X_transformed, **params.get(f"step_15").transform
            )
        self.steps_.append(
            clone(self.steps[-1]).fit(X_transformed, y, **params.predictor.fit)
        )
        return self
    def get_metadata_routing(self):
        router = MetadataRouter(owner=self)
        for i, step in enumerate(self.steps[:-1]):
            router.add(
                **{f"step_15": step},
                method_mapping=MethodMapping()
                .add(caller="fit", callee="fit")
                .add(caller="fit", callee="transform")
                .add(caller="predict", callee="transform"),
            )
        router.add(
            predictor=self.steps[-1],
            method_mapping=MethodMapping()
            .add(caller="fit", callee="fit")
            .add(caller="predict", callee="predict"),
        )
        return router

手动实现路由分发,用于测试 process_routing 的嵌套路由能力。test_nested_routing 验证了嵌套路由的参数传递链。

_Registry 深拷贝共享引用机制与 record_metadata 调用栈记录层级

flowchart TD subgraph Registry["_Registry 深拷贝共享引用"] R1[_Registry 继承 list] R2[__deepcopy__ 返回 self] R3[__copy__ 返回 self] R4[克隆子估计器时 registry 保持同一引用] R5[ConsumingTransformer.fit 时<br/>self.registry.append(self)] R6[测试后检查 registry 内容<br/>验证子估计器被正确克隆与调用] end subgraph Record["record_metadata 调用栈记录层级"] M1[inspect.stack()[1].function -> callee] M2[inspect.stack()[2].function -> caller] M3[obj._records[callee][caller].append(kwargs)] M4[记录方法调用层级关系<br/>如 fit->fit, transform->fit] end Registry --> Record

15.16 测试辅助类与边界行为 —— 验证 Pipeline 的“压力测试工具箱”

核心辅助类

源码路径:sklearn/tests/test_pipeline.py - Mult/Transf/FitParamT 等(第75-175行)

class Mult(TransformerMixin, BaseEstimator):
    def __init__(self, mult=1):
        self.mult = mult
    def __sklearn_is_fitted__(self):
        return True
    def fit(self, X, y=None):
        return self
    def transform(self, X):
        return np.asarray(X) * self.mult
    def inverse_transform(self, X):
        return np.asarray(X) / self.mult
    def predict(self, X):
        return (np.asarray(X) * self.mult).sum(axis=1)
    # ... 多种方法别名

class FitParamT(BaseEstimator):
    def __init__(self):
        self.successful = False
    def fit(self, X, y, should_succeed=False):
        self.successful = should_succeed
        self.fitted_ = True
    def predict(self, X):
        return self.successful
    def fit_predict(self, X, y, should_succeed=False):
        self.fit(X, y, should_succeed=should_succeed)
        return self.predict(X)

Mult 验证乘法变换、加权与 inverse_transformTransf 系列验证可逆变换与参数传递;FitParamT 验证 fit_predictscore 行为;DummyTransf.timestamp_ 用于检测缓存命中。

passthrough 与步骤替换

源码路径:sklearn/pipeline.py - Pipeline._iter() 配合测试 test_set_pipeline_step_passthrough

def _iter(self, with_final=True, filter_passthrough=True):
    stop = len(self.steps)
    if not with_final:
        stop -= 1
    for idx, (name, trans) in enumerate(islice(self.steps, 0, stop)):
        if not filter_passthrough:
            yield idx, name, trans
        elif trans is not None and trans != "passthrough":
            yield idx, name, trans

_iterfilter_passthrough 参数控制遍历:fit 阶段 filter_passthrough=False 遍历所有步骤(含 passthrough),遇到 passthrough 直接 continue;transform 阶段 filter_passthrough=True 跳过 passthrough。测试验证了各位置的 passthrough 替换后,fit_transform/predict/inverse_transform 结果变化。

缓存与克隆交互

测试 test_pipeline_memory:启用缓存后 DummyTransf.timestamp_ 不变,证明缓存命中跳过了 fit。_fithasattr(memory, "location") and memory.location is None 判断缓存是否启用,不启用缓存时不 clone,维持向后兼容。

Array API 与 set_output 集成

测试 test_feature_union_array_api_compliance:在 config_context(array_api_dispatch=True) 下验证 FeatureUnion 跨后端行为一致。test_pipeline_set_output_integration 验证 DataFrame 特征名链:pipe[:-1].get_feature_names_out() 与最终分类器 feature_names_in_ 一致。

passthrough 替换前后 _iter 遍历差异与缓存命中判定流程

flowchart TD subgraph Iter_Diff["passthrough 替换前后 _iter 遍历差异"] I1[_iter(with_final=False, filter_passthrough=False)] I2[遍历所有步骤含 passthrough] I3[遇到 passthrough/None -> continue] I4[_iter(with_final=True, filter_passthrough=True)] I5[跳过 passthrough/None 只遍历有效步骤] I1 --> I2 --> I3 I4 --> I5 end subgraph Cache["缓存命中判定流程"] C1[_fit 开始] C2{memory.location is None?} C2 -- 是 --> C3[不启用缓存<br/>cloned_transformer = transformer<br/>保持原对象引用] C2 -- 否 --> C4[启用缓存<br/>cloned_transformer = clone(transformer)<br/>防止多流水线共享状态] C3 --> C5[执行 fit_transform_one_cached] C4 --> C5 C5 --> C6{cache 命中?} C6 -- 是 --> C7[跳过 fit 直接返回缓存结果<br/>timestamp_ 不变] C6 -- 否 --> C8[执行 fit_transform<br/>timestamp_ 更新] end
posted @ 2026-09-04 04:08  绝不原创的飞龙  阅读(4)  评论(0)    收藏  举报