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

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

要点predict_proba 用逆链接 expit 恢复 类别 1 的概率,然后 1‑p 得到 类别 0 的概率。这在二分类预测接口(如 predict_proba)中至关重要。

63.6.6 代码解读:HalfMultinomialLoss(第 687‑780 行)

class HalfMultinomialLoss(BaseLoss):
    is_multiclass = True                                             # 标记为多分类任务

    def __init__(self, sample_weight=None, n_classes=3):
        super().__init__(closs=CyHalfMultinomialLoss(),              # 使用 Cython 实现的半多项损失
                         link=MultinomialLogit(),                    # 使用对数几率链接(对称 softmax)
                         n_classes=n_classes)                        # 设置类别数
        self.interval_y_true = Interval(0, np.inf, True, False)      # y_true 为非负整数(标签编码)
        self.interval_y_pred = Interval(0, 1, False, False)          # y_pred 每元素在 (0, 1) 开区间
        self.class_indexing_offsets = None                           # 用于 Array API 的索引偏移
        self.y_true_int = None                                       # 用于 Array API 的整数标签
        self.y_true_one_hot = None                                   # 用于 Array API 的 one-hot 编码

    def in_y_true_range(self, y):
        return self.interval_y_true.includes(y) and np.all(y.astype(int) == y)  # 检查 y_true 为整数且在区间内

    def fit_intercept_only(self, y_true, sample_weight=None):
        out = np.zeros(self.n_classes, dtype=y_true.dtype)           # 初始化输出数组
        eps = np.finfo(y_true.dtype).eps                             # 数值稳定性微小偏移
        for k in range(self.n_classes):                              # 遍历每个类别
            out[k] = np.average(y_true == k, weights=sample_weight, axis=0)  # 计算加权频率
            out[k] = np.clip(out[k], eps, 1 - eps)                   # 裁剪到 [eps, 1-eps] 避免极端值
        return self.link.link(out[None, :]).reshape(-1)              # 应用链接函数并展平

    def predict_proba(self, raw_prediction):
        return self.link.inverse(raw_prediction)                     # 直接返回 softmax 概率

    def gradient_proba(self, y_true, raw_prediction,
                       sample_weight=None,
                       gradient_out=None, proba_out=None,
                       n_threads=1):
        if gradient_out is None:                                     # 处理梯度输出
            if proba_out is None:                                    # 且未提供概率输出
                gradient_out = np.empty_like(raw_prediction)         # 创建梯度数组
                proba_out = np.empty_like(raw_prediction)            # 创建概率数组
            else:                                                    # 仅提供了概率输出
                gradient_out = np.empty_like(proba_out)              # 梯度数组 dtype/shape 随概率数组对齐
        elif proba_out is None:                                      # 仅提供了梯度输出
            proba_out = np.empty_like(gradient_out)                  # 概率数组 dtype/shape 随梯度数组对齐

        self.closs.gradient_proba(y_true=y_true,                     # 调用 Cython 实现计算梯度与概率
                                  raw_prediction=raw_prediction,
                                  sample_weight=sample_weight,
                                  gradient_out=gradient_out,
                                  proba_out=proba_out,
                                  n_threads=n_threads)
        return gradient_out, proba_out                               # 返回梯度与概率

要点MultinomialLogit向量化的概率(每行之和为 1)映射到 中心化的原始预测(行和为 0),从而保证模型参数可唯一识别。fit_intercept_only 计算每个类别的 加权频率,再对数化并减去几何均值,实现 截距模型 的最小化。gradient_proba 同时返回 梯度softmax 概率,在 HistGradientBoostingClassifier 中用于 类概率的二次近似,提升预测精度。

63.6.7 代码解读:ExponentialLoss(第 802‑840 行)

class ExponentialLoss(BaseLoss):
    def __init__(self, sample_weight=None):
        super().__init__(closs=CyExponentialLoss(),      # 使用 Cython 实现的指数损失
                         link=HalfLogitLink(),           # 使用半对数几率链接(y_pred = expit(2*raw_prediction))
                         n_classes=2)                    # 二分类任务
        self.interval_y_true = Interval(0, 1, True, True)  # y_true 必须在 [0, 1] 区间

    def constant_to_optimal_zero(self, y_true, sample_weight=None):
        term = -2 * np.sqrt(y_true * (1 - y_true))         # 计算指数损失常数项
        if sample_weight is not None:                      # 若有样本权重
            term *= sample_weight                          # 则乘以样本权重
        return term                                        # 返回常数项

    def predict_proba(self, raw_prediction):
        if raw_prediction.ndim == 2 and raw_prediction.shape[1] == 1:  # 处理 (n_samples, 1) 输入
            raw_prediction = raw_prediction.squeeze(1)                 # 自动降维
        proba = np.empty((raw_prediction.shape[0], 2), dtype=raw_prediction.dtype)  # 创建概率输出数组
        proba[:, 1] = self.link.inverse(raw_prediction)                # p = expit(2*raw) 为类别 1 概率
        proba[:, 0] = 1 - proba[:, 1]                                  # 类别 0 概率为 1 - p
        return proba                                                   # 返回概率数组

要点:与二元对数损失不同,指数损失使用 HalfLogitLinkraw → 0.5·logit(p)),对应 AdaBoost 中的 加权投票constant_to_optimal_zero 计算 -2·√(y·(1‑y)),确保完美预测时损失为零。

63.7 链接函数与区间定义 —— 预测空间的坐标变换与有效域约束

63.7.1 生活类比

链接函数就像 地图的投影变换

  • IdentityLink 等于“原样投影”——直接在平面上绘制(适用于实数目标);

  • LogLink 是“对数投影”——把密度大的区域压缩、密度小的区域放大(适用于指数增长的目标);

  • LogitLink 是“sigmoid 投影”——把整个实数轴压缩到 (0,1) 区间(适用于概率);

  • MultinomialLogit 是“对称投影”——保证多类预测的和为 1(避免冗余参数)。

Interval 则是“地图的有效范围标注”——告诉用户哪些区域是可用的。

63.7.2 源码地图

sklearn/_loss/link.py
├── Interval              # 定义 [low, high] 区间的开闭性
├── BaseLink              # 抽象基类:link() 与 inverse()
├── IdentityLink          # y_pred = raw
├── LogLink               # y_pred = exp(raw), y_pred > 0
├── LogitLink             # y_pred = expit(raw), y_pred ∈ (0, 1)
├── HalfLogitLink         # 0.5·logit
└── MultinomialLogit      # 对称 softmax / log‑geometric‑mean

63.7.3 架构图

                  ┌──────────────┐
                  │   BaseLink   │  ← 抽象接口 (link / inverse)
                  └──────┬───────┘
                         │
   ┌───────────────┬─────┼───────┬──────────────────┐
   ▼               ▼     ▼       ▼                  ▼
IdentityLink  LogLink  LogitLink  HalfLogitLink    MultinomialLogit
interval=ℝ   (0, ∞)   (0, 1)     (0, 1)            (0, 1)
class BaseLink(ABC):
    interval_y_pred = Interval(-np.inf, np.inf, False, False)  # 默认预测区间为全实数

    @abstractmethod
    def link(self, y_pred, out=None): ...                        # 链接函数:y_pred → raw_prediction

    @abstractmethod
    def inverse(self, raw_prediction, out=None): ...             # 逆链接函数:raw_prediction → y_pred

class IdentityLink(BaseLink):
    def link(self, y_pred, out=None):                            # 身份链接:g(x) = x
        if out is not None:                                      # 若提供输出数组
            np.copyto(out, y_pred)                               # 则复制数据
            return out                                           # 返回输出数组
        return y_pred                                            # 否则直接返回输入
    inverse = link                                               # 逆链接等于链接函数自身

class LogLink(BaseLink):
    interval_y_pred = Interval(0, np.inf, False, False)          # 对数链接要求 y_pred > 0
    def link(self, y_pred, out=None):                            # 对数链接:g(x) = log(x)
        return np.log(y_pred, out=out)                           # 使用 numpy.log 计算
    def inverse(self, raw_prediction, out=None):                 # 逆对数链接:h(x) = exp(x)
        return np.exp(raw_prediction, out=out)                   # 使用 numpy.exp 计算

class LogitLink(BaseLink):
    interval_y_pred = Interval(0, 1, False, False)               # 对数几率链接要求 y_pred ∈ (0, 1)
    def link(self, y_pred, out=None):                            # 对数几率链接:g(x) = logit(x)
        return logit(y_pred, out=out)                            # 使用 scipy.special.logit 计算
    def inverse(self, raw_prediction, out=None):                 # 逆对数几率链接:h(x) = expit(x)
        return expit(raw_prediction, out=out)                    # 使用 scipy.special.expit 计算

class HalfLogitLink(BaseLink):
    interval_y_pred = Interval(0, 1, False, False)               # 半对数几率链接要求 y_pred ∈ (0, 1)
    def link(self, y_pred, out=None):                            # 半对数几率链接:g(x) = 0.5 * logit(x)
        out = logit(y_pred, out=out)                             # 先计算 logit
        out *= 0.5                                               # 再乘以 0.5
        return out                                               # 返回结果
    def inverse(self, raw_prediction, out=None):                 # 逆半对数几率链接:h(x) = expit(2 * x)
        return expit(2 * raw_prediction, out=out)                # 先乘以 2 再计算 expit

class MultinomialLogit(BaseLink):
    is_multiclass = True                                         # 标记为多分类链接
    interval_y_pred = Interval(0, 1, False, False)               # 预测概率在 (0, 1) 开区间

    def symmetrize_raw_prediction(self, raw_prediction):         # 对称化原始预测(使行和为零)
        return raw_prediction - np.mean(raw_prediction, axis=1)[:, np.newaxis]

    def link(self, y_pred, out=None):                            # 对数几率链接:g(y) = log(y / gmean(y))
        gm = gmean(y_pred, axis=1)                               # 计算几何均值
        return np.log(y_pred / gm[:, np.newaxis], out=out)       # 返回对数比值

    def inverse(self, raw_prediction, out=None):                 # 逆链接是 softmax
        if out is None:                                          # 若未提供输出数组
            return softmax(raw_prediction, copy=True)            # 则返回复制后的 softmax
        else:                                                    # 若提供了输出数组
            np.copyto(out, raw_prediction)                       # 则复制原始预测
            softmax(out, copy=False)                             # 原地计算 softmax
            return out                                           # 返回输出数组

实现要点:每个子类都明确了 合法的 y_pred 区间interval_y_pred),并在 BaseLoss.__init__ 中复制给相应损失实例,实现 输入合法性验证IdentityLinkinverse 直接等于 link,因为恒等映射的逆仍是恒等。MultinomialLogit.link 使用 几何均值 作为参考类别,使得 对称多项 Logitsoftmax 形成互逆关系。

63.7.5 代码解读:Interval.includes(第 15‑35 行)

@dataclass
class Interval:
    low: float
    high: float
    low_inclusive: bool
    high_inclusive: bool

    def includes(self, x):
        if self.low_inclusive:                                 # 若下界包含等于
            low = np.greater_equal(x, self.low)                # 则使用 >= 比较
        else:                                                  # 若下界不包含等于
            low = np.greater(x, self.low)                      # 则使用 > 比较

        if not np.all(low):                                    # 若任一元素不满足下界条件
            return False                                       # 则直接返回 False

        if self.high_inclusive:                                # 若上界包含等于
            high = np.less_equal(x, self.high)                 # 则使用 <= 比较
        else:                                                  # 若上界不包含等于
            high = np.less(x, self.high)                       # 则使用 < 比较

        return bool(np.all(high))                              # 若所有元素均满足上界条件则返回 True

实现要点BaseLoss.in_y_true_rangein_y_pred_range 直接调用此方法,实现 向量化的合法性检查,在训练前可以快速捕获非法输入(如负数计数、概率超界)。Interval@dataclass 自动生成 __init____repr__,简洁明了。

63.8 Cython 高性能实现 —— 数值计算的底层引擎

63.8.1 代码概览

sklearn/_loss/_loss.pxd       # 类型声明、fusion types、CyLossFunction 基类
sklearn/_loss/_loss.pyx       # 具体实现(如 CyHalfSquaredError、CyAbsoluteError …)

63.8.2 架构图

                         ┌──────────────────────────┐
                         │      CyLossFunction      │
                         │    (cdef 抽象类)         │
                         └────────────┬─────────────┘
                                      │ cdef class (继承)
       ┌─────────────┬───────────┬────┴────┬────────────┬────────────┐
       ▼             ▼           ▼         ▼            ▼            ▼
CyHalfSquaredError CyAbsoluteError CyPinballLoss CyHuberLoss CyHalfPoissonLoss …
   cy_loss/        cy_gradient/    (含 quantile)   (含 delta)   (exp)
   cy_grad_hess    cy_grad_hess

CyHalfMultinomialLoss 是唯一 不继承 CyLossFunction 的 Cython 类,因为它的输入是二维的,且需要一次性返回梯度和概率(gradient_proba),因此采用了自定义签名。

63.8.2.1 示例:CyHalfSquaredError

cdef class CyHalfSquaredError(CyLossFunction):
    cdef double cy_loss(self, double y_true, double raw_prediction) noexcept nogil:
        cdef double diff = y_true - raw_prediction                    # 计算残差
        return 0.5 * diff * diff                                      # 返回 0.5 * (y - raw)^2

    cdef double cy_gradient(self, double y_true, double raw_prediction) noexcept nogil:
        return raw_prediction - y_true                                # 返回梯度:raw - y

    cdef double_pair cy_grad_hess(self, double y_true, double raw_prediction) noexcept nogil:
        cdef double_pair out                                          # 定义返回结构体
        out.val1 = self.cy_gradient(y_true, raw_prediction)           # val1 = 梯度
        out.val2 = 1.0                                                # val2 = 海森(常数 1)
        return out                                                    # 返回结构体

实现要点:所有数值计算均在 C 级别 完成,nogil 释放全局解释器锁,配合 OpenMP 的 n_threads 参数实现 多核并行float_pair 结构一次返回梯度与海森,避免两次遍历输入数据。

63.8.2.2 示例:CyPinballLoss

cdef class CyPinballLoss(CyLossFunction):
    cdef readonly double quantile   # 只读属性,Python 端可读取但不可修改

    cdef double cy_loss(self, double y_true, double raw_prediction) noexcept nogil:
        cdef double diff = raw_prediction - y_true                    # 计算残差
        if diff >= 0:                                                 # 若残差非负
            return diff * self.quantile                               # 返回 quantile * diff
        else:                                                         # 若残差为负
            return -diff * (1 - self.quantile)                        # 返回 (1-quantile) * |diff|

    cdef double cy_gradient(self, double y_true, double raw_prediction) noexcept nogil:
        cdef double diff = raw_prediction - y_true                    # 计算残差
        if diff >= 0:                                                 # 若残差非负
            return self.quantile                                      # 返回梯度:quantile
        else:                                                         # 若残差为负
            return -(1 - self.quantile)                               # 返回梯度:-(1-quantile)

    cdef double_pair cy_grad_hess(self, double y_true, double raw_prediction) noexcept nogil:
        cdef double_pair out                                          # 定义返回结构体
        out.val1 = self.cy_gradient(y_true, raw_prediction)           # val1 = 梯度
        out.val2 = 1.0                                                # val2 = 近似海森(固定为 1)
        return out                                                    # 返回结构体

实现要点readonly 修饰符保证 quantile 在 Python 层只能读取,防止意外修改导致数值不一致。cy_gradientcy_loss 使用分支实现分位数的不同斜率。

63.8.2.3 示例:CyHalfMultinomialLoss(不继承 CyLossFunction

cdef class CyHalfMultinomialLoss():
    cdef void cy_gradient(                                              # 无返回值,直接写入梯度数组
        self,
        const floating_in y_true,                                       # 输入:标签(内存视图)
        const floating_in[::1] raw_prediction,                          # 输入:原始预测(内存视图)
        const floating_in sample_weight,                                # 输入:样本权重(标量)
        floating_out[::1] gradient_out,                                 # 输出:梯度数组(内存视图)
    ) noexcept nogil:
        # 1. 计算 softmax 概率
        # 2. 生成 one‑hot 编码的 y_true
        # 3. gradient = prob - one_hot
        # 4. 如有 sample_weight, 乘以对应权重
        # 该实现直接在 C 级别完成,能够对任意后端(NumPy、CuPy、torch)使用统一的 memoryview 接口。

实现要点floating_in / floating_out 融合类型让输入和输出可以拥有不同的 dtype(例如 float32 输入、float64 输出),从而支持 Array API 的跨后端数值计算。cy_gradient 直接返回 梯度矩阵prob - one_hot),无需额外的海森,因为多分类交叉熵的全局海森在 GBDT 中通常只使用对角线近似。

63.9 Array API 兼容层的数值稳定技巧 —— 跨后端的通用接口

63.9.1 源码地图

sklearn/_loss/loss.py (Array API 部分)
├── _log1pexp            # 数值稳定的 log(1+exp(x))
├── ArrayAPILossMixin    # 为损失提供 Array API 兼容的 __call__
├── HalfBinomialLossArrayAPI      # 兼容 Array API 的二元对数损失
└── HalfMultinomialLossArrayAPI   # 兼容 Array API 的多分类交叉熵

63.9.2 架构图

                       ┌─────────────────────────┐
                       │    ArrayAPILossMixin    │  ← 提供 xp‑aware __call__
                       └─────────────┬───────────┘
                                     │
       ┌────────────────────┬───────┴────────────┬───────────────────────┐
       ▼                    ▼                    ▼                       ▼
HalfBinomialLossArrayAPI …  HalfMultinomialLossArrayAPI   (其他 Array API 损失)
       │
       ├─ loss()           ──→ _compute_loss(xp, …)
       ├─ gradient()       ──→ _compute_gradient(xp, …)
       └─ loss_gradient()  ──→ _compute_loss + _compute_gradient

63.9.3 代码解读:ArrayAPILossMixin(第 897‑920 行)

class ArrayAPILossMixin:
    """Mixin for loss classes that are compatible with the array API."""
    def __call__(self, y_true, raw_prediction, sample_weight=None,
                 n_threads=1, xp=None):
        xp, _ = get_namespace(y_true, raw_prediction, sample_weight, xp=xp)  # 推断 Array API 命名空间
        loss_xp = self.loss(y_true=y_true,                                  # 调用损失计算(如 _compute_loss)
                            raw_prediction=raw_prediction,
                            sample_weight=None)
        return float(_average(loss_xp, weights=sample_weight, xp=xp))       # 使用 Array API 的加权平均并返回 Python 标量

实现要点:Mixin 将 xp(Array API 命名空间)推断出来,然后调用对应的 loss 实现(如 _compute_loss),最后使用 _average 完成加权平均,返回 Python 标量。在 Array API 环境下,n_threads 被忽略,因为底层实现(如 CuPy)会自行调度。

63.9.4 代码解读:数值安全的 _log1pexp(第 850‑895 行)

def _log1pexp(raw_prediction, raw_prediction_exp, xp):
    """Numerically stable version of log(1 + exp(x)) compatible with Array API."""
    constants = ([-37, -2, 18, 33.3] if raw_prediction.dtype == xp.float64   # float64 下的阈值
                 else [-17, -1, 9, 14.6])                                   # float32 下的阈值
    return xp.where(                                                # 使用 where 实现多分支
        raw_prediction <= constants[0],                             # 极端负值:x <= -37(float64)或 <= -17(float32)
        raw_prediction_exp,                                         # 则近似为 exp(x)
        xp.where(                                                   # 否则进入下一层
            raw_prediction <= constants[1],                         # 中等负值:-37 < x <= -2(float64)或 -17 < x <= -1(float32)
            xp.log1p(raw_prediction_exp),                           # 则使用 log1p(exp(x)) 计算 log(1+exp(x))
            xp.where(                                               # 否则进入下一层
                raw_prediction <= constants[2],                     # 中等正值:-2 < x <= 18(float64)或 -1 < x <= 9(float32)
                xp.log(1.0 + raw_prediction_exp),                   # 则直接计算 log(1+exp(x))
                xp.where(                                           # 否则进入下一层
                    raw_prediction <= constants[3],                 # 较大正值:18 < x <= 33.3(float64)或 9 < x <= 14.6(float32)
                    raw_prediction + 1 / raw_prediction_exp,        # 则使用 x + exp(-x) 避免 exp(x) 溢出
                    raw_prediction,                                 # 极端正值:x > 33.3(float64)或 x > 14.6(float32)则近似为 x
                ),
            ),
        ),
    )

实现要点log(1+exp(x))极端负值x << -37)时直接近似为 exp(x),在 极端正值x >> 33.3)时近似为 x,而在 中间区间 使用 log1plog(1+exp(x)) 以保持精度。常数在 float64float32 上不同,以适配各自的数值范围,确保 相对误差 在机器精度以内。

63.9.5 代码解读:HalfBinomialLossArrayAPI._compute_loss_compute_gradient

def _compute_loss(self, xp, y_true, raw_prediction,
                  raw_prediction_exp, sample_weight=None):
    log1pexp = _log1pexp(raw_prediction, raw_prediction_exp, xp)      # 使用数值稳定的 log(1+exp(x))
    loss = log1pexp - y_true * raw_prediction                         # 计算损失:log(1+exp(raw)) - y*raw
    if sample_weight is not None:                                     # 若有样本权重
        loss *= sample_weight                                         # 则乘以样本权重
    return loss                                                       # 返回损失数组

def _compute_gradient(self, xp, y_true, raw_prediction,
                      raw_prediction_exp, sample_weight=None):
    neg_raw_prediction_exp = 1 / raw_prediction_exp                   # 计算 exp(-raw)
    grad = xp.where(                                                # 使用 where 实现多分支
        raw_prediction > (-37 if raw_prediction.dtype == xp.float64 else -17),  # 极端负值条件
        ((1 - y_true) - y_true * neg_raw_prediction_exp)              # 安全区间内的代数等价形式
        / (1 + neg_raw_prediction_exp),                               # 即 (1-y-y*exp(-x))/(1+exp(-x))
        raw_prediction_exp - y_true,                                  # 极端负值时直接使用 exp(x) - y
    )
    if sample_weight is not None:                                     # 若有样本权重
        grad *= sample_weight                                         # 则乘以样本权重
    return grad                                                       # 返回梯度数组

实现要点_compute_loss 使用 _log1pexp 保证数值安全。_compute_gradient极端负值 时直接使用 exp(x) - y,而在 安全区间 内采用 代数等价 的形式 (1‑y‑y·exp(-x))/(1+exp(-x)),从而避免 exp(x) 溢出。

63.9.6 代码解读:HalfMultinomialLossArrayAPI._compute_loss_compute_gradient

def _compute_loss(self, xp, device_, y_true, raw_prediction, sample_weight=None):
    log_sum_exp = _logsumexp(raw_prediction, axis=1, xp=xp)           # 计算数值稳定的 log-sum-exp
    if self.y_true_int is None:                                       # 首次运行时转换标签为整数
        self.y_true_int = xp.asarray(y_true, dtype=xp.int64, device=device_)
    if self.class_indexing_offsets is None:                           # 首次运行时计算索引偏移
        self.class_indexing_offsets = (
            xp.arange(y_true.shape[0], device=device_) * self.n_classes
        )
    true_label_probs = xp.take(_ravel(raw_prediction),                # 通过索引偏移提取真实标签对数概率
                              self.y_true_int + self.class_indexing_offsets)
    loss = log_sum_exp - true_label_probs                             # 计算损失:log-sum-exp - 对数概率
    if sample_weight is not None:                                     # 若有样本权重
        loss *= sample_weight                                         # 则乘以样本权重
    return loss                                                       # 返回损失数组

def _compute_gradient(self, xp, device_, y_true, raw_prediction,
                      sample_weight=None):
    if self.y_true_one_hot is None:                                   # 首次运行时生成 one-hot 编码
        if self.y_true_int is None:                                   # 若尚未转换标签
            self.y_true_int = xp.asarray(y_true, dtype=xp.int64, device=device_)
        self.y_true_one_hot = self.y_true_int[:, None] == xp.arange(  # 生成布尔型 one-hot
            self.n_classes, device=device_
        )
        self.y_true_one_hot = xp.astype(                              # 转换为与 raw_prediction 同 dtype
            self.y_true_one_hot, raw_prediction.dtype, copy=False
        )
    grad = softmax(raw_prediction)                                    # 计算 softmax 概率
    grad -= self.y_true_one_hot                                       # 梯度 = softmax - one-hot
    if sample_weight is not None:                                     # 若有样本权重
        grad *= sample_weight[:, None]                                # 则按样本权重广播后相乘
    return grad                                                       # 返回梯度数组

实现要点_compute_loss 先计算 log‑sum‑exp(数值稳定的 softmax 分母),随后通过 索引偏移class_indexing_offsets)一次性提取对应标签的对数概率,实现 向量化无显式循环 的高效计算。_compute_gradient 通过一次性生成 one‑hot 编码(利用 broadcasting),然后 softmax 减去 one‑hot 得到梯度矩阵。注释中提到的 增量赋值grad[xp.arange(...), y_true] -= 1)将在未来的 Array API 规范中得到支持,从而进一步消除临时大矩阵的分配。

63.10 设计取舍的深度思考

在实现损失函数体系时,团队面临一个关键设计抉择:是通过多重继承将 Cython 损失类直接继承到 Python 层,还是采用组合(composition)将两者解耦。Cython 在 2023 年的 Issue #4350 中明确指出,多重继承会导致方法解析顺序(MRO)异常,进而导致 Cython 方法无法正确定位。为规避这一限制,scikit‑learn 采用了组合模式BaseLoss 持有一个 closs 实例(Cython 损失对象)和一个 link 实例(链接函数对象),所有高层 API(如 lossgradient 等)均委托closs。这种设计的最大优势在于:

  • 可维护性提升:Cython 与 Python 层的职责清晰分离,修改链接函数或 Cython 实现不会相互影响。

  • 跨后端兼容:组合方式天然支持后续的 Array API 实现,只需要在 closs 旁边提供对应的 Array API 包装即可。

  • 明确的错误定位:当出现数值异常时,可以直接定位到 Cython 实现或链接函数,便于调试。

然而,这种组合也带来了一些细微的代价:

  • 每次调用时需要通过属性访问 self.closs,在极端的微基准测试中会产生极少量的 Python 级别开销。

  • 子类必须在构造函数中显式传入 closslink,代码略显冗长。

综合来看,可维护性、跨后端适配以及对 Cython 已知局限的规避的收益远远超过了这点微小的运行时开销。因此,组合模式成为了 scikit‑learn 损失函数体系的最佳实现方案。

63.11 动手练习

练习中的每一项都已改写为完整的段落描述,避免使用简短列表形式。

63.11.1 练习 1:阅读 BaseLoss 基类的完整实现

阅读 sklearn/_loss/loss.pyBaseLoss 类(第 40‑270 行)的完整代码。重点思考以下问题:

  1. 为什么在构造函数中使用组合而不是多重继承来绑定 Cython 损失与链接函数?(提示:Cython Issue #4350)

  2. lossgradientloss_gradientgradient_hessian 四个方法是如何通过 self.closs 将计算委托给 Cython 实现的?

  3. fit_intercept_only 如何利用 self.link 把目标的均值/中位数/分位数映射到链接空间(raw_prediction)?

  4. constant_to_optimal_zero 在损失计算中起什么作用?为什么要把它单独抽离为一个方法?

  5. init_gradient_and_hessian 如何根据 self.constant_hessian 判断是分配完整的海森矩阵还是仅仅分配一个标量?

63.11.2 练习 2:比较回归损失函数的数学特性与实现差异

阅读 sklearn/_loss/loss.py 第 280‑450 行,比较以下回归损失的实现细节:

  1. HalfSquaredErrorAbsoluteErrorPinballLossHuberLossdifferentiableneed_update_leaves_valuesapprox_hessianconstant_hessian 四个属性上的区别。

  2. HalfPoissonLossHalfGammaLossHalfTweedieLoss 如何通过 LogLink 处理正实数约束?它们的 constant_to_optimal_zero 分别补全了哪些常数项?

  3. HalfTweedieLossIdentityHalfTweedieLoss 的区别是什么?在 power 参数变化时,interval_y_pred 如何自适应?

  4. 为什么 AbsoluteErrorPinballLossfit_intercept_only 返回加权中位数/分位数,而 HalfSquaredError 返回加权平均?

63.11.3 练习 3:深入分类损失与链接函数的数学原理

阅读 sklearn/_loss/loss.py 第 452‑670 行以及 sklearn/_loss/link.py 第 92‑115 行,分析以下内容:

  1. HalfBinomialLoss 的损失公式 log(1+exp(raw)) - y*raw 与交叉熵 -y·log(p) - (1-y)·log(1-p) 如何等价?请给出完整的代数推导。

  2. LogitLinkHalfLogitLink 的区别是什么?为什么 ExponentialLoss 使用后者?

  3. MultinomialLogit.link 为什么使用几何均值作为参考类别?这如何导致原始预测的每行求和为零的约束?

  4. HalfMultinomialLoss.gradient_proba 为什么需要同时返回梯度和概率?在 HistGradientBoostingClassifier 中这些信息如何被使用?

63.11.4 练习 4:探究 Cython 底层实现的高性能机制

阅读 sklearn/_loss/_loss.pxd,弄清以下细节:

  1. floating_infloating_out 这两组融合类型的设计意图是什么?它们如何实现输入输出 dtype 的不一致?

  2. CyLossFunction 中的三个纯虚函数 cy_losscy_gradientcy_grad_hess 对应的计算粒度分别是什么?

  3. CyHalfMultinomialLoss 为什么不继承 CyLossFunction?它的 cy_gradient 签名有何特点?

  4. readonlypublic 修饰符在 CyPinballLoss.quantileCyHuberLoss.delta 上的区别是什么?它们对 Python 侧属性访问有什么影响?

63.11.5 练习 5:实现自定义损失函数并验证数值正确性

参考 sklearn/_loss/loss.py 中已有的损失实现,尝试实现一个自定义损失类 MyCustomLoss,要求满足以下条件:

  1. 继承 BaseLoss,组合一个新的 Cython 损失类(可以在 _loss.pxd 中声明,或直接复用已有 Cython 类)。

  2. 为其选择或实现一个合适的链接函数(可以参考 link.py 中的 BaseLink 子类)。

  3. 正确设置 differentiableneed_update_leaves_valuesapprox_hessianconstant_hessian 等属性。

  4. 实现 fit_intercept_onlyconstant_to_optimal_zero,确保截距模型与常数项与数学定义一致。

  5. 编写单元测试,验证:损失非负、梯度在最优点为零、数值微分验证梯度与海森、样本权重正确广播、支持 float32/float64、可 pickle 序列化。请参考 tests/test_loss.py 中的对应测试用例。

63.11.6 练习 6:分析 Array API 兼容层的数值稳定技巧

阅读 sklearn/_loss/loss.py 第 672‑780 行,重点分析:

  1. _log1pexp 为何需要四段分支处理?针对 float64 与 float32,常数 -37/-2/18/33.3-17/-1/9/14.6 分别对应哪种数值近似?

  2. HalfBinomialLossArrayAPI._compute_gradient 中,xp.where 的条件分支如何避免在极端正/负值下的数值溢出?

  3. HalfMultinomialLossArrayAPI._compute_loss 如何利用 class_indexing_offsetsy_true_int 实现 无 one‑hot 编码 的真实标签概率提取?

  4. 注释中提到当前无法使用增量赋值(grad[xp.arange(...), y_true_int] -= 1)。在未来的 Array API 规范中,这一特性将如何改进?

63.12 本章小结

本章我们从宏观架构微观实现,层层剖析了 scikit‑learn 损失函数体系的设计与实现。首先,BaseLoss 通过组合模式统一了 API,并把 Cython 高效计算链接函数 严格解耦;随后,回归与分类损失的数学推导与源码实现被逐一对应,展示了 平方‑绝对‑分位‑Huber‑Poisson‑Gamma‑Tweedie 以及 二元/多元‑对数‑指数 损失的完整路径。接着,链接函数区间约束的设计使得预测空间能够在合法域内自由变换;随后,Cython 层通过 ** fused types、nogil 并行、结构体返回** 实现了 毫秒级 的数值计算。最后,Array API 兼容层通过 数值安全的 _log1pexp向量化的 softmax一次性索引偏移等技巧,实现了 跨后端(NumPy、CuPy、PyTorch) 的一致行为。

接下来,您可以继续阅读 第64章 参数调优与模型选择,学习交叉验证、网格搜索、随机搜索以及 Successive Halving 等先进策略,帮助您在庞大的模型与超参数空间中快速定位最佳配置。祝学习愉快!

63.13 架构与数据流图

graph TD A[loss] --> B[loss] B --> C[loss]
sequenceDiagram participant U as 调用者 participant E as loss participant C as loss U->>E: 调用入口 E->>C: 传递参数 C-->>U: 返回结果
graph LR I[输入] --> P[参数校验] P --> T[核心处理] T --> O[输出]
graph TD L1[用户 API 层] --> L2[算法/服务层] L2 --> L3[数据结构层] L3 --> L4[运行时与依赖层]

上述图分别展示模块依赖、调用时序、数据流和架构分层。

63.14 设计取舍的深度思考

为什么采用当前方案而不是更复杂的替代方案? 本章实现优先保证与既有 API 的一致性、可维护性与运行效率。这意味着在少数极端场景下,调用者需要自行在灵活性、内存与速度之间做取舍,换取默认路径的清晰与稳定。

第 64 章 —— sklearn.externals._arff 源码解析

64.1 学习目标

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

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

本节旨在帮助读者:

  1. 了解 ARFF(Attribute‑Relation File Format)文件的结构及其在机器学习实验中的作用。

  2. 掌握 sklearn.externals._arff 模块提供的 读取写入 接口,包括不同的数据结构(密集、稀疏、生成器)对应的返回类型。

  3. 能够根据实际需求选择合适的矩阵表示方式,并理解模块内部的 编码/解码 逻辑与异常处理机制。

  4. 认识模块关键类(ArffDecoder, ArffEncoder, EncodedNominalConversor, NominalConversor, Data, COOData, LODData)的职责与实现细节。

温馨提示:ARFF 常用于 Weka、ML‑lib 等工具之间的数据交换,熟悉其解析流程有助于在跨平台实验中避免格式错误。


64.2 模块整体架构

flowchart TD A[用户调用 load / loads] -->|字符串或文件| B[ArffDecoder] B --> C{解析阶段} C -->|头部: description| D[_decode_comment] C -->|头部: relation| E[_decode_relation] C -->|头部: attribute| F[_decode_attribute] C -->|数据: @DATA| G[_parse_values] G --> H{数据结构} H -->|DENSE| I[Data.decode_rows] H -->|COO| J[COOData.decode_rows] H -->|LOD| K[LODData.decode_rows] H -->|生成器| L[DenseGeneratorData / LODGeneratorData] I --> M[返回 Python 对象] J --> M K --> M L --> M M --> N[用户获得 dict] N --> O[调用 dump / dumps] O --> P[ArffEncoder] P --> Q{编码阶段} Q -->|属性编码| R[EncodedNominalConversor / NominalConversor] Q -->|数据编码| S[Data.encode_data / COOData.encode_data / LODData.encode_data] S --> T[生成 ARFF 文本]

说明:图中虚线表示可选路径,例如稀疏矩阵的 COOLOD 表示,或使用 生成器 逐行写入,以降低内存占用。


64.3 关键常量与正则表达式

| 常量 | 含义 |

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

| _SIMPLE_TYPES | 支持的基本属性类型:NUMERIC, REAL, INTEGER, STRING |

| _TK_DESCRIPTION, _TK_COMMENT | 用 % 标记的说明与注释行 |

| _TK_RELATION, _TK_ATTRIBUTE, _TK_DATA | ARFF 关键字,分别对应 @RELATION, @ATTRIBUTE, @DATA |

| _RE_RELATION, _RE_ATTRIBUTE | 验证 relation 与 attribute 语法的正则表达式 |

| _RE_DENSE_VALUES, _RE_SPARSE_KEY_VALUES | 通过 _build_re_values() 生成,用于解析密集与稀疏数据行 |

64.3.1 _build_re_values() 细节

  1. quoted_re:匹配双引号包围的值,支持转义(\"\\)以及嵌套单引号。

  2. value_re:允许三种形式的值——双引号、单引号或不含特殊字符的裸字符串。

  3. dense 正则:捕获逗号分隔的值,包括空值(?)与行尾。

  4. sparse 正则:匹配 {index value} 形式的稀疏数据,确保索引是整数且值遵循 value_re

这些正则表达式构成了解码过程的 词法 层,能够在不完整或异常的行上提供明确的错误定位。


64.4 数据结构常量

DENSE = 0       # 完全密集矩阵
COO   = 1       # (row, col, value) 三元组的稀疏坐标格式
LOD   = 2       # 每行一个 dict 的稀疏列表格式
DENSE_GEN = 3   # 密集数据的生成器(逐行产出)
LOD_GEN   = 4   # 稀疏字典列表的生成器
  • 密集 适合特征全部非缺失(如图像、表格)。

  • COOLOD 适用于大部分为零的稀疏特征(如文本词袋)。

  • 生成器 在处理 GB 级别 ARFF 文件时可以显著降低峰值内存。


64.5 异常层次结构

所有自定义异常均继承自 ArffException,其 __str__ 方法会自动引用出现错误的 行号self.line),从而帮助用户快速定位问题。

| 异常类 | 触发条件 |

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

| BadRelationFormat | @RELATION 行语法不符合 _RE_RELATION |

| BadAttributeFormat | @ATTRIBUTE 行不匹配 _RE_ATTRIBUTE |

| BadAttributeType | 属性类型既不是基本类型也不是合法的 nominal 列表 |

| BadAttributeName | 两个属性使用了相同的名称 |

| BadNominalValue | 数据行出现了未在 nominal 列表中声明的取值 |

| BadNominalFormatting | nominal 值包含空格却未加引号 |

| BadNumericalValue | 数值属性的字符串无法转换为 float |

| BadStringValue | 字符串属性出现未加引号的空格 |

| BadLayout | 整体结构不符合 ARFF 规范(例如缺少 @DATA) |

| BadObject | 用户提供的 Python 对象不符合编码要求 |

设计取舍:采用专属异常而非统一 ValueError,使得错误信息更具可读性,尤其在自动化数据流水线中可以通过异常类快速分流处理。


64.6 编码/解码核心类

64.6.1 EncodedNominalConversorNominalConversor

class EncodedNominalConversor:
    def __init__(self, values):
        # 将 nominal 列表映射为整数索引,0 预留为缺失值
        self.values = {v: i for i, v in enumerate(values)}
        self.values[0] = 0

    def __call__(self, value):
        try:
            return self.values[value]
        except KeyError:
            raise BadNominalValue(value)

class NominalConversor:
    def __init__(self, values):
        # 使用集合保持原始字符串,0 表示稀疏矩阵中的默认值
        self.values = set(values)
        self.zero_value = values[0]

    def __call__(self, value):
        if value not in self.values:
            if value == 0:               # 稀疏解码的隐式值
                return self.zero_value
            raise BadNominalValue(value)
        return str(value)
  • EncodedNominalConversor:在读取 ARFF 时将 nominal 值直接 编码为整数,常用于机器学习模型需要数值标签的场景(encode_nominal=True)。

  • NominalConversor:保持原始字符串,仅在稀疏解码时提供默认值。

这两者的差别体现在 ArffDecoder._decode 中对每个属性的 conversor 选择逻辑。

64.6.2 数据解码助手

| 类 | 负责的矩阵类型 | 关键方法 |

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

| Data (DenseGeneratorData + _DataListMixin) | 密集 | decode_rows, _decode_values, encode_data |

| COOData | 稀疏 COO | decode_rows, encode_data |

| LODData (LODGeneratorData + _DataListMixin) | 稀疏 LOD | decode_rows, encode_data |

这些类均实现了统一的 decode_rows(stream, conversors) 接口,接受行迭代器 stream 与属性对应的 conversors,返回对应的 Python 数据结构。

64.6.2.1 示例:密集解码过程

def decode_rows(self, stream, conversors):
    for row in stream:
        values = _parse_values(row)          # 解析为 list / dict
        if isinstance(values, dict):       # 稀疏行 → 转为完整向量
            values = [values[i] if i in values else 0
                      for i in range(len(conversors))]
        else:
            if len(values) != len(conversors):
                raise BadDataFormat(row)
        yield self._decode_values(values, conversors)

64.6.3 ArffDecoder

  • 构造:初始化 _conversors 与行号计数器。

  • 主要流程decode()_decode()

    • 逐行扫描,依据当前 状态(DESCRIPTION → RELATION → ATTRIBUTE → DATA)分派内部解码函数。

    • 在遇到 @DATA 前收集所有属性并为每个属性创建对应的 conversorEncodedNominalConversorNominalConversor,或基本类型的 lambda)。

    • 通过 _get_data_object_for_decoding(matrix_type) 取得对应的数据解码器(Data, COOData, LODData 等),随后把 stream(除去空行与注释)喂给解码器得到 obj['data']

异常包装:外层捕获 ArffException,在抛出前填充错误的 行号,提升调试体验。

64.6.4 ArffEncoder

  • 主要方法

    • _encode_comment_encode_relation_encode_attribute 负责把元信息转成 ARFF 文本。

    • encode()iter_encode() 负责把 数据 部分转换为对应的 ARFF 行。

  • 属性编码:若属性类型是 nominal(列表/元组),会调用 encode_string 对每个取值进行转义,以确保空格、逗号等特殊字符被安全包装。

  • 数据编码:根据对象实际类型自动选择 Data, COODataLODData;稀疏矩阵会被转化为 {index value, ...} 形式。


64.7 关键函数逐行解读

| 函数 | 作用 | 关键实现细节 |

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

| _parse_values(s) | 将单行 ARFF 数据拆解为 Python 列表或稀疏字典 | - 首先检查是否包含需要正则处理的特殊字符
- 对于稠密行使用 _RE_DENSE_VALUES.findall 捕获 valueerror
- 若发现稀疏行(_RE_SPARSE_LINE)则使用 _RE_SPARSE_KEY_VALUES 生成 {index: value} |

| encode_string(s) | 对可能包含特殊字符的字符串进行转义与单引号包装 | - 若包含空格、逗号、百分号等,使用 _RE_ESCAPE_CHARS 替换为 \ 转义序列 |

| _unquote(v) | 去除外层引号并反转转义序列 | - 通过 _escape_sub_callback\t, \n 等恢复为真实字符
- ? 与空字符串统一映射为 None |

| _get_data_object_for_decoding(matrix_type) | 根据返回类型返回对应的解码器实例 | 简单的 if/elif 分支,抛出 ValueError 表示不支持的矩阵类型 |

| load(fp, ...) / loads(s, ...) | 对外的简洁 API,内部实例化 ArffDecoder 并调用 decode | 兼容文件对象与字符串两种输入方式 |


64.8 设计取舍分析

| 维度 | 方案 | 优点 | 缺点 | 适用场景 |

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

| 数据表示 | 密集 list | 读取速度快,易于直接切片 | 大量缺失值会浪费内存 | 小至中等规模、稠密特征 |

| | 稀疏 COO (data, rows, cols) | 与 SciPy 稀疏矩阵兼容,内存占用线性 | 需要额外的 row/col 索引数组 | 大规模稀疏特征(文本、One‑Hot) |

| | 稀疏 LOD list[dict] | 逐行访问且易于序列化为 ARFF 稀疏格式 | 随机访问速度慢,字典开销略大 | 行数不多但列数极大(基因表达) |

| | 生成器 DENSE_GEN / LOD_GEN | 零内存 读取,大文件流式处理 | 只能一次遍历,无法随机回溯 | 超大 ARFF 文件(>10GB) |

| 异常处理 | 细粒度自定义异常 | 直观定位错误行,便于单元测试 | 类层级略增代码量 | 开发库或教学演示时的友好提示 |

| 正则实现 | 单一复合正则 | 仅一次扫描即可完成 token 化,性能好 | 正则难以阅读,维护成本高 | 对性能敏感的生产环境 |

| 兼容性 | 支持 Python 2.7+ 与 PyPy | 适配老旧环境 | 需要保留老式字符串处理逻辑 | 需要在多平台部署的库 |


64.9 完整类与函数概览(含简要注释)

# 第 64 章 —— ---------- 公共常量 ----------
_SIMPLE_TYPES = ['NUMERIC', 'REAL', 'INTEGER', 'STRING']
_TK_DESCRIPTION = '%'
_TK_COMMENT     = '%'
_TK_RELATION    = '@RELATION'
_TK_ATTRIBUTE   = '@ATTRIBUTE'
_TK_DATA        = '@DATA'

# 第 64 章 —— ---------- 正则构建 ----------
_RE_RELATION   = re.compile(r'^([^\{\}%,\s]*|\".*\"|\'.*\')$', re.UNICODE)
_RE_ATTRIBUTE  = re.compile(r'^(\".*\"|\'.*\'|[^\{\}%,\s]*)\s+(.+)$', re.UNICODE)
_RE_QUOTE_CHARS = re.compile(r'["\'\\\s%,\000-\031]', re.UNICODE)
_RE_ESCAPE_CHARS = re.compile(r'(?=["\'\\%])|[\n\r\t\000-\031]')
_RE_SPARSE_LINE = re.compile(r'^\s*\{.*\}\s*$', re.UNICODE)
_RE_NONTRIVIAL_DATA = re.compile('["\'{}\\s]', re.UNICODE)

# 第 64 章 —— ---------- 辅助正则 ----------
_RE_DENSE_VALUES, _RE_SPARSE_KEY_VALUES = _build_re_values()

# 第 64 章 —— ---------- 转义映射 ----------
_ESCAPE_SUB_MAP = {...}
_UNESCAPE_SUB_MAP = {...}
# 第 64 章 —— ---------- 转义/反转义 ----------
def _escape_sub_callback(match): ...
def _unquote(v): ...

# 第 64 章 —— ---------- 解析值 ----------
def _parse_values(s): ...

# 第 64 章 —— ---------- 数据结构常量 ----------
DENSE = 0; COO = 1; LOD = 2; DENSE_GEN = 3; LOD_GEN = 4
_SUPPORTED_DATA_STRUCTURES = [DENSE, COO, LOD, DENSE_GEN, LOD_GEN]

# 第 64 章 —— ---------- 自定义异常 ----------
class ArffException(Exception): ...
class BadRelationFormat(ArffException): ...
class BadAttributeFormat(ArffException): ...
class BadDataFormat(ArffException): ...
class BadAttributeType(ArffException): ...
class BadAttributeName(ArffException): ...
class BadNominalValue(ArffException): ...
class BadNominalFormatting(ArffException): ...
class BadNumericalValue(ArffException): ...
class BadStringValue(ArffException): ...
class BadLayout(ArffException): ...
class BadObject(ArffException): ...

# 第 64 章 —— ---------- 编码/解码类 ----------
def encode_string(s): ...

class EncodedNominalConversor:
    ...

class NominalConversor:
    ...

class DenseGeneratorData:
    ...

class _DataListMixin:
    ...

class Data(_DataListMixin, DenseGeneratorData):
    ...

class COOData:
    ...

class LODGeneratorData:
    ...

class LODData(_DataListMixin, LODGeneratorData):
    ...

def _get_data_object_for_decoding(matrix_type): ...

def _get_data_object_for_encoding(matrix): ...

# 第 64 章 —— ---------- 高层接口 ----------
class ArffDecoder:
    def __init__(self): ...
    def _decode_comment(self, s): ...
    def _decode_relation(self, s): ...
    def _decode_attribute(self, s): ...
    def _decode(self, s, encode_nominal=False, matrix_type=DENSE): ...
    def decode(self, s, encode_nominal=False, return_type=DENSE): ...

class ArffEncoder:
    def _encode_comment(self, s=''): ...
    def _encode_relation(self, name): ...
    def _encode_attribute(self, name, type_): ...
    def encode(self, obj): ...
    def iter_encode(self, obj): ...

# 第 64 章 —— ---------- 基础 API ----------
def load(fp, encode_nominal=False, return_type=DENSE): ...
def loads(s, encode_nominal=False, return_type=DENSE): ...
def dump(obj, fp): ...
def dumps(obj): ...

64.10 使用示例

from sklearn.externals import _arff

# 第 64 章 —— 1️⃣ 读取 ARFF 文件为字典(密集返回)
data_dict = _arff.load(open('iris.arff'), return_type=_arff.DENSE)

# 第 64 章 —— 2️⃣ 读取为稀疏 COO(适用于大特征空间)
sparse_tuple = _arff.load(open('text_data.arff'), return_type=_arff.COO)
values, rows, cols = sparse_tuple  # 与 scipy.sparse.coo_matrix 兼容

# 第 64 章 —— 3️⃣ 写入 ARFF(带描述、属性、稀疏数据)
obj = {
    'description': 'Demo dataset',
    'relation': 'demo',
    'attributes': [('x1', 'REAL'), ('x2', ['A', 'B', 'C'])],
    'data': [{0: 0.5, 1: 'B'}, {0: 1.2, 1: 'A'}]  # LOD 格式
}
_arff.dump(obj, open('out.arff', 'w'))

64.11 小结

sklearn.externals._arff 通过 严谨的正则解析分层异常体系多种矩阵表示,为机器学习工作流提供了可靠的 ARFF 读写能力。掌握其内部类(尤其是 EncodedNominalConversorNominalConversor)的职责,有助于在实际项目中:

  • 选择合适的返回结构以平衡 内存计算速度

  • 在出现数据不一致时快速定位并修正错误。

  • 将自定义 Python 数据结构无缝导出为符合 Weka/ARFF 规范的文件。

后续阅读:了解 sklearn.externals._numpydocsklearn.externals._packaging.version 也是提升源码阅读能力的好路径,它们分别演示了 文档抽取PEP‑440 版本解析 的实现技巧。

64.12 生活类比

想象 scikit-learn 的 externals 模块是一个自给自足的“瑞士军刀工具箱”ARFF 解析器 = 通用翻译官:能读懂 Weka 的“方言”(ARFF),转成 Python 通用的字典/数组,还能把结果写回 ARFF 方言 NumPy 文档解析器 = 结构化提取器:把松散的文档字符串“拆解”成参数、返回值、See Also 等标准化零件,方便自动生成 API 文档 PEP 440 版本解析器 = 版本法庭裁判:把乱七八糟的版本号(1.0a1、1.0.post1、legacy 版本)统一翻译成可比较的“排序键”,判定新旧优劣 SciPy 稀疏图拉普拉斯 = 图论计算器:输入邻接矩阵,输出拉普拉斯矩阵(或矩阵向量乘积函数),支持归一化、对称化,且不必显式构造大矩阵省内存 vendoring 机制 = 应急物资仓库:关键第三方库(liac-arff、numpydoc、packaging、scipy.csgraph)的精简副本随 scikit-learn 打包分发,用户无需额外 pip install 即可使用核心功能,避免依赖地狱

64.13 模块地图/架构图

sklearn/externals/__init__.py
├── (空文件,标记包存在)
sklearn/externals/conftest.py
├── pytest_ignore_collect()  # 禁止收集外部依赖中的测试
sklearn/externals/_array_api_compat_vendor.py
├── 从 .array_api_compat 导入所有符号  # 协同 vendor array_api_compat
sklearn/externals/_arff.py
├── 常量与正则定义
│   ├── _TK_DESCRIPTION/COMMENT/RELATION/ATTRIBUTE/DATA
│   ├── _RE_RELATION/_ATTRIBUTE/_QUOTE_CHARS/_ESCAPE_CHARS/_SPARSE_LINE/_NONTRIVIAL_DATA
│   ├── _RE_DENSE_VALUES, _RE_SPARSE_KEY_VALUES  # 由 _build_re_values 构建
├── 异常体系
│   ├── ArffException (基类)
│   ├── BadRelationFormat / BadAttributeFormat / BadDataFormat
│   ├── BadAttributeType / BadAttributeName / BadNominalValue
│   ├── BadNominalFormatting / BadNumericalValue / BadStringValue
│   ├── BadLayout / BadObject
├── 核心工具函数
│   ├── _unescape_sub_callback / encode_string
│   ├── _unquote / _parse_values
│   ├── _build_re_values / _escape_sub_callback
├── 数据结构适配器 (解码侧)
│   ├── DenseGeneratorData.decode_rows() / _decode_values() / encode_data()
│   ├── Data (继承 _DataListMixin + DenseGeneratorData)
│   ├── COOData.decode_rows() / encode_data()
│   ├── LODGeneratorData.decode_rows() / encode_data()
│   ├── LODData (继承 _DataListMixin + LODGeneratorData)
│   ├── _get_data_object_for_decoding() / _get_data_object_for_encoding()
├── ArffDecoder
│   ├── __init__()
│   ├── _decode_comment() / _decode_relation() / _decode_attribute()
│   ├── _decode()  # 核心状态机解析
│   ├── decode()  # 公共入口,支持 encode_nominal/return_type
├── ArffEncoder
│   ├── _encode_comment() / _encode_relation() / _encode_attribute()
│   ├── encode() / iter_encode()
├── 基础接口
│   ├── load() / loads() / dump() / dumps()
sklearn/externals/_numpydoc/docscrape.py
├── Reader  # 基于行的字符串读取器
│   ├── read() / seek_next_non_empty_line() / eof()
│   ├── read_to_condition() / read_to_next_empty_line() / read_to_next_unindented_line() / peek()
│   ├── __init__() / reset() / __getitem__() / is_empty()
├── NumpyDocString (Mapping)
│   ├── sections  # 标准节定义
│   ├── __init__() -> _parse()
│   ├── _is_at_section() / _strip() / _read_to_next_section() / _read_sections()
│   ├── _parse_param_list()  # 解析 Parameters/Returns 等键值列表
│   ├── _parse_see_also()  # 复杂正则解析交叉引用
│   ├── _parse_index() / _parse_summary() / _parse()
│   ├── _str_header() / _str_indent() / _str_signature() / _str_summary()
│   ├── _str_param_list() / _str_section() / _str_see_also() / _str_index()
│   ├── __str__() / __getitem__() / __setitem__() / __iter__() / __len__()
│   ├── _error_location()
├── FunctionDoc / ObjDoc / ClassDoc  # 面向函数/对象/类的专用解析器
│   ├── FunctionDoc.__init__() / get_func() / __str__()
│   ├── ObjDoc.__init__()
│   ├── ClassDoc.__init__() / methods / properties / _is_show_member() / _should_skip_member()
├── 工具函数
│   ├── get_doc_object()  # 工厂函数
│   ├── dedent_lines() / strip_blank_lines()
├── ParseError  # 解析异常
sklearn/externals/_packaging/version.py
├── 常量与正则
│   ├── VERSION_PATTERN  # PEP 440 完整正则
│   ├── _legacy_version_component_re / _legacy_version_replacement_map
├── 核心数据结构
│   ├── _Version (namedtuple: epoch, release, dev, pre, post, local)
├── InvalidVersion / _BaseVersion (比较运算符实现)
│   ├── __hash__() / __lt__() / __le__() / __eq__() / __ne__() / __gt__() / __ge__()
├── LegacyVersion
│   ├── __init__() -> _legacy_cmpkey()
│   ├── __str__() / __repr__()
│   ├── epoch/release/pre/post/dev/local/public/base_version/is_*release 属性
│   ├── _parse_version_parts() 生成器
├── Version
│   ├── __init__() -> 正则匹配 -> _Version 实例化 -> _cmpkey()
│   ├── __repr__() / __str__()
│   ├── epoch/release/pre/post/dev/local/public/base_version/is_*release 属性
│   ├── major/minor/micro 属性
├── 解析辅助函数
│   ├── _parse_letter_version()  # 归一化 a/b/rc/post/dev
│   ├── _parse_local_version()  # 分割字母数字段
│   ├── _cmpkey()  # 生成排序键,处理 Infinity/NegativeInfinity
│   ├── _legacy_cmpkey()
├── parse()  # 入口函数,优先尝试 Version 回退 LegacyVersion
sklearn/externals/_packaging/_structures.py
├── InfinityType / NegativeInfinityType  # 单例用于比较边界
│   ├── __repr__() / __hash__() / __lt__() / __le__() / __eq__() / __ne__() / __gt__() / __ge__() / __neg__()
sklearn/externals/_packaging/__init__.py
├── (空文件)
sklearn/externals/_scipy/sparse/csgraph/_laplacian.py
├── laplacian()  # 主入口,参数校验与分发
│   ├── 根据 issparse 与 form 选择 _laplacian_sparse/dense/_flo
├── 密集矩阵实现
│   ├── _laplacian_dense() -> 返回 (ndarray, diag)
│   ├── _laplacian_dense_flo() -> 返回 (callable/LinearOperator, diag)
│   ├── _laplace() / _laplace_normed() / _laplace_sym() / _laplace_normed_sym()
│   ├── _setdiag_dense() / _linearoperator()
├── 稀疏矩阵实现
│   ├── _laplacian_sparse() -> 返回 (spmatrix, diag)
│   ├── _laplacian_sparse_flo() -> 返回 (callable/LinearOperator, diag)
sklearn/externals/_scipy/sparse/csgraph/__init__.py
├── from ._laplacian import laplacian
sklearn/externals/_scipy/sparse/__init__.py
├── (空文件)
sklearn/externals/_scipy/__init__.py
├── (空文件)

以上地图列出本章源码模块及其职责,后文将按数据流逐一解析。

64.14 动手练习

64.14.1 ARFF 解析器的稀疏数据支持

阅读 sklearn/externals/_arff.pyCOOData.decode_rowsLODGeneratorData.decode_rows 方法。

回答问题:

  • COO 格式解码时如何构造 (data, rows, cols) 三元组?为何要求行索引有序?

  • LOD (List of Dicts) 格式解码时,隐式零值如何处理?对比 EncodedNominalConversor 与 NominalConversor 对稀疏隐式零值的不同语义。

  • _get_data_object_for_encoding 如何根据输入对象类型自动选择编码器?若传入 CSC 矩阵会发生什么?

64.14.2 NumPy 文档字符串解析的边界情况

阅读 sklearn/externals/_numpydoc/docscrape.pyNumpyDocString._parse_see_also_parse_param_list

回答问题:

  • _line_rgx 正则如何同时匹配 func1, func2: description:meth:func1: description 两种语法?捕获组 allfuncsmorefuncstrailingdesc 分别起什么作用?

  • single_element_is_type=True 在解析 Returns/Yields 时如何改变 name: typetype 两种写法的识别逻辑?

  • ClassDoc 如何通过 inspect.getmembers 发现公共方法与属性?_should_skip_member 为何要特殊处理 namedtuple 的 _fields

64.14.3 PEP 440 版本比较键的生成细节

阅读 sklearn/externals/_packaging/version.pyVersion.__init___cmpkey 函数。

回答问题:

  • _cmpkey 中为何对 release 元组进行 reversed -> dropwhile(==0) -> reversed 操作?举例说明 1.0.01.0 的比较结果。

  • pre 为 None 且 dev 不为 None 时,为何将 _pre 设为 NegativeInfinity?这如何保证 1.0.dev0 < 1.0a0

  • local 版本段的排序键为何将字符串段包装为 (NegativeInfinity, str)、数字段包装为 (int, '')?这如何实现“字母段 < 数字段”及“前缀匹配时短版本优先”?

  • LegacyVersion 如何通过 _parse_version_parts1.0.dev 转换为可比较的元组?zfill(8)* 前缀的作用是什么?

64.14.4 稀疏图拉普拉斯的矩阵自由计算

阅读 sklearn/externals/_scipy/sparse/csgraph/_laplacian.py_laplacian_sparse_flo_laplace_normed_sym

回答问题:

  • form='lo' 时,_linearoperator 如何封装矩阵向量乘积函数?matvecmatmat 为何指向同一个 mv

  • 归一化对称拉普拉斯 L_sym = D^(-1/2) L D^(-1/2)_laplace_normed_sym 中如何分解为两次对角缩放与一次 _laplace_sym 调用?零度节点的 w=1 替代在数学上等效于什么操作?

  • 对比 _laplacian_sparse_laplacian_sparse_flo:前者显式构造 COO 矩阵修改 data/setdiag,后者仅构造闭包函数。在大规模稀疏图上后者的内存优势体现在哪里?

第 65 章 —— 外部工具与依赖管理 —— scikit-learn 的"自给自足工具箱"

65.1 学习目标

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

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

  • 理解外部依赖管理与打包机制,掌握 scikit-learn 如何实现自给自足的工具箱

  • 深入 ARFF 文件格式解析器,理解与 Weka 生态系统的互操作实现

  • 掌握 NumPy 风格文档字符串解析器,理解自动提取参数、返回值等章节的结构化引擎

  • 熟悉 PEP 440 版本解析与比较工具,支持标准版本号及传统版本号的完整排序功能

  • 了解 SciPy 稀疏图拉普拉斯计算的独立封装,支持多种输出格式与归一化选项

  • 掌握 Array API 兼容层核心架构,理解跨后端数组分发与命名空间识别机制

  • 熟悉通用别名与函数包装,理解统一接口层如何处理参数差异与语义对齐

  • 了解线性代数与 FFT 标准化封装,掌握跨后端数值一致性的保障机制

  • 深入后端专属适配器,理解 NumPy、CuPy、PyTorch、Dask 等后端的定制化实现

  • 掌握类型系统与构建工具,理解跨后端契约与自动化模块构建机制

65.2 生活类比

想象 scikit-learn 的外部工具箱是一个自给自足的瑞士军刀工坊,每一个工具都为远为征任务精心锻造。

工坊首先需要一座原材料仓库——这就是 _array_api_compat_vendor.py 的角色。它将 array-api-compat、packaging、numpydoc 等第三方库的源码完整复制到 sklearn/externals/ 目录下。这样做的好处是远为征队出发前不必再四处采购补给,不会出现运行时依赖冲突,也不会因网络问题而陷入安装困境。版本锁定后,远为征队可以精确掌控每个工具的版本与行为。

仓库旁的桌面上摆着一本通用翻译官——_arff.py。它是连接 scikit-learn 与 Weka 生态的桥梁,能读写 ARFF 文件格式,无论是密集数据、稀疏数据 {0 1.0, 2 3.0} 还是多标签属性,它都能准确翻译成 Python 对象。让两套使用不同"语言"的机器学习工具能顺畅交流数据。

紧挨着翻译官的是一台结构化提取引擎——_numpydoc/docscrape.py。它的任务是把非结构化的 docstring 字符串解析成 Parameters、Returns、See Also 等结构化章节。如同从矿石中提炼金属,这台引擎将散落的文字提炼为机器可读的元数据,支撑后续的自动文档生成与参数链接。

角落里坐着一位严格的版本裁判员——_packaging/version.py。他能解析 1.0a1.post2.dev3 这类复杂版本号,按 PEP 440 规范比较大小,确保 1.0.dev0 < 1.0a0 < 1.0 < 1.0.post0 的排序逻辑无误。传统的非标准版本号则由 LegacyVersion 兜底,避免遗漏任何依赖。

货架上还摆着一套独立算法组件——_scipy/sparse/csgraph/_laplacian.py。这是从 SciPy 完整剥离出的图拉普拉斯计算模块,支持归一化、对称化选项,能输出密集数组、稀疏矩阵或 LinearOperator 多种格式。谱聚类、流形学习等算法可以直接调用它,而不必背负 SciPy 的重量。

整个工坊的核心,是一座多语言同声传译中心——array_api_compat/ 目录。这里驻扎着 scikit-learn 最庞大的工程:核心分发器 array_namespace()语言识别雷达,能自动识别输入数组来自 NumPy、CuPy、PyTorch、JAX、Dask 还是 Sparse。通用别名层标准化翻译模板,统一参数名、默认值、返回类型,消除各后端的"方言差异"。后端适配器专属翻译官,针对每个后端的特殊语法(如 PyTorch 的 dim→axis、Dask 的惰性约束)做定制化修正。类型契约 common/_typing.py共同词典,用 TypedDict/Protocol 定义跨后端的类型接口规范,确保所有人"说法一致"。

就像一支装备精良的远为征队,scikit-learn 自带所有必需工具,不依赖外部补给站,在任何环境下都能独立完成从数据加载、文档生成、版本管理到图计算、跨后端数组计算的全流程任务。

65.3 源码地图

sklearn/externals/
├── __init__.py                      # 包标记:External, bundled dependencies
├── conftest.py                      # pytest 配置:忽略 externals 测试
├── _array_api_compat_vendor.py     # vendor 钩子:共置 array_api_compat
│
├── _arff.py                        # ARFF 格式读写(Weka 生态桥头堡)
│
├── _numpydoc/
│   └── docscrape.py                # NumPy 风格 docstring 结构化解析
│
├── _packaging/
│   ├── __init__.py
│   ├── _structures.py              # Infinity/NegativeInfinity 类型
│   └── version.py                  # PEP 440 版本解析与比较
│
├── _scipy/
│   └── sparse/
│       └── csgraph/
│           ├── __init__.py         # 导出 laplacian
│           └── _laplacian.py       # 稀疏图拉普拉斯(从 SciPy 1.12 复制)
│
└── array_api_compat/               # Array API 兼容层(同声传译中心)
    ├── __init__.py                 # 版本声明与公共 re-export
    ├── _internal.py                # get_xp 装饰器 + clone_module 工具
    │
    ├── common/                     # 跨后端通用逻辑
    │   ├── __init__.py             # 从 _helpers 导出全部符号
    │   ├── _helpers.py             # array_namespace 核心分发器
    │   ├── _aliases.py             # 通用别名层(clip/argsort/sort 等)
    │   ├── _linalg.py              # 线性代数标准化(cholesky/svd/qr 等)
    │   ├── _fft.py                 # FFT 精度保持包装
    │   └── _typing.py              # 跨后端类型契约 TypedDict/Protocol
    │
    ├── numpy/                      # NumPy 后端适配
    │   ├── __init__.py
    │   ├── _aliases.py             # asarray copy 语义/astype 等
    │   ├── _info.py                # __array_namespace_info__
    │   ├── _typing.py
    │   ├── linalg.py               # 线性代数专用包装
    │   └── fft.py              # FFT 包装
    │
    ├── cupy/                       # CuPy 后端适配
    │   ├── __init__.py
    │   ├── _aliases.py             # 设备上下文管理
    │   ├── _info.py
    │   ├── _typing.py
    │   ├── linalg.py
    │   └── fft.py
    │
    ├── torch/                      # PyTorch 后端适配
    │   ├── __init__.py
    │   ├── _aliases.py             # _fix_promotion 核心修正
    │   ├── _info.py
    │   ├── _typing.py
    │   ├── linalg.py               # cross/vecdot/solve/vector_norm 修正
    │   └── fft.py                  # axes→dim 重命名
    │
    └── dask/                       # Dask 后端适配
        ├── __init__.py             # 空文件,标记包存在
        └── array/
            ├── __init__.py
            ├── _aliases.py         # 惰性约束下的 clip/sort/argsort
            ├── _info.py            # __array_namespace_info__
            ├── linalg.py           # qr/svd 受限
            └── fft.py

65.4 外部依赖打包机制

这一节我们来看 scikit-learn 的"原材料仓库"如何避免运行时依赖地狱。

scikit-learn 选择将关键第三方库源码直接复制进仓库。这种做法看似简单,却解决了多个棘手问题:避免运行时依赖冲突、确保版本版本行为稳定、支持离线安装。

conftest.py 的设计尤其精妙——它不是用 --ignore 命令行参数(因为 --ignore 需要路径,且 site-packages 中的 externals 路径会很长且依赖安装环境),而是用 pytest_ignore_collect 钩子让 pytest 直接跳过这个子树:

源码路径:sklearn/externals/conftest.py - pytest_ignore_collect()

# 第 65 章 —— Do not collect any tests in externals. This is more robust than using
# 第 65 章 —— --ignore because --ignore needs a path and it is not convenient to pass in
# 第 65 章 —— the externals path (very long install-dependent path in site-packages) when
# 第 65 章 —— using --pyargs
def pytest_ignore_collect(collection_path, config):
    # 始终跳过 externals 子树下的测试发现
    return True

这段代码定义了一个始终返回 True 的 pytest_ignore_collect 钩子函数。当 pytest 试图收集这个目录下的测试时,函数立即返回 True,pytest 会跳过整个目录的测试发现流程。这种"一刀切"的方式比精确指定 --ignore 路径更鲁棒,因为它对安装路径完全不敏感。

_array_api_compat_vendor.py 的设计更加巧妙。它不是单纯的复制文件,而是预留了一个钩子array_api_extra 在 vendor 时可以覆盖函数:

源码路径:sklearn/externals/_array_api_compat_vendor.py - *

# 第 65 章 —— DO NOT RENAME THIS FILE
# 第 65 章 —— This is a hook for array_api_extra/_lib/_compat.py
# 第 65 章 —— to co-vendor array_api_compat and potentially override its functions.

# 第 65 章 —— 重新导出 vendor 后的 array_api_compat 包
from .array_api_compat import *  # noqa: F403

这一行 from .array_api_compat import * 将 vendor 的 array_api_compat 包重新导出。这里的关键在于文件名是约定——array_api_extra_lib/_compat.py 会优先检查这个文件是否存在,如果存在则通过它来覆盖某些函数。注释 DO NOT RENAME THIS FILE 强调这是一个不可破坏的契约。

65.5 ARFF 解析器核心机制

ARFF 解析器是连接 scikit-learn 与 Weka 生态的桥梁。它支持密集格式 1.0,2.0,?、稀疏格式 {0 1.0, 3 2.5}、多种属性类型(NUMERIC、REAL、INTEGER、STR、NOMINAL)以及缺失值标记 ?

下面是核心的状态机驱动的解析流程:

源码路径:sklearn/externals/_arff.py - ArffDecoder._decode()

def _decode(self, s, encode_nominal=False, matrix_type=DENSE):
    '''Do the job the ``encode``.'''

    # 确保方法幂等
    self._current_line = 0

    # 将字符串转换为行列表
    if isinstance(s, str):
        # 统一换行符并按行分割
        s = s.strip('\r\n '").replace('\r\n', '\n').split('\n')

    # 创建返回对象
    obj: ArffContainerType = {
        'description': '',  # 描述部分
        'relation': '',     # 关系名
        'attributes': [],   # 属性列表
        'data': []          # 数据部分
    }
    attribute_names = {}    # 用于检查属性名重复

    # 根据返回类型创建数据辅助对象
    data = _get_data_object_for_decoding(matrix_type)

    # 状态机:依次处理 DESCRIPTION → RELATION → ATTRIBUTE → DATA
    STATE = _TK_DESCRIPTION  # 初始状态为描述部分
    s = iter(s)
    for row in s:
        self._current_line += 1  # 记录行号,便于错误定位
        row = row.strip(' \r\n')
        if not row: continue     # 跳过空行

        u_row = row.upper()

        # 描述部分:以 % 开头
        if u_row.startswith(_TK_DESCRIPTION) and STATE == _TK_DESCRIPTION:
            obj['description'] += self._decode_comment(row) + '\n'

        # 关系声明:@RELATION
        elif u_row.startswith(_TK_RELATION):
            if STATE != _TK_DESCRIPTION:
                raise BadLayout()
            STATE = _TK_RELATION
            obj['relation'] = self._decode_relation(row)

        # 属性定义:@ATTRIBUTE
        elif u_row.startswith(_TK_ATTRIBUTE):
            if STATE != _TK_RELATION and STATE != _TK_ATTRIBUTE:
                raise BadLayout()
            STATE = _TK_ATTRIBUTE

            attr = self._decode_attribute(row)
            if attr[0] in attribute_names:  # 检查重名
                raise BadAttributeName(attr[0], attribute_names[attr[0]])
            else:
                attribute_names[attr[0]] = self._current_line
            obj['attributes'].append(attr)
            # 根据属性类型选择 conversor(类型转换器)
            if isinstance(attr[1], (list, tuple)):  # nominal 类型
                if encode_nominal:
                    conversor = EncodedNominalConversor(attr[1])  # 整数编码
                else:
                    conversor = NominalConversor(attr[1])       # 字符串保留
            else:
                CONVERSOR_MAP = {'STRING': str,
                                 'INTEGER': lambda x: int(float(x)),
                                 'NUMERIC': float,
                                 'REAL': float}
                conversor = CONVERSOR_MAP[attr[1]]
            self._conversors.append(conversor)

        # 数据段开始:@DATA
        elif u_row.startswith(_TK_DATA):
            if STATE != _TK_ATTRIBUTE:
                raise BadLayout()
            break  # 跳出头部解析循环

        # 注释行:在头部允许,在数据段被跳过
        elif u_row.startswith(_TK_COMMENT):
            pass
    else:
        # 循环正常结束说明没找到 @DATA,布局错误
        raise BadLayout()

这段代码定义了 ARFF 解析的核心状态机。_decode 方法按行扫描 ARFF 文件,通过 STATE 变量追踪解析状态(DESCRIPTION → RELATION → ATTRIBUTE → DATA),任何违反状态转移顺序的情况都会抛出 BadLayout 异常。属性类型通过 conversor 对象延迟绑定——nominal 类型使用 EncodedNominalConversor(整数编码)或 NominalConversor(字符串保留),数值类型直接用 Python 内置类型构造。这种设计将词法分析与类型转换解耦,使得稀疏数据与密集数据可以用不同的数据辅助对象处理。

_parse_values 是处理单行数据的核心,它同时支持密集与稀疏两种格式:

源码路径:sklearn/externals/_arff.py - _parse_values()

def _parse_values(s):
    '''(INTERNAL) Split a line into a list of values'''
    # 快速路径:处理无特殊字符的简单行
    if not _RE_NONTRIVIAL_DATA.search(s):
        # 无引号、花括号等特殊字符,直接用 csv 模块拆分
        return [None if s in ('?', '') else s
                for s in next(csv.reader([s]))]

    # 复杂路径:处理引号、转义、空格等情况
    values, errors = zip(*_RE_DENSE_VALUES.findall(',' + s))
    if not any(errors):
        return [_unquote(v) for v in values]

    # 稀疏格式:{index value, index value}
    if _RE_SPARSE_LINE.match(s):
        try:
            # 解析为 {列索引: 值} 字典
            return {int(k): _unquote(v)
                    for k, v in _RE_SPARSE_KEY_VALUES.findall(s)}
        except ValueError:
            # ARFF 语法错误
            for match in _RE_SPARSE_KEY_VALUES.finditer(s):
                if not match.group(1):
                    raise BadLayout('Error parsing %r' % match.group())
            raise BadLayout('Unknown parsing error')
    else:
        # 密集格式语法错误
        for match in _RE_DENSE_VALUES.finditer(s):
            if match.group(2):
                raise BadLayout('Error parsing %r' % match.group())
        raise BadLayout('Unknown parsing error')

这段代码实现了 ARFF 数据行的解析。s 是单个数据样本字符串(如 "1.0,2.0,?""{0 1.0, 3 2.5}")。首先尝试快速路径——如果没有引号、花括号等特殊字符,直接用 csv 模块拆分。然后用 _RE_DENSE_VALUES 正则处理含引号的复杂情况。如果行以 { 开头,则按稀疏格式解析,返回 {列索引: 值} 字典。缺失值 ?_unquote 中统一转为 None

ARFF 解析器与 Weka 互操作的工程细节体现在以下几方面:严格遵循 ARFF 规范(关键字大小写不敏感)、支持转义字符与引号包裹、处理带空格的属性名(自动加引号)、区分数据缺失值与字符串空值。

65.6 NumPy 文档字符串解析器

numpydoc 是 scikit-learn 文档系统的结构化引擎。它将非结构化的 docstring 解析为 Parameters、Returns、See Also 等章节,为自动文档生成与参数文档链接提供数据支撑。

下面用 Mermaid 流程图展示 NumpyDocString 的解析流程:

flowchart TB Input["docstring 字符串"] --> Dedent["textwrap.dedent 去除缩进"] Dedent --> Split["按行分割"] Split --> Reader["Reader 行读取器"] Reader --> Loop["逐章节循环"] Loop --> DetSec{"_is_at_section<br/>检测下划线?"} DetSec -->|是| Parse["解析章节"] DetSec -->|否| Next["读取下一行"] Parse --> End{"还有内容?"} End -->|是| Loop End -->|否| Result["返回章节字典"] Parse -.->|"Parameters/Attributes"| ParamList["_parse_param_list"] Parse -.->|"Returns/Yields/Raises"| SingleType["_parse_param_list single_element_is_type=True"] Parse -.->|"See Also"| SeeAlso["_parse_see_also 正则匹配"] Parse -.->|".. index::"| Index["_parse_index"] Parse -.->|"Notes/Examples 等"| Raw["直接存入 raw lines"]

65.6.1 NumpyDocString 核心架构

NumpyDocString 类继承自 Mapping,实例本身就是章节名到结构化数据的字典映射。这种设计让用户可以通过 docstring["Parameters"] 这种自然方式访问章节。

源码路径:sklearn/externals/_numpydoc/docscrape.py - NumpyDocString.__init__()

class NumpyDocString(Mapping):
    """Parses a numpydoc string to an abstract representation"""

    # 所有支持的章节及其默认值
    sections = {
        "Signature": "",
        "Summary": [""],
        "Extended Summary": [],
        "Parameters": [],
        "Attributes": [],
        "Methods": [],
        "Returns": [],
        "Yields": [],
        "Receives": [],
        "Other Parameters": [],
        "Raises": [],
        "Warns": [],
        "Warnings": [],
        "See Also": [],
        "Notes": [],
        "References": "",
        "Examples": "",
        "index": {},
    }

    def __init__(self, docstring, config=None):
        orig_docstring = docstring  # 保存原始字符串用于错误定位
        # textwrap.dedent 去除公共缩进,split 按行分割
        docstring = textwrap.dedent(docstring).split("\n")

        self._doc = Reader(docstring)        # 行读取器
        self._parsed_data = copy.deepcopy(self.sections)  # 初始化所有章节

        try:
            self._parse()                    # 执行解析
        except ParseError as e:
            e.docstring = orig_docstring     # 在异常中附加原始字符串
            raise

这段代码定义了 NumpyDocString 类的初始化逻辑。sections 类属性列出所有支持的章节及默认值。__init__textwrap.dedent 去除 docstring 的公共缩进(处理多行内字符串嵌套问题),然后用内部 Reader 类按行管理解析过程。copy.deepcopy(self.sections) 避免类属性被实例共享。解析过程中如果出现错误,会将原始 docstring 附加到异常对象,便于用户定位。

65.6.2 章节检测机制

numpydoc 的章节检测基于下划线(---===)识别:

源码路径:sklearn/externals/_numpydoc/docscrape.py - _is_at_section()

def _is_at_section(self):
    self._doc.seek_next_non_empty_line()  # 跳到下一个非空行

    if self._doc.eof():
        return False

    l1 = self._doc.peek().strip()        # 当前行:章节标题
    if l1.startswith(".. index::"):
        return True

    l2 = self._doc.peek(1).strip()       # 下一行:下划线
    if len(l2) >= 3 and (set(l2) in ({"-"}, {"="})) and len(l2) != len(l1):
        snip = "\n".join(self._doc._str[:2]) + "..."
        self._error_location(
            f"potentially wrong underline length... \n{l1} \n{l2} in \n{snip}",
            error=False,
        )
    # 下划线全部由 - 或 = 组成,且长度匹配,则判定为章节边界
    return l2.startswith("-" * len(l1)) or l2.startswith("=" * len(l1))

这段代码实现章节检测。numpydoc 格式约定:章节标题独占一行,紧跟一行下划线,下划线的字符数必须等于或长于标题长度。检测时先获取当前行(章节标题),再 peek 第二行(下划线),如果下划线全部由 -= 组成,且长度匹配,则判定为章节边界。如果下划线长度不匹配,会发出警告(error=False 表示只警告不报错)。

65.6.3 参数列表解析

Parameters 章节的解析最为复杂,需要处理 name : type 这种格式:

源码路径:sklearn/externals/_numpydoc/docscrape.py - _parse_param_list()

def _parse_param_list(self, content, single_element_is_type=False):
    content = dedent_lines(content)  # 去除缩进
    r = Reader(content)
    params = []
    while not r.eof():
        header = r.read().strip()          # 读取一行参数标题
        if " : " in header:
            # 标准格式:name : type
            arg_name, arg_type = header.split(" : ", maxsplit=1)
        else:
            # 单元素可能是类型(用于 Returns/Yields/Raises 等)
            header = header.removesuffix(" :")
            if single_element_is_type:
                arg_name, arg_type = "", header
            else:
                arg_name, arg_type = header, ""

        desc = r.read_to_next_unindented_line()  # 读取描述(缩进的段落)
        desc = dedent_lines(desc)
        desc = strip_blank_lines(desc)

        params.append(Parameter(arg_name, arg_type, desc))
    return params

这段代码解析参数列表。content 是章节内容行列表。逐行读取参数标题(如 x : array_like),用 " : " 拆分名称与类型。single_element_is_type 参数用于 Returns/Yields/Raises 等章节(这些章节没有参数名,整行就是类型)。然后读取缩进的描述段落直到非缩进行。Parameter 是一个 namedtuple,包含 nametypedesc 三个字段。dedent_linesstrip_blank_lines 清理描述段落的缩进与空行。

65.6.4 See Also 章节解析

See Also 章节的解析则使用正则匹配函数名与角色:

源码路径:sklearn/externals/_numpydoc/docscrape.py - _parse_see_also()

def _parse_see_also(self, content):
    """
    func_name : Descriptive text
        continued text
    another_func_name : Descriptive text
    func_name1, func_name2, :meth:`func_name`, func_name3
    """
    content = dedent_lines(content)

    items = []

    def parse_item_name(text):
        """Match ':role:`name`' or 'name'."""
        m = self._func_rgx.match(text)
        if not m:
            self._error_location(f"Error parsing See Also entry {line!r}")
        role = m.group("role")
        name = m.group("name") if role else m.group("name2")
        return name, role, m.end()

    rest = []
    for line in content:
        if not line.strip():
            continue

        line_match = self._line_rgx.match(line)
        description = None
        if line_match:
            description = line_match.group("desc")
            if line_match.group("trailing") and description:
                self._error_location(
                    "Unexpected comma or period after function list at index %d of "
                    'line "%s"' % (line_match.end("trailing"), line),
                    error=False,
                )
        if not description and line.startswith(" "):
            rest.append(line.strip())
        elif line_match:
            funcs = []
            text = line_match.group("allfuncs")
            while True:
                    if not text.strip():
                        break
                    # 逐个解析函数名(可能带角色)
                    name, role, match_end = parse_item_name(text)
                    funcs.append((name, role))
                    text = text[match_end:].strip()
                    if text and text[0] == ",":
                        text = text[1:].strip()
            rest = list(filter(None, [description]))
            items.append((funcs, rest))
        else:
            self._error_location(f"Error parsing See Also entry {line!r}")
    return items

这段代码解析 See Also 章节。支持多种格式:纯函数名、func_name : description、带角色标记 :meth:\func`、多函数逗号分隔 func1, func2_func_rgx正则匹配:role:`name` 或纯函数名,_line_rgx 匹配整行。parse_item_name` 是内嵌辅助函数,提取函数名与角色。

65.7 PEP 440 版本解析与比较

_packaging/version.py 是 scikit-learn 自带的版本号裁判员,完整复制自 PyPA 的 packaging 库。它能解析 1.0a1.post2.dev3 这类 PEP 440 规范版本,也能兜底处理传统非标准版本号。

下面用 Mermaid 流程图展示 _cmpkey 的核心逻辑:

flowchart TB Input["版本号字符串"] --> Regex["VERSION_PATTERN 命名捕获组"] Regex --> Epoch["epoch"] Regex --> Release["release"] Regex --> Pre["pre"] Regex --> Post["post"] Regex --> Dev["dev"] Regex --> Local["local"] Epoch --> Key["_cmpkey 元组键"] Release --> Strip["去除末尾零"] Strip --> Key Pre --> PreLogic{"pre/post 是否为空?"} PreLogic -->|dev 存在| NegInf["_pre = NegativeInfinity"] PreLogic -->|pre 空| Inf["_pre = Infinity"] PreLogic -->|其他| Pre NegInf --> Key Inf --> Key Post --> PostLogic{"post 是否为空?"} PostLogic -->|否| NegInf2["_post = NegativeInfinity"] PostLogic -->|是| Post NegInf2 --> Key Dev --> DevLogic{"dev 是否为空?"} DevLogic -->|否| Inf2["_dev = Infinity"] DevLogic -->|是| Dev Inf2 --> Key Local --> Key Key --> Cmp["元组比较"]

65.7.1 Version 类的正则解析

PEP 440 版本号由五部分组成:[epoch!]release[.pre[.pre_n]][.post[.post_n]][.dev[.dev_n]][+local]VERSION_PATTERN 正则用命名捕获组一次性提取所有部分:

源码路径:sklearn/externals/_packaging/version.py - VERSION_PATTERN

VERSION_PATTERN = r"""
    v?                                                            # 可选的前缀 v
    (?:
        (?:(?P<epoch>[0-9]+)!)?                                   # epoch 段(数字!)
        (?P<release>[0-9]+(?:\.[0-9]+)*)                          # release 段(必需)
        (?P<pre>                                                  # pre-release 段
            [-_\.]?                                               # 可选分隔符
            (?P<pre_l>(a|b|c|rc|alpha|beta|pre|preview))          # pre 标签
            [-_\.]?                                               # 可选分隔符
            (?P<pre_n>[0-9]+)?                                    # 数量 pre 编号
        )?
        (?P<post>                                                 # post release 段
            (?:-(?P<post_n1>[0-9]+))                              # -N 形式
            |
            (?:
                [-_\.]?
                (?P<post_l>post|rev|r))                          # post/rev/r 标签
                [-_\.]?
                (?P<post_n2>[0-9]+)?
            )
        )?
        (?P<dev>                                                  # dev release 段
            [-_\.]?
            (?P<dev_l>dev)                                        # dev 标签
            [-_\.]?
            (?P<dev_n>[0-9]+)?
        )?
    )
    (?:\+(?P<local>[a-z0-9]+(?:[-_\.][a-z0-9]+)*))?               # local version(+后缀)
"""

这段代码定义了 PEP 440 版本号的正则表达式。(?P<name>...) 是命名捕获组,便于后续按名称提取。epoch 是可选的 epoch 段(用 ! 分隔),release 是必需的 release 段(数字与点),pre 是预发布段(a/b/rc/alpha/beta/pre/preview),post 是后发布段(数字后缀或 post/rev/r 后缀),dev 是开发版段(dev 后缀),local 是本地版本标识(+ 后缀)。可选段用 ? 表示,[-_\.]? 允许分隔符为 -_.

65.7.2 比较键生成与排序逻辑

版本比较的核心是 _cmpkey 函数——它把版本号各部分转换为一个可元组比较的键:

源码路径:sklearn/externals/_packaging/version.py - _cmpkey()

def _cmpkey(
    epoch, release, pre, post, dev, local,
) -> CmpKey:
    # 去除 release 末尾的零(保持简洁)
    _release = tuple(
        reversed(list(itertools.dropwhile(lambda x: x == 0, reversed(release))))
    )

    # 关键技巧:让 1.0.dev0 < 1.0a0 < 1.0 < 1.0.post0
    if pre is None and post is None and dev is not None:
        # 有 dev 无 pre/post → pre 视为负无穷,让 dev 版本"假装有 pre"
        _pre = NegativeInfinity
    elif pre is None:
        # 无 pre → pre 视为正无穷,让无预发布的版本"假装有最大的 pre"
        _pre = Infinity
    else:
        _pre = pre

    if post is None:
        # 无 post → 视为负无穷,让无后发布的版本排序靠前
        _post = NegativeInfinity
    else:
        _post = post

    if dev is None:
        # 无 dev → 视为正无穷,让非开发版本排序靠后
        _dev = Infinity
    else:
        _dev = dev

    if local is None:
        # 无 local 段 → 视为负无穷
        _local = NegativeInfinity
    else:
        # local 段解析:字母段排前,数字段排后
        _local = tuple(
            (i, "") if isinstance(i, int) else (NegativeInfinity, i) for i in local
        )

    # 返回可比较的元组键
    return epoch, _release, _pre, _post, _dev, _local

这段代码生成了版本比较的可哈希键。_release 通过反转+dropwhile+再反转去除末尾零(如 (1, 0, 0)(1,))。_pre_post_dev 的处理使用了 Infinity/NegativeInfinity 哨兵类型实现巧妙的排序:

  • 有 dev 无 pre/post:_pre = NegativeInfinity → 让 dev 版本"假装有 pre",确保 1.0.dev0 < 1.0a0

  • 无 pre:_pre = Infinity → 让无预发布的版本"假装有最大的 pre",确保 1.0a0 < 1.0

  • 无 post:_post = NegativeInfinity → 让无后发布的版本排序靠前

  • 无 dev:_dev = Infinity → 让非开发版本排序靠后

_local 段将数字转为 (int, "")、字母转为 (NegativeInfinity, str),让字母段排在数字段前面。_structures.py 中的 InfinityType/NegativeInfinityType 通过重载所有比较运算符实现"比任何值都大/小"的语义。

源码路径:sklearn/externals/_packaging/_structures.py - InfinityType

class InfinityType:
    def __repr__(self) -> str:
        return "Infinity"

    def __hash__(self) -> int:
        return hash(repr(self))

    def __lt__(self, other: object) -> bool:
        return False       # ∞ 不小于任何东西

    def __le__(self, other: object) -> bool:
        return False

    def __eq__(self, other: object) -> bool:
        return isinstance(other, self.__class__)  # 仅与同类相等

    def __ne__(self, other: object) -> bool:
        return not isinstance(other, self.__class__)

    def __gt__(self, other: object) -> bool:
        return True        # ∞ 大于任何东西

    def __ge__(self, other: object) -> bool:
        return True

    def __neg__(self: object) -> "NegativeInfinityType":
        return NegativeInfinity  # -∞ 是一反向

这段代码定义了 InfinityType 类。它重载了所有 six 个比较运算符:</<= 永远返回 False(∞ 不小于任何东西),>/>= 永远返回 True(∞ 大于任何东西),== 仅在同类间成立。__neg__-Infinity 返回 NegativeInfinity 单例。这种设计让 Infinity 可以作为元组元素参与比较,无需特殊处理。

65.7.3 LegacyVersion 兜底策略

对于不符合 PEP 440 规范的版本号,LegacyVersion 提供了兜底:

源码路径:sklearn/externals/_packaging/version.py - LegacyVersion.__init__()

class LegacyVersion(_BaseVersion):
    def __init__(self, version: str) -> None:
        self._version = str(version)
        self._key = _legacy_cmpkey(self._version)

        warnings.warn(
            "Creating a LegacyVersion has been deprecated and will be "
            "removed in the next major release",
            DeprecationWarning,
        )

    @property
    def epoch(self) -> int:
        return -1

    @property
    def release(self) -> None:
        return None

    @property
    def is_prerelease(self) -> bool:
        return False

LegacyVersion_legacy_cmpkey 生成比较键,硬编码 epoch = -1,让传统版本号永远排在 PEP 440 版本之前。parse 函数会先尝试 Version(version),失败则降级为 LegacyVersion(version),确保任何版本字符串都能被处理。

65.8 SciPy 稀疏图拉普拉斯计算

_laplacian.py 是从 SciPy 1.12 完整剥离的图拉普拉斯计算模块。scikit-learn 选择 vendor 这个文件是为了支持 sparse arrays(SciPy 1.11 仅支持 sparse matrix),且不增加对 SciPy 1.12+ 的硬性依赖。

下面用 Mermaid 流程图展示 laplacian 函数的多输出格式选择:

flowchart TB Input["csgraph 输入"] --> Check{"是否为稀疏?"} Check -->|是| Sparse["_laplacian_sparse 系列"] Check -->|否| Dense["_laplacian_dense 系列"] Form["form 参数"] --> CheckForm{"form 值"} CheckForm -->|"array"| Array["返回 array"] CheckForm -->|"function"| Func["返回 lambda"] CheckForm -->|"lo"| LO["返回 LinearOperator"] Sparse --> Array Sparse --> Func Sparse --> LO Dense --> Array Dense --> Func Dense --> LO

65.8.1 laplacian 函数的多输出格式

laplacian 函数是模块的核心入口,支持密集数组、稀疏矩阵、LinearOperator 三种输出格式:

源码路径:sklearn/externals/_scipy/sparse/csgraph/_laplacian.py - laplacian()

def laplacian(
    csgraph,                # 输入图(密集 ndarray 或稀疏矩阵)
    normed=False,           # 是否启用归一化拉普拉斯
    return_diag=False,      # 是否同时返回度向量
    use_out_degree=False,   # 使用出度还是入度
    *,
    copy=True,              # 是否复制输入
    form="array",           # 输出格式:array/function/lo
    dtype=None,             # 输出 dtype
    symmetrized=False,      # 是否强制对称化
):
    # 输入必须是方阵
    if csgraph.ndim != 2 or csgraph.shape[0] != csgraph.shape[1]:
        raise ValueError("csgraph must be a square matrix or array")

    # 归一化时强制转为 float64(避免整数溢出)
    if normed and (
        np.issubdtype(csgraph.dtype, np.signedinteger)
        or np.issubdtype(csgraph.dtype, np.uint)
    ):
        csgraph = csgraph.astype(np.float64)

    # 根据 form 参数选择实现路径
    if form == "array":
        # 数组形式:选择稀疏或密集实现
        create_lap = _laplacian_sparse if issparse(csgraph) else _laplacian_dense
    else:
        # 算子形式:选择稀疏或密集的 _flo(function/lo)实现
        create_lap = (
            _laplacian_sparse_flo if issparse(csgraph) else _laplacian_dense_flo
        )

    # 选择入度还是出度
    degree_axis = 1 if use_out_degree else 0

    # 调用选定的实现
    lap, d = create_lap(
        csgraph,
        normed=normed,
        axis=degree_axis,
        copy=copy,
        form=form,
        dtype=dtype,
        symmetrized=symmetrized,
    )
    # 根据 return_diag 决定是否同时返回度向量
    if return_diag:
        return lap, d
    return lap

这段代码定义了 laplacian 函数的主入口。csgraph 是输入图(密集 ndarray 或稀疏矩阵),normed=True 启用归一化拉普拉斯,symmetrized=True 强制对称化(处理有向图),form 选择输出格式("array"/"function"/"lo"),return_diag=True 同时返回度向量。函数先做输入校验(必须是 2D 方阵),再根据 form 选择不同的实现路径(密集 vs 稀疏,array vs 算子),最后根据 use_out_degree 决定度计算的方向。

65.8.2 稀疏实现的代数技巧

_laplacian_sparse 用稀疏矩阵的就地修改实现非归一化拉普拉斯:

源码路径:sklearn/externals/_scipy/sparse/csgraph/_laplacian.py - _laplacian_sparse()

def _laplacian_sparse(graph, normed, axis, copy, form, dtype, symmetrized):
    del form  # form 在此路径无用

    if dtype is None:
        # 默认使用输入的 dtype
        dtype = graph.dtype

    needs_copy = False
    # LIL 和 DOK 不支持算术运算,先转 Coo
    if graph.format in ("lil", "dok"):
        m = graph.tocoo()
    else:
        m = graph
        if copy:
            needs_copy = True

    if symmetrized:
        # 对称化:A + A.T.conj()(不除以 2,保留整数 dtype)
        m += m.T.conj()

    # 计算度:对角线元素(自环)被排除
    w = np.asarray(m.sum(axis=axis)).ravel() - m.diagonal()

    if normed:
        # 归一化分支:实现 D^(-1/2) A D^(-1/2)
        m = m.tocoo(copy=needs_copy)
        isolated_node_mask = w == 0
        # 隔离节点(度=0)的缩放系数设为 1(避免除零)
        w = np.where(isolated_node_mask, 1, np.sqrt(w))
        m.data /= w[m.row]       # 行归一化
        m.data /= w[m.col]       # 列归一化
        m.data *= -1             # 边权加
        # 隔离节点度为 1,普通节点度为 0
        m.setdiag(1 - isolated_node_mask)
    else:
        # 非归一化分支:L = D - A
        if m.format == "dia":
            m = m.copy()
        else:
            m = m.tocoo(copy=needs_copy)
        m.data *= -1             # 边权加
        m.setdiag(w)             # 对角线设为度

    return m.astype(dtype, copy=False), w.astype(dtype)

这段代码实现稀疏图的拉普拉斯计算。graph.format 是稀疏矩阵的存储格式(CSR/CSC/COO/LIL/DOK/LIA),LIL 和 DOK 不支持算术运算,先转 COO。如果 symmetrized=True,执行 m += m.T.conj()(不除以 2 保持精度)。w = sum(axis) - diagonal 计算每个节点的度(排除自环)。

归一化分支的巧妙之处:先用 np.where(isolated_node_mask, 1, np.sqrt(w)) 让隔离节点(度=0)的缩放系数为 1(避免除零),然后对边的两端节点分别除以 sqrt(w)(即 m.data /= w[m.row]; m.data /= w[m.col]),等价于 D^(-1/2) A D^(-1/2)。最后取负并设置对角线为 1(孤立节点)或 0(普通节点)。

非归一化分支更直接:m.data *= -1 让所有边权变为负值(因为 L = D - A,A 是邻接矩阵),然后 m.setdiag(w) 把对角线设为度向量。这种就地修改节省内存(needs_copy 控制是否需要先复制)。

65.9 Array API 兼容层核心架构

array_api_compat 是 scikit-learn 最庞大的外部依赖——它让 scikit-learn 的算法能无缝运行在 NumPy、CuPy、PyTorch、JAX、Dask、Sparse 等不同数组后端上。

下面用 Mermaid 流程图描述 Array API 兼容层在调用 array_namespace() 时的整体数据流:

flowchart TB Input["输入数组 xs"] --> Loop["遍历每个输入"] Loop --> Classify{"类型分类"} Classify -->|NumPy| N["NumPy namespace"] Classify -->|CuPy| C["CuPy namespace"] Classify -->|PyTorch| T["PyTorch namespace"] Classify -->|Dask| D["Dask namespace"] Classify -->|JAX| J["JAX namespace"] Classify -->|Python 标量| S["跳过(SCALAR)"] N --> Collect["收集命名空间到集合"] C --> Collect T --> Collect D --> Collect J --> Collect Collect --> Check{"集合是否唯一?"} Check -->|是| Return["返回唯一 namespace"] Check -->|否,0 个| Error1["抛 TypeError:至少一个非标量"] Check -->|否,>1 个| Error2["抛 TypeError:多后端混合"]

65.9.1 包初始化与版本声明

array_api_compat/__init__.py 是整个兼容层的统一入口,负责版本声明与符号再导出:

源码路径:sklearn/externals/array_api_compat/__init__.py

"""
NumPy Array API compatibility library

This is a small wrapper around NumPy, CuPy, JAX, sparse and others that are
compatible with the Array API standard https://data-apis.org/array-api/latest/.
See also NEP 47 https://numpy.org/neps/nep-0047-array-api-standard.html.
...
"""
__version__ = '1.13.0'

from .common import *  # noqa: F401, F403

这个文件非常短——它只做两件事:声明 __version__ = '1.13.0' 以及 from .common import *。后者将 common 子包的所有公开符号(如 array_namespacedeviceto_deviceis_*_array 等)重新导出到 array_api_compat 顶层命名空间。这样用户既可以 from array_api_compat import array_namespace,也可以 from array_api_compat.common import array_namespace

65.9.2 common 子包初始化模式

common/__init__.py 同样简短——它把 _helpers 中所有符号暴露到 common 命名空间:

源码路径:sklearn/externals/array_api_compat/common/__init__.py

from ._helpers import *  # noqa: F403

这种"层层 re-export"的模式让用户可以从任意层级导入符号,同时保持模块内部的私有性(_helpers.py 内部的辅助函数仍可通过下划线前缀标记)。

65.9.3 核心分发器 array_namespace

array_namespace 是所有兼容层的入口,它接收任意多个数组,验证它们来自同一后端,返回对应的命名空间:

源码路径:sklearn/externals/array_api_compat/common/_helpers.py - array_namespace()

def array_namespace(
    *xs: Array | complex | None,         # 可变参数:接受任意数量的数组或标量
    api_version: str | None = None,       # 数组 API 规范版本
    use_compat: bool | None = None,       # 是否使用 compat 包装
) -> Namespace:
    """获取数组 xs 对应的数组 API 兼容命名空间。"""
    namespaces: set[Namespace] = set()    # 用集合收集命名空间
    for x in xs:
        # 通过类型查找命名空间
        xp, info = _cls_to_namespace(cast(Hashable, type(x)), api_version, use_compat)

        # Python 标量(int/float/complex/None)透传
        if info is _ClsToXPInfo.SCALAR:
            continue

        # NumPy 数组可能是 JAX 零梯度数组的伪装
        if (
            info is _ClsToXPInfo.MAYBE_JAX_ZERO_GRADIENT
            and _is_jax_zero_gradient_array(x)
        ):
            xp = _jax_namespace(api_version, use_compat)

        if xp is None:
            # 后备方案:检查对象自己的 __array_namespace__ 方法
            get_ns = getattr(x, "__array_namespace__", None)
            if get_ns is None:
                raise TypeError(f"{type(x).__name__} is not a supported array type")
            if use_compat:
                raise ValueError(
                    "The given array does not have an array-api-compat wrapper"
                )
            xp = get_ns(api_version=api_version)

        namespaces.add(xp)

    try:
        (xp,) = namespaces    # 解包:必须只有 1 个
        return xp
    except ValueError:
        if not namespaces:
            raise TypeError(
                "array_namespace requires at least one non-scalar array input"
            )
        raise TypeError(f"Multiple namespaces for array inputs: {namespaces}")

这段代码定义了 array_namespace 函数。它遍历所有输入数组,通过 _cls_to_namespace 获取每个数组对应的命名空间。Python 标量被跳过(_ClsToXPInfo.SCALAR)。NumPy 数组会被二次检查是否是否是 JAX 的零梯度数组(一种特殊 dtype,用 np.void 伪装),如果是是返回数组命名空间。最后用集合去重,如果只有一种命名空间则返回它,否则抛出 TypeError(混合后端)。

65.9.4 类型到命名空间的映射

_cls_to_namespace 是后端识别的核心映射表:

源码路径:sklearn/externals/array_api_compat/common/_helpers.py - _cls_to_namespace()

@lru_cache(100)    # 缓存最近 100 次调用结果
def _cls_to_namespace(
    cls: type,
    api_version: str | None,
    use_compat: bool | None,
) -> tuple[Namespace | None, _ClsToXPInfo | None]:
    # 验证 use_compat 参数
    if use_compat not in (None, True, False):
        raise ValueError("use_compat must be None, True, or False")
    # None 与 True 共享的归的导入逻辑
    _use_compat = use_compat in (None, True)
    cls_ = cast(Hashable, cls)

    # NumPy 数组或标量(必须在 Python 标量之前检查!)
    if (
        _issubclass_fast(cls_, "numpy", "ndarray")
        or _issubclass_fast(cls_, "numpy", "generic")
    ):
        if use_compat is True:
            _check_api_version(api_version)
            from .. import numpy as xp
        elif use_compat is False:
            import numpy as xp
        else:
            # NumPy 2.0+ 有 __array_namespace__ 但尚未完全兼容
            from .. import numpy as xp
        # 标记需要二次验证是否为 JAX 零梯度数组
        return xp, _ClsToXPInfo.MAYBE_JAX_ZERO_GRADIENT

    # Python 标量(必须在 np.generic 之后,因为 np.float64 是 float 子类)
    if issubclass(cls, int | float | complex | type(None)):
        return None, _ClsToXPInfo.SCALAR

    # CuPy 数组
    if _issubclass_fast(cls_, "cupy", "ndarray"):
        if _use_compat:
            _check_api_version(api_version)
            from .. import cupy as xp
        else:
            import cupy as xp
        return xp, None

    # PyTorch 张量
    if _issubclass_fast(cls_, "torch", "Tensor"):
        if _use_compat:
            _check_api_version(api_version)
            from .. import torch as xp
        else:
            import torch as xp
        return xp, None

    # Dask 数组
    if _issubclass_fast(cls_, "dask.array", "Array"):
        if _use_compat:
            _check_api_version(api_version)
            from ..dask import array as xp
        else:
            import dask.array as xp
        return xp, None

    # JAX 数组(jnp.ndarray 才有 __array_namespace__)
    if _issubclass_fast(cls_, "jax", "Array"):
        return _jax_namespace(api_version, use_compat), None

    return None, None

这段代码实现类型到命名空间的映射。@lru_cache(100) 装饰器缓存最近 100 次调用结果(cls 哈希可能,但实际场景同一类型会被反复查询)。_use_compat = use_compat in (None, True) 是因为 None 与 True 共享相同的导入逻辑。

关键顺序:NumPy 检查在 Python 标量检查之前,否则 np.float64(是 float 的子类)会被错误地识别为 Python 标量。_ClsToXPInfo.MAYBE_JAX_ZERO_GRADIENT 标记告诉 array_namespace 需要二次验证——NumPy 数组可能实际上是 jax.float0 伪装的零梯度数组。_check_api_version 仅支持 2024.12 版本,其他版本(旧或新)会发出警告或报错。

65.9.5 is_*_namespace 系列函数

除了 array_namespace 之外,兼容层还提供了一系列命名空间类型判断函数:

源码路径:sklearn/externals/array_api_compat/common/_helpers.py - is_*_namespace()

@lru_cache(100)
def is_numpy_namespace(xp: Namespace) -> bool:
    """判断 xp 是否为 NumPy 命名空间(包括原生 NumPy 与 compat 包装)。"""
    # 通过模块名匹配:原生 "numpy" 或 compat 包装路径
    return xp.__name__ in {"numpy", _compat_module_name() + ".numpy"}

@lru_cache(100)
def is_cupy_namespace(xp: Namespace) -> bool:
    """判断 xp 是否为 CuPy 命名空间。"""
    return xp.__name__ in {"cupy", _compat_module_name() + ".cupy"}

@lru_cache(100)
def is_torch_namespace(xp: Namespace) -> bool:
    """判断 xp 是否为 PyTorch 命名空间。"""
    return xp.__name__ in {"torch", _compat_module_name() + ".torch"}

def is_listernamespace(xp: Namespace) -> bool:
    """判断 xp 是否为 NDONNX 命名空间(仅原生,无 compat 包装)。"""
    return xp.__name__ == "ndonnx"

@lru_cache(100)
def is_dask_namespace(xp: Namespace) -> bool:
    """判断 xp 是否为 Dask 命名空间。"""
    return xp.__name__ in {"dask.array", _compat_module_name() + ".dask.array"}

def is_jax_namespace(xp: Namespace) -> bool:
    """判断 xp 是否为 JAX 命名空间。"""
    # JAX 早期有 jax.experimental.array_api,新版本直接用 jax.numpy
    return xp.__name__ in {"jax.numpy", "jax.experimental.array_api"}

def is_pydata_sparse_namespace(xp: Namespace) -> bool:
    """判断 xp 是否为 pydata/sparse 命名空间。"""
    return xp.__name__ == "sparse"

def is_array_api_strict_namespace(xp: Namespace) -> bool:
    """判断 xp 是否为 array-api-strict 命名空间(用于严格测试)。"""
    return xp.__name__ == "array_api_strict"

这段代码定义了 8 个 is_*_namespace 函数,每个都通过比较 xp.__name__ 与已知模块名来判断。_compat_module_name() 返回当前 compat 库的完整模块路径(用于 vendor 后能正确识别)。这些函数为多后端代码提供运行时分支能力——例如某些算法只对 GPU 存在端优化,就可以在 is_cupy_namespace(xp) or is_torch_namespace(xp) 时启用特殊路径。

65.9.6 is_*_array 系列函数

与命名空间判断对应,兼容层也提供了对数组对象本身的类型判断:

源码路径:sklearn/externals/array_api_compat/common/_helpers.py - is_*_array()

def is_numpy_array(x: object) -> TypeIs[npt.NDArray[Any]:
    """Return True if `x` is a NumPy array."""
    cls = cast(Hashable, type(x))
    return (
        _issubclass_fast(cls, "numpy", "ndarray")
        or _issubclass_fast(cls, "numpy", "generic")
    ) and not _is_jax_zero_gradient_array(x)

def is_cupy_array(x: object) -> bool:
    cls = cast(Hashable, type(x))
    return _issubclass_fast(cls, "cupy", "ndarray")

def is_torch_array(x: object) -> TypeIs[torch.Tensor]:
    cls = cast(Hashable, type(x))
    return _issubclass_fast(cls, "torch", "Tensor")

def is_dask_array(x: object) -> TypeIs[da.Array]:
    cls = cast(Hashable, type(x))
    return _issubclass_fast(cls, "dask.array", "Array")

def is_jax_array(x: object) -> TypeIs[jax.Array]:
    cls = cast(Hashable, type(x))
    return (
        _issubclass_fast(cls, "jax", "Array")
        or _issubclass_fast(cls, "jax.core", "Tracer")
        or _is_jax_zero_gradient_array(x)
    )

def is_pydata_sparse_array(x: object) -> TypeIs[sparse.SparseArray]:
    cls = cast(Hashable, type(x))
    return _issubclass_fast(cls, "sparse", "SparseArray")

def is_array_api_obj(x: object) -> TypeGuard[_ArrayApiObj]:
    return (
        hasattr(x, '__array_namespace__')
        or _is_array_api_cls(cast(Hashable, type(x)))
    )

is_*_namespace 不同,is_*_array 检查的是对象本身的类型(如 np.ndarraycp.ndarraytorch.Tensor)。is_jax_array 额外检查 jax.core.Tracer,因为从 JAX 0.8.2 开始,tracer 不再是 jax.Array 的子类。_is_jax_zero_gradient_array 用于识别 jax.float0 这种用 NumPy void dtype 伪装的零梯度数组。

65.9.7 设备抽象与数据转移

device()to_device() 提供了统一的设备抽象:

源码路径:sklearn/externals/array_api_compat/common/_helpers.py - device()

def device(x: _ArrayApiObj, /) -> Device:
    """Hardware device the array data resides on."""
    if is_numpy_array(x):
        return "cpu"
    elif is_dask_array(x):
        # Peek at the metadata of the Dask array to determine type
        if is_numpy_array(x._meta):
            # Must be on CPU if backed by numpy
            return "cpu"
        return _DASK_DEVICE
    elif is_jax_array(x):
        # FIXME Jitted JAX arrays do not have a device attribute
        x_device = getattr(x, "device", None)
        if inspect.ismethod(x_device):
            return x_device()
        else:
            return x_device
    elif is_pydata_sparse_array(x):
        x_device = getattr(x, "device", None)
        if x_device is not None:
            return x_device
        # Everything but DOK has this attr.
        try:
            inner = x.data
        except AttributeError:
            return "cpu"
        return device(inner)
    return x.device

NumPy 和 Dask 没有真正的多设备概念,所以 NumPy 永远返回 "cpu",Dask 用 _DASK_DEVICE 特殊对象表示"非 CPU 设备"。JAX 在 jax.jit 上下文中 .device 可能返回 None,此时函数返回 None(虽然违反标准,但避免崩溃)。Sparse 数组会回退到其内部 .data 数组的设备。

源码路径:sklearn/externals/array_api_compat/common/_helpers.py - to_device()

def to_device(x: Array, device: Device, /, *, stream: int | Any | None = None) -> Array:
    """Copy the array from the device on which it currently resides to the specified ``device``."""
    if is_numpy_array(x):
        if stream is not None:
            raise ValueError("The stream argument to to_device() is not supported")
        if device == "cpu":
            return x
        raise ValueError(f"Unsupported device {device!r}")
    elif is_cupy_array(x):
        return _cupy_to_device(x, device, stream=stream)
    elif is_torch_array(x):
        return _torch_to_device(x, device, stream=stream)
    elif is_dask_array(x):
        if stream is not None:
            raise ValueError("The stream argument to to_device() is not supported")
        if device == "cpu":
            return x
        raise ValueError(f"Unsupported device {device!r}")
    elif is_jax_array(x):
        if not hasattr(x, "__array_namespace__"):
            import jax.experimental.array_api
            if not hasattr(x, "to_device"):
                return x
        return x.to_device(device, stream=stream)
    elif is_pydata_sparse_array(x) and device == _device(x):
        return x
    return x.to_device(device, stream=stream)

to_device 将数组从当前设备转移到目标设备。NumPy 和 Dask 不支持真正的多设备,对非 CPU 设备报错。CuPy 调用内部 _cupy_to_device(支持 stream 参数)。PyTorch 调用 _torch_to_device(不支持 stream)。JAX 在 jax.jit 上下文中 to_device 可能无效,函数会做兜底处理。

源码路径:sklearn/externals/array_api_compat/common/_helpers.py - _cupy_to_device()

def _cupy_to_device(
    x: cp.ndarray,
    device: Device,
    /,
    stream: int | Any | None = None,
) -> cp.ndarray:
    if device == "cpu":
        return x.get()
    if not isinstance(device, cp.cuda.Device):
        raise TypeError(f"Unsupported device type {device!r}")

    if stream is None:
        with device:
            return cp.asarray(x)

    if isinstance(stream, int):
        stream = cp.cuda.ExternalStream(stream)
    elif not isinstance(stream, cp.cuda.Stream):
        raise TypeError(f"Unsupported stream type {stream!r}")

    with device, stream:
        return cp.asarray(x)

CuPy 的设备转移支持 cp.cuda.Device 上下文管理器与 cp.cuda.Stream 流对象。stream 可以是 int(dlpack 格式)或 CuPy Stream。

源码路径:sklearn/externals/array_api_compat/common/_helpers.py - _torch_to_device()

def _torch_to_device(
    x: torch.Tensor,
    device: torch.device | str | int,
    /,
    stream: int | Any | None = None,
) -> torch.Tensor:
    if stream is not None:
        raise NotImplementedError
    return x.to(device)

PyTorch 的 .to() 方法本身支持流参数,但简单起见,这里直接抛 NotImplementedError

65.9.8 零开销类型检测

_issubclass_fast 是整个兼容层性能的关键。它通过 sys.modules 缓存避免重复导入:

源码路径:sklearn/externals/array_api_compat/common/_helpers.py - _issubclass_fast()

@lru_cache(100)
def _issubclass_fast(cls: type, modname: str, clsname: str) -> bool:
    try:
        # 通过 sys.modules 缓存避免重复导入
        mod = sys.modules[modname]
    except KeyError:
        # 模块未加载则直接返回 False,不触发 import
        return False
    # 获取父类引用
    parent_cls = getattr(mod, clsname)
    # 标准 issubclass 检查
    return issubclass(cls, parent_cls)

这段代码定义了高性能的 issubclass 检查。cls 是待检测类型,modname 是父类所在模块名,clsname 是父类名。通过 sys.modules.get(modname) 避免触发 import(如果模块未加载则直接返回 False),然后用标准 issubclass 检查继承关系。@lru_cache(100) 缓存最近 100 个组合的查询结果。

这种设计避免了两种性能陷阱:一是无意义的模块导入(用户代码可能只用 NumPy,不应该触发 CuPy 导入),二是反复的属性查找(缓存了 mod 和 parent_cls 的引用)。

65.9.9 惰性与可写性判断

is_lazy_arrayis_writeable_array 提供了跨后端的惰性求值与可写性判断:

源码路径:sklearn/externals/array_api_compat/common/_helpers.py - is_lazy_array()is_writeable_array()

@lru_cache(100)
def _is_writeable_cls(cls: type) -> bool | None:
    if (
        _issubclass_fast(cls, "numpy", "generic")
        or _issubclass_fast(cls, "jax", "Array")
        or _issubclass_fast(cls, "jax.core", "Tracer")
        or _issubclass_fast(cls, "sparse", "SparseArray")
    ):
        return False
    if _is_array_api_cls(cls):
        return True
    return None

def is_writeable_array(x: object) -> TypeGuard[_ArrayApiObj]:
    cls = cast(Hashable, type(x))
    if _issubclass_fast(cls, "numpy", "ndarray"):
        return cast("npt.NDArray", x).flags.writeable
    res = _is_writeable_cls(cls)
    if res is not None:
        return res
    return hasattr(x, '__array_namespace__')

@lru_cache(100)
def _is_lazy_cls(cls: type) -> bool | None:
    if (
        _issubclass_fast(cls, "numpy", "ndarray")
        or _issubclass_fast(cls, "numpy", "generic")
        or _issubclass_fast(cls, "cupy", "ndarray")
        or _issubclass_fast(cls, "torch", "Tensor")
        or _issubclass_fast(cls, "sparse", "SparseArray")
    ):
        return False
    if (
        _issubclass_fast(cls, "jax", "Array")
        or _issubclass_fast(cls, "jax.core", "Tracer")
        or _issubclass_fast(cls, "dask.array", "Array")
        or _issubclass_fast(cls, "ndonnx", "Array")
    ):
        return True
    return None

def is_lazy_array(x: object) -> TypeGuard[_ArrayApiObj]:
    cls = cast(Hashable, type(x))
    res = _is_lazy_cls(cls)
    if res is not None:
        return res
    if not hasattr(x, "__array_namespace__"):
        return False
    s = size(cast("HasShape[Collection[SupportsIndex | None]]", x))
    if s is None:
        return True
    xp = array_namespace(x)
    if s > 1:
        x = xp.reshape(x, (-1,))[0]
    x = xp.any(x)
    try:
        bool(x)
        return False
    except Exception:
        return True

NumPy 标量、JAX 数组、Sparse 数组默认不可写;NumPy ndarray 需要检查 .flags.writeable。NumPy/CuPy/PyTorch/Sparse 默认是急切执行的;JAX/Dask/NDonnx 默认是惰性的。对于未知类型,is_lazy_array 会尝试 bool(x) 触发求值,若抛异常则为视为 lazy(这是对 Dask 的一种特殊情况处理——Dask 在 __bool__ 上会触发求图计算,但应该被视为 lazy)。

65.9.10 size 函数

size 是跨后端的元素总数查询(处理 None 与 NaN):

源码路径:sklearn/externals/array_api_compat/common/_helpers.py - size()

def size(x: HasShape[Collection[SupportsIndex | None]], /) -> int | None:
    """Return the total number of elements of x."""
    if None in x.shape:
        return None
    out = math.prod(cast("Collection[SupportsIndex]", x.shape))
    return None if math.isnan(out) else out

PyTorch 的 Tensor.size() 与 Array API 的 array.size 语义不一致(前者返回 shape,后者返回元素总数)。Dask 对未知形状返回 NaN。这里统一为:含 None 的 shape 返回 None(符合 Array API),NaN 也转 None。

65.9.11 dir 自定义

每个 _helpers.py 模块末尾都定义了 __dir__

源码路径:sklearn/externals/array_api_compat/common/_helpers.py - __dir__()

__all__ = [
    "array_namespace",
    "device",
    "get_namespace",
    ...
]

def __dir__() -> list[str]:
    return __all__

__dir__ 影响 IDE 自动补全与 dir() 的输出。通过自定义 __dir__,可以确保只暴露公开 API(避免内部辅助函数污染补全列表)。这是 Python 跨库兼容性的常见技巧。

65.10 通用别名与函数包装

common/_aliases.py 是跨后端函数语义对齐的标准化车间。所有后端适配器都可以复用这里的实现,仅在必要时做定制。

下面用 Mermaid 流程图展示通用别名层的设计模式:

flowchart LR subgraph "通用别名层" Clip["clip dtype 保持"] Unique["unique_* 返回 NamedTuple"] Sort["sort/argsort descending"] CumSum["cumulative_sum include_initial"] Create["arange/empty/eye 添加 device"] Reshape["reshape copy 语义"] MatT["matrix_transpose 最后两轴"] Unstack["unstack 2023.12 新增"] end Clip --> Reuse["get_xp(backend)(通用)"] Unique --> Reuse Sort --> Reuse CumSum --> Reuse Create --> Reuse Reshape --> Reuse MatT --> Reuse Unstack --> Reuse Reuse --> Backend["后端模块"]

65.10.1 创建函数族

arangeemptyeyelinspaceoneszerosfull 等创建函数都添加了 device 参数:

源码路径:sklearn/externals/array_api_compat/common/_aliases.py - arange/empty/eye/linspace/ones/zeros/full

def arange(
    start: float, /, stop: float | None = None, step: float = 1,
    *, xp: Namespace, dtype: DType | None = None, device: Device | None = None,
    **kwargs: object,
) -> Array:
    _check_device(xp, device)
    return xp.arange(start, stop=stop, step=step, dtype=dtype, **kwargs)

def empty(
    shape: int | tuple[int, ...], xp: Namespace,
    *, dtype: DType | None = None, device: Device | None = None, **kwargs: object,
) -> Array:
    _check_device(xp, device)
    return xp.empty(shape, dtype=dtype, **kwargs)

def eye(
    n_rows: int, n_cols: int | None = None, /,
    *, xp: Namespace, k: int = 0, dtype: DType | None = None,
    device: Device | None = None, **kwargs: object,
) -> Array:
    _check_device(xp, device)
    return xp.eye(n_rows, M=n_cols, k=k, dtype=dtype, **kwargs)

def linspace(
    start: float, stop: float, /, num: int,
    *, xp: Namespace, dtype: DType | None = None, device: Device | None = None,
    endpoint: bool = True, **kwargs: object,
) -> Array:
    _check_device(xp, device)
    return xp.linspace(start, stop, num, dtype=dtype, endpoint=endpoint, **kwargs)

def ones/zeros/full 类似,调用 xp.ones/zeros/full 并传 device

每个函数都用 _check_device(xp, device) 验证 device 对当前后端的合法性(NumPy/Dask 只接受 "cpu",CuPy/PyTorch 接受 GPU),然后调用底层 xp.arange/xp.empty 等原生函数。注意 eyeM= 关键字(NumPy 用法),而 Array API 用 n_cols

65.10.2 clip 函数的 dtype 保持

clip 是 Array API 中语义最复杂的函数之一——它要求输出与输入同 dtype,但 NumPy 原生 np.clip 会做类型提升:

源码路径:sklearn/externals/array_api_compat/common/_aliases.py - clip()

def clip(
    x: Array, /, min: float | Array | None = None, max: float | Array | None = None,
    *, xp: Namespace, out: Array | None = None,
) -> Array:
    def _isscalar(a: object) -> TypeIs[float | None]:
        return isinstance(a, int | float) or a is None

    min_shape = () if _isscalar(min) else min.shape
    max_shape = () if _isscalar(max) else max.shape

    wrapped_xp = array_namespace(x)
    result_shape = xp.broadcast_shapes(x.shape, min_shape, max_shape)

    # 关键技巧:手动分配与 x 同 dtype 的输出缓冲区
    dev = _get_device(x)
    if out is None:
        out = wrapped_xp.empty(result_shape, dtype=x.dtype, device=dev)
    assert out is not None
    out[()] = x  # 先复制 x 的内容(保持 dtype)

    # 处理 Python 整数溢出:int8 范围 [-128, 127]
    if wrapped_xp.isdtype(x.dtype, "integral"):
        if type(min) is int and min <= wrapped_xp.iinfo(x.dtype).min:
            min = None
        if type(max) is int and max >= wrapped_xp.iinfo(x.dtype).max:
            max = None

    if min is not None:
        a = wrapped_xp.asarray(min, dtype=x.dtype, device=dev)
        a = xp.broadcast_to(a, result_shape)
        ia = (out < a) | xp.isnan(a)
        out[ia] = a[ia]

    if max is not None:
        b = wrapped_xp.asarray(max, dtype=x.dtype, device=dev)
        b = xp.broadcast_to(b, result_shape)
        ib = (out > b) | xp.isnan(b)
        out[ib] = b[ib]

    return out[()]

这段代码定义了 Array API 兼容的 clip 函数。核心技巧是手动分配输出缓冲区out = wrapped_xp.empty(result_shape, dtype=x.dtype, device=dev) 创建与 x 同 dtype 的空数组,然后 out[()] = x 复制 x 的内容。这样绕过了 NumPy np.clip 的类型提升(int8 + int64 提升为 int64)。

Python 整数溢出处理:当 x 是整数类型且边界值超出范围时(如 min=128, x.dtype=int8),直接将边界设为 None,避免溢出截断。xp.isnan(a) 处理 NaN 边界值——任何值与 NaN 比较都是 False,所以 NaN 边界实际上不会限制 x。

65.10.3 unique 系列返回类型标准化

unique_all 等函数返回 UniqueAllResult NamedTuple 而不是普通 tuple:

源码路径:sklearn/externals/array_api_compat/common/_aliases.py - unique_all()

class UniqueAllResult(NamedTuple):
    values: Array
    indices: Array
    inverse_indices: Array
    counts: Array

def unique_all(x: Array, /, xp: Namespace) -> UniqueAllResult:
    kwargs = _unique_kwargs(xp)
    values, indices, inverse_indices, counts = xp.unique(
        x,
        return_counts=True,
        return_index=True,
        return_inverse=True,
        **kwargs,
    )
    # 关键修复:np.unique() 会将 inverse_indices 展平
    # 但 Array API 要求它保持 x 的原始 shape
    inverse_indices = inverse_indices.reshape(x.shape)
    return UniqueAllResult(
        values,
        indices,
        inverse_indices,
        counts,
    )

这段代码定义了 unique_all 函数。它调用 xp.unique 同时返回 values、indices、inverse_indices、counts,然后做关键的 shape 修复——inverse_indices.reshape(x.shape)。这是因为 np.unique 会将所有返回值展平(即使输入是多维),但数组 API 要求 inverse_indices 与输入同 shape。UniqueAllResult NamedTuple 让用户可以通过 .values.counts 等属性访问,比元组解包更清晰。

源码路径:sklearn/externals/array_api_compat/common/_aliases.py - _unique_kwargs()

def _unique_kwargs(xp: Namespace) -> dict[str, bool]:
    # 老版本 NumPy/CuPy 没有 equal_nan。检查签名而非解析版本号
    if "equal_nan" in inspect.signature(xp.unique).parameters:
        return {"equal_nan": False}
    return {}

这段代码动态检查 xp.unique 是否支持 equal_nan 参数。这种运行时检查避免了版本号解析的复杂性——它直接通过 inspect.signature 查看参数列表。这种模式在兼容层中很常见。

65.10.4 std/var correction 参数重命名

stdvarddof 重命名为 correction

源码路径:sklearn/externals/array_api_compat/common/_aliases.py - std() / var()

def std(
    x: Array, /, xp: Namespace,
    *, axis: int | tuple[int, ...] | None = None,
    correction: float = 0.0,  # correction instead of ddof
    keepdims: bool = False,
    **kwargs: object,
) -> Array:
    return xp.std(x, axis=axis, ddof=correction, keepdims=keepdims, **kwargs)

def var(
    x: Array, /, xp: Namespace,
    *, axis: int | tuple[int, ...] | None = None,
    correction: float = 0.0,
    keepdims: bool = False,
    **kwargs: object,
) -> Array:
    return xp.var(x, axis=axis, ddof=correction, keepdims=keepdims, **kwargs)

这是 Array API 与 NumPy 的命名差异——correction 代替 ddof。函数内部做名称映射,对用户隐藏差异。

65.10.5 累加函数 cumulative_sum / cumulative_prod

cumulative_sumnp.cumsum 的 Array API 兼容版本,扩展了 include_initial 参数:

源码路径:sklearn/externals/array_api_compat/common/_aliases.py - cumulative_sum()

def cumulative_sum(
    x: Array, /, xp: Namespace,
    *, axis: int | None = None, dtype: DType | None = None,
    include_initial: bool = False, **kwargs: object,
) -> Array:
    wrapped_xp = array_namespace(x)

    if axis is None:
        if x.ndim > 1:
            raise ValueError(
                "axis must be specified in cumulative_sum for more than one dimension"
            )
        axis = 0

    res = xp.cumsum(x, axis=axis, dtype=dtype, **kwargs)

    if include_initial:
        initial_shape = list(x.shape)
        initial_shape[axis] = 1
        res = xp.concatenate(
            [
                wrapped_xp.zeros(
                    shape=initial_shape, dtype=res.dtype, device=_get_device(res)
                ),
                res,
            ],
            axis=axis,
        )
    return res

这段代码定义了 cumulative_sum。核心逻辑是调用 xp.cumsum,然后用 concatenate 手动实现 include_initial=True 的语义——即在结果前拼接一个 0cumulative_prod 实现类似,只是用 ones 而非 zeros,并调用 cumprod 而非 cumsum

65.10.6 排序与参数重命名

sortargsort 添加了 Array API 特有的 descendingstable 参数:

源码路径:sklearn/externals/array_api_compat/common/_aliases.py - sort() / argsort()

def argsort(
    x: Array, /, xp: Namespace,
    *, axis: int = -1, descending: bool = False, stable: bool = True,
    **kwargs: object,
) -> Array:
    if stable:
        kwargs["kind"] = "stable"
    if not descending:
        res = xp.argsort(x, axis=axis, **kwargs)
    else:
        # 降序:翻转输入 → 升序排序 → 翻转结果 → 取反索引
        res = xp.flip(
            xp.argsort(xp.flip(x, axis=axis), axis=axis, **kwargs),
            axis=axis,
        )
        normalised_axis = axis if axis >= 0 else x.ndim + axis
        max_i = x.shape[normalised_axis] - 1
        res = max_i - res
    return res


def sort(
    x: Array, /, xp: Namespace,
    *, axis: int = -1, descending: bool = False, stable: bool = True,
    **kwargs: object,
) -> Array:
    if stable:
        kwargs["kind"] = "stable"
    res = xp.sort(x, axis=axis, **kwargs)
    if descending:
        res = xp.flip(res, axis=axis)
    return res

这段代码实现了支持 descendingstable 参数的 sort/argsortsort 简单——升序排序后翻转。argsort 复杂一些——不能简单翻转索引(会破坏相对顺序),所以采用"翻转输入→升序→翻转结果→取反索引"的方式,确保稳定的降序排序。

65.10.7 维度变换与索引函数

permute_dimsreshapenonzerotensordot 等函数处理 Array API 与 NumPy 的参数名/语义差异:

源码路径:sklearn/externals/array_api_compat/common/_aliases.py - permute_dims/reshape/nonzero/matrix_transpose

def permute_dims(x: Array, /, axes: tuple[int, ...], xp: Namespace) -> Array:
    return xp.transpose(x, axes)

def reshape(
    x: Array, /, shape: tuple[int, ...], xp: Namespace,
    *, copy: bool | None = None, **kwargs: object,
) -> Array:
    if copy is True:
        x = x.copy()
    elif copy is False:
        y = x.view()
        y.shape = shape
        return y
    return xp.reshape(x, shape, **kwargs)

def nonzero(x: Array, /, xp: Namespace, **kwargs: object) -> tuple[Array, ...]:
    if x.ndim == 0:
        raise ValueError("nonzero() does not support zero-dimensional arrays")
    return xp.nonzero(x, **kwargs)

def matrix_transpose(x: Array, /, xp: Namespace) -> Array:
    if x.ndim < 2:
        raise ValueError("x must be at least 2-dimensional for matrix_transpose")
    return xp.swapaxes(x, -1, -2)

def tensordot(
    x1: Array, x2: Array, /, xp: Namespace,
    *, axes: int | tuple[Sequence[int], Sequence[int]] = 2,
    **kwargs: object,
) -> Array:
    return xp.tensordot(x1, x2, axes=axes, **kwargs)

这段代码展示了 Array API 与 NumPy 在参数名(newshape vs shape)、必需性(axes 必需 vs 可选)、语义(matrix_transpose vs transpose)上的差异处理。

65.10.8 线性代数核心包装

matmulvecdotoutercross 是 Array API 中跨后端语义统一的线性代数基础函数:

源码路径:sklearn/externals/array_api_compat/common/_aliases.py - matmul/vecdot/outer/cross

def matmul(x1: Array, x2: Array, /, xp: Namespace, **kwargs: object) -> Array:
    return xp.matmul(x1, x2, **kwargs)

def vecdot(x1: Array, x2: Array, /, xp: Namespace, *, axis: int = -1) -> Array:
    if x1.shape[axis] != x2.shape[axis]:
        raise ValueError("x1 and x2 must have the same size along the given axis")
    if hasattr(xp, "broadcast_tensors"):
        _broadcast = xp.broadcast_tensors
    else:
        _broadcast = xp.broadcast_arrays
    x1_ = xp.moveaxis(x1, axis, -1)
    x2_ = xp.moveaxis(x2, axis, -1)
    x1_, x2_ = _broadcast(x1_, x2_)
    res = xp.conj(x1_[..., None, :]) @ x2_[..., None]
    return res[..., 0, 0]

vecdotmatmul 手动实现——移动收缩轴到最后一维,构造单元素矩阵相乘,取 [..., 0, 0] 得到标量。这样可以处理任意 dtype(PyTorch 的 vecdot 不支持整数 dtype)。

65.10.9 类型判断与设备/精度信息

isdtypefinfoiinfo 提供跨后端的类型元信息查询:

源码路径:sklearn/externals/array_api_compat/common/_aliases.py - isdtype/finfo/iinfo/sign/unstack

def isdtype(
    dtype: DType, kind: DType | str | tuple[DType | str, ...],
    xp: Namespace, *, _tuple: bool = True,
) -> bool:
    if isinstance(kind, tuple) and _tuple:
        return any(
            isdtype(dtype, k, xp, _tuple=False)
            for k in cast("tuple[DType | str, ...]", kind)
        )
    elif isinstance(kind, str):
        if kind == "bool":
            return dtype == xp.bool_
        elif kind == "signed integer":
            return xp.issubdtype(dtype, xp.signedinteger)
        elif kind == "unsigned integer":
            return xp.issubdtype(dtype, xp.unsignedinteger)
        elif kind == "integral":
            return xp.issubdtype(dtype, xp.integer)
        elif kind == "real floating":
            return xp.issubdtype(dtype, xp.floating)
        elif kind == "complex floating":
            return xp.issubdtype(dtype, xp.complexfloating)
        elif kind == "numeric":
            return xp.issubdtype(dtype, xp.number)
        else:
            raise ValueError(f"Unrecognized data type kind: {kind!r}")
    else:
        return dtype == kind

def finfo(type_: DType | Array, /, xp: Namespace) -> Any:
    try:
        return xp.finfo(type_)
    except (ValueError, TypeError):
        return xp.finfo(type_.dtype)

def iinfo(type_: DType | Array, /, xp: Namespace) -> Any:
    try:
        return xp.iinfo(type_)
    except (ValueError, TypeError):
        return xp.iinfo(type_.dtype)

def sign(x: Array, /, xp: Namespace, **kwargs: object) -> Array:
    if isdtype(x.dtype, "complex floating", xp=xp):
        out = (x / xp.abs(x, **kwargs))[...]
        out[x == 0j] = 0j
    else:
        out = xp.sign(x, **kwargs)
    if is_cupy_namespace(xp) and isdtype(x.dtype, "real floating", xp=xp):
        out[xp.isnan(x)] = xp.nan
    return out[()]

def unstack(x: Array, /, xp: Namespace, *, axis: int = 0) -> tuple[Array, ...]:
    if x.ndim == 0:
        raise ValueError("Input array must be at least 1-d.")
    return tuple(xp.moveaxis(x, axis, 0))

isdtype 是 Array API 2022.12 规范的新函数,提供跨后端的类型分类查询。finfo/iinfo 处理 dtype 与 array 双输入的差异。sign 在复数上需要特殊处理(NumPy 1.26 与 Array API 语义不同)。unstack 是 2023.12 新增的函数。

65.11 线性代数与 FFT 标准化

下面用 Mermaid 流程图展示线性代数分解函数的统一返回类型设计:

flowchart LR subgraph "返回类型标准化" SVD["SVDResult<br/>U, S, Vh"] Eigh["EighResult<br/>eigenvalues, eigenvectors"] QR["QRResult<br/>Q, R"] Slog["SlogdetResult<br/>sign, logabsdet"] end Input["调用 svd(x)"] --> SVD Input2["调用 eigh(x)"] --> Eigh Input3["调用 qr(x)"] --> QR Input4["调用 slogdet(x)"] --> Slog

65.11.1 命名元组返回类型

eighqrsvdslogdet 等分解函数统一返回命名元组:

源码路径:sklearn/externals/array_api_compat/common/_linalg.py - svd/eigh/qr/slogdet/cross/outer

class EighResult(NamedTuple):
    eigenvalues: Array
    eigenvectors: Array

class QRResult(NamedTuple):
    Q: Array
    R: Array

class SlogdetResult(NamedTuple):
    sign: Array
    logabsdet: Array

class SVDResult(NamedTuple):
    U: Array
    S: Array
    Vh: Array

def svd(x: Array, /, xp: Namespace, *, full_matrices: bool = True, **kwargs: object) -> SVDResult:
    return SVDResult(*xp.linalg.svd(x, full_matrices=full_matrices, **kwargs))

def eigh(x: Array, /, xp: Namespace, **kwargs: object) -> EighResult:
    return EighResult(*xp.linalg.eigh(x, **kwargs))

def qr(x: Array, /, xp: Namespace, *, mode: Literal["reduced", "complete"] = "reduced", **kwargs: object) -> QRResult:
    return QRResult(*xp.linalg.qr(x, mode=mode, **kwargs))

def slogdet(x: Array, /, xp: Namespace, **kwargs: object) -> SlogdetResult:
    return SlogdetResult(*xp.linalg.slogdet(x, **kwargs))

def cross(x1: Array, x2: Array, /, xp: Namespace, *, axis: int = -1, **kwargs: object) -> Array:
    return xp.cross(x1, x2, axis=axis, **kwargs)

def outer(x1: Array, x2: Array, /, xp: Namespace, **kwargs: object) -> Array:
    return xp.outer(x1, x2, **kwargs)
posted @ 2026-09-04 04:07  绝不原创的飞龙  阅读(5)  评论(0)    收藏  举报