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 # 返回概率数组
要点:与二元对数损失不同,指数损失使用 HalfLogitLink(raw → 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)
63.7.4 代码解读:BaseLink 与子类(第 87‑165 行)
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__ 中复制给相应损失实例,实现 输入合法性验证。IdentityLink 的 inverse 直接等于 link,因为恒等映射的逆仍是恒等。MultinomialLogit.link 使用 几何均值 作为参考类别,使得 对称多项 Logit 与 softmax 形成互逆关系。
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_range 与 in_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_gradient 与 cy_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,而在 中间区间 使用 log1p 或 log(1+exp(x)) 以保持精度。常数在 float64 与 float32 上不同,以适配各自的数值范围,确保 相对误差 在机器精度以内。
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(如 loss、gradient 等)均委托给 closs。这种设计的最大优势在于:
-
可维护性提升:Cython 与 Python 层的职责清晰分离,修改链接函数或 Cython 实现不会相互影响。
-
跨后端兼容:组合方式天然支持后续的 Array API 实现,只需要在
closs旁边提供对应的 Array API 包装即可。 -
明确的错误定位:当出现数值异常时,可以直接定位到 Cython 实现或链接函数,便于调试。
然而,这种组合也带来了一些细微的代价:
-
每次调用时需要通过属性访问
self.closs,在极端的微基准测试中会产生极少量的 Python 级别开销。 -
子类必须在构造函数中显式传入
closs与link,代码略显冗长。
综合来看,可维护性、跨后端适配以及对 Cython 已知局限的规避的收益远远超过了这点微小的运行时开销。因此,组合模式成为了 scikit‑learn 损失函数体系的最佳实现方案。
63.11 动手练习
练习中的每一项都已改写为完整的段落描述,避免使用简短列表形式。
63.11.1 练习 1:阅读 BaseLoss 基类的完整实现
阅读 sklearn/_loss/loss.py 中 BaseLoss 类(第 40‑270 行)的完整代码。重点思考以下问题:
-
为什么在构造函数中使用组合而不是多重继承来绑定 Cython 损失与链接函数?(提示:Cython Issue #4350)
-
loss、gradient、loss_gradient、gradient_hessian四个方法是如何通过self.closs将计算委托给 Cython 实现的? -
fit_intercept_only如何利用self.link把目标的均值/中位数/分位数映射到链接空间(raw_prediction)? -
constant_to_optimal_zero在损失计算中起什么作用?为什么要把它单独抽离为一个方法? -
init_gradient_and_hessian如何根据self.constant_hessian判断是分配完整的海森矩阵还是仅仅分配一个标量?
63.11.2 练习 2:比较回归损失函数的数学特性与实现差异
阅读 sklearn/_loss/loss.py 第 280‑450 行,比较以下回归损失的实现细节:
-
HalfSquaredError、AbsoluteError、PinballLoss、HuberLoss在differentiable、need_update_leaves_values、approx_hessian、constant_hessian四个属性上的区别。 -
HalfPoissonLoss、HalfGammaLoss、HalfTweedieLoss如何通过LogLink处理正实数约束?它们的constant_to_optimal_zero分别补全了哪些常数项? -
HalfTweedieLossIdentity与HalfTweedieLoss的区别是什么?在power参数变化时,interval_y_pred如何自适应? -
为什么
AbsoluteError与PinballLoss的fit_intercept_only返回加权中位数/分位数,而HalfSquaredError返回加权平均?
63.11.3 练习 3:深入分类损失与链接函数的数学原理
阅读 sklearn/_loss/loss.py 第 452‑670 行以及 sklearn/_loss/link.py 第 92‑115 行,分析以下内容:
-
HalfBinomialLoss的损失公式log(1+exp(raw)) - y*raw与交叉熵-y·log(p) - (1-y)·log(1-p)如何等价?请给出完整的代数推导。 -
LogitLink与HalfLogitLink的区别是什么?为什么ExponentialLoss使用后者? -
MultinomialLogit.link为什么使用几何均值作为参考类别?这如何导致原始预测的每行求和为零的约束? -
HalfMultinomialLoss.gradient_proba为什么需要同时返回梯度和概率?在HistGradientBoostingClassifier中这些信息如何被使用?
63.11.4 练习 4:探究 Cython 底层实现的高性能机制
阅读 sklearn/_loss/_loss.pxd,弄清以下细节:
-
floating_in与floating_out这两组融合类型的设计意图是什么?它们如何实现输入输出 dtype 的不一致? -
CyLossFunction中的三个纯虚函数cy_loss、cy_gradient、cy_grad_hess对应的计算粒度分别是什么? -
CyHalfMultinomialLoss为什么不继承CyLossFunction?它的cy_gradient签名有何特点? -
readonly与public修饰符在CyPinballLoss.quantile与CyHuberLoss.delta上的区别是什么?它们对 Python 侧属性访问有什么影响?
63.11.5 练习 5:实现自定义损失函数并验证数值正确性
参考 sklearn/_loss/loss.py 中已有的损失实现,尝试实现一个自定义损失类 MyCustomLoss,要求满足以下条件:
-
继承
BaseLoss,组合一个新的 Cython 损失类(可以在_loss.pxd中声明,或直接复用已有 Cython 类)。 -
为其选择或实现一个合适的链接函数(可以参考
link.py中的BaseLink子类)。 -
正确设置
differentiable、need_update_leaves_values、approx_hessian、constant_hessian等属性。 -
实现
fit_intercept_only与constant_to_optimal_zero,确保截距模型与常数项与数学定义一致。 -
编写单元测试,验证:损失非负、梯度在最优点为零、数值微分验证梯度与海森、样本权重正确广播、支持 float32/float64、可 pickle 序列化。请参考
tests/test_loss.py中的对应测试用例。
63.11.6 练习 6:分析 Array API 兼容层的数值稳定技巧
阅读 sklearn/_loss/loss.py 第 672‑780 行,重点分析:
-
_log1pexp为何需要四段分支处理?针对 float64 与 float32,常数-37/-2/18/33.3与-17/-1/9/14.6分别对应哪种数值近似? -
在
HalfBinomialLossArrayAPI._compute_gradient中,xp.where的条件分支如何避免在极端正/负值下的数值溢出? -
HalfMultinomialLossArrayAPI._compute_loss如何利用class_indexing_offsets与y_true_int实现 无 one‑hot 编码 的真实标签概率提取? -
注释中提到当前无法使用增量赋值(
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 架构与数据流图
上述图分别展示模块依赖、调用时序、数据流和架构分层。
63.14 设计取舍的深度思考
为什么采用当前方案而不是更复杂的替代方案? 本章实现优先保证与既有 API 的一致性、可维护性与运行效率。这意味着在少数极端场景下,调用者需要自行在灵活性、内存与速度之间做取舍,换取默认路径的清晰与稳定。
第 64 章 —— sklearn.externals._arff 源码解析
64.1 学习目标
-
难度:★★★☆☆(3/5)
-
预备知识:Python 基础、面向对象编程与 Markdown/代码阅读基础
本节旨在帮助读者:
-
了解 ARFF(Attribute‑Relation File Format)文件的结构及其在机器学习实验中的作用。
-
掌握
sklearn.externals._arff模块提供的 读取 与 写入 接口,包括不同的数据结构(密集、稀疏、生成器)对应的返回类型。 -
能够根据实际需求选择合适的矩阵表示方式,并理解模块内部的 编码/解码 逻辑与异常处理机制。
-
认识模块关键类(
ArffDecoder,ArffEncoder,EncodedNominalConversor,NominalConversor,Data,COOData,LODData)的职责与实现细节。
温馨提示:ARFF 常用于 Weka、ML‑lib 等工具之间的数据交换,熟悉其解析流程有助于在跨平台实验中避免格式错误。
64.2 模块整体架构
说明:图中虚线表示可选路径,例如稀疏矩阵的 COO 与 LOD 表示,或使用 生成器 逐行写入,以降低内存占用。
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() 细节
-
quoted_re:匹配双引号包围的值,支持转义(
\"、\\)以及嵌套单引号。 -
value_re:允许三种形式的值——双引号、单引号或不含特殊字符的裸字符串。
-
dense 正则:捕获逗号分隔的值,包括空值(
?)与行尾。 -
sparse 正则:匹配
{index value}形式的稀疏数据,确保索引是整数且值遵循value_re。
这些正则表达式构成了解码过程的 词法 层,能够在不完整或异常的行上提供明确的错误定位。
64.4 数据结构常量
DENSE = 0 # 完全密集矩阵
COO = 1 # (row, col, value) 三元组的稀疏坐标格式
LOD = 2 # 每行一个 dict 的稀疏列表格式
DENSE_GEN = 3 # 密集数据的生成器(逐行产出)
LOD_GEN = 4 # 稀疏字典列表的生成器
-
密集 适合特征全部非缺失(如图像、表格)。
-
COO 与 LOD 适用于大部分为零的稀疏特征(如文本词袋)。
-
生成器 在处理 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 EncodedNominalConversor 与 NominalConversor
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前收集所有属性并为每个属性创建对应的 conversor(EncodedNominalConversor或NominalConversor,或基本类型的 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,COOData或LODData;稀疏矩阵会被转化为{index value, ...}形式。
64.7 关键函数逐行解读
| 函数 | 作用 | 关键实现细节 |
|------|------|--------------|
| _parse_values(s) | 将单行 ARFF 数据拆解为 Python 列表或稀疏字典 | - 首先检查是否包含需要正则处理的特殊字符
- 对于稠密行使用 _RE_DENSE_VALUES.findall 捕获 value 与 error
- 若发现稀疏行(_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 读写能力。掌握其内部类(尤其是 EncodedNominalConversor 与 NominalConversor)的职责,有助于在实际项目中:
-
选择合适的返回结构以平衡 内存 与 计算速度。
-
在出现数据不一致时快速定位并修正错误。
-
将自定义 Python 数据结构无缝导出为符合 Weka/ARFF 规范的文件。
后续阅读:了解
sklearn.externals._numpydoc与sklearn.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.py 中 COOData.decode_rows 与 LODGeneratorData.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.py 中 NumpyDocString._parse_see_also 与 _parse_param_list。
回答问题:
-
_line_rgx正则如何同时匹配func1, func2: description和:meth:func1: description两种语法?捕获组allfuncs、morefuncs、trailing、desc分别起什么作用? -
single_element_is_type=True在解析 Returns/Yields 时如何改变name: type与type两种写法的识别逻辑? -
ClassDoc如何通过inspect.getmembers发现公共方法与属性?_should_skip_member为何要特殊处理 namedtuple 的_fields?
64.14.3 PEP 440 版本比较键的生成细节
阅读 sklearn/externals/_packaging/version.py 中 Version.__init__ 与 _cmpkey 函数。
回答问题:
-
_cmpkey中为何对release元组进行reversed -> dropwhile(==0) -> reversed操作?举例说明1.0.0与1.0的比较结果。 -
pre为 None 且dev不为 None 时,为何将_pre设为NegativeInfinity?这如何保证1.0.dev0 < 1.0a0? -
local版本段的排序键为何将字符串段包装为(NegativeInfinity, str)、数字段包装为(int, '')?这如何实现“字母段 < 数字段”及“前缀匹配时短版本优先”? -
LegacyVersion如何通过_parse_version_parts将1.0.dev转换为可比较的元组?zfill(8)与*前缀的作用是什么?
64.14.4 稀疏图拉普拉斯的矩阵自由计算
阅读 sklearn/externals/_scipy/sparse/csgraph/_laplacian.py 中 _laplacian_sparse_flo 与 _laplace_normed_sym。
回答问题:
-
form='lo'时,_linearoperator如何封装矩阵向量乘积函数?matvec与matmat为何指向同一个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 的解析流程:
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,包含 name、type、desc 三个字段。dedent_lines 与 strip_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 的核心逻辑:
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 函数的多输出格式选择:
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() 时的整体数据流:
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_namespace、device、to_device、is_*_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.ndarray、cp.ndarray、torch.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_array 和 is_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 流程图展示通用别名层的设计模式:
65.10.1 创建函数族
arange、empty、eye、linspace、ones、zeros、full 等创建函数都添加了 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 等原生函数。注意 eye 用 M= 关键字(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 参数重命名
std 和 var 把 ddof 重命名为 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_sum 是 np.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 的语义——即在结果前拼接一个 0。cumulative_prod 实现类似,只是用 ones 而非 zeros,并调用 cumprod 而非 cumsum。
65.10.6 排序与参数重命名
sort 和 argsort 添加了 Array API 特有的 descending 和 stable 参数:
源码路径: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
这段代码实现了支持 descending 和 stable 参数的 sort/argsort。sort 简单——升序排序后翻转。argsort 复杂一些——不能简单翻转索引(会破坏相对顺序),所以采用"翻转输入→升序→翻转结果→取反索引"的方式,确保稳定的降序排序。
65.10.7 维度变换与索引函数
permute_dims、reshape、nonzero、tensordot 等函数处理 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 线性代数核心包装
matmul、vecdot、outer、cross 是 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]
vecdot 用 matmul 手动实现——移动收缩轴到最后一维,构造单元素矩阵相乘,取 [..., 0, 0] 得到标量。这样可以处理任意 dtype(PyTorch 的 vecdot 不支持整数 dtype)。
65.10.9 类型判断与设备/精度信息
isdtype、finfo、iinfo 提供跨后端的类型元信息查询:
源码路径: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 流程图展示线性代数分解函数的统一返回类型设计:
65.11.1 命名元组返回类型
eigh、qr、svd、slogdet 等分解函数统一返回命名元组:
源码路径: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)

浙公网安备 33010602011771号