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

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

这段代码定义了 SVDResult 等 NamedTuple 与对应的包装函数。SVDResult(*xp.linalg.svd(...)) 用星号解析原生返回的元组。用户可以写 U, S, Vh = svd(x) 解析,也可以写 result.U 访问单个值。

65.11.2 cholesky 的复共轭处理

cholesky 支持 upper 选项返回上三角,复杂之处在于复数矩阵需要共轭转置:

源码路径:sklearn/externals/array_api_compat/common/_linalg.py - cholesky()

def cholesky(x: Array, /, xp: Namespace, *, upper: bool = False, **kwargs: object) -> Array:
    L = xp.linalg.cholesky(x, **kwargs)
    if upper:
        U = get_xp(xp)(matrix_transpose)(L)
        if get_xp(xp)(isdtype)(U.dtype, 'complex floating'):
            U = xp.conj(U)
        return U
    return L

这段代码实现 cholesky 函数支持 upper 参数。xp.linalg.cholesky 仅返回下三角,要得到上三角就用 matrix_transpose 转置。但复数矩阵的"上三角"是下三角的共轭转置(不是普通转置),所以需要 xp.conj(U) 取共轭。

65.11.3 矩阵与向量范数

matrix_rankpinvmatrix_normvector_norm 处理 Array API 的 rtol 语义:

源码路径:sklearn/externals/array_api_compat/common/_linalg.py - matrix_rank/pinv/matrix_norm/vector_norm

def matrix_rank(
    x: Array, /, xp: Namespace, *, rtol: float | Array | None = None, **kwargs: object,
) -> Array:
    if x.ndim < 2:
        raise xp.linalg.LinAlgError("1-dimensional array given. Array must be at least two-dimensional")
    S = get_xp(xp)(svdvals)(x, **kwargs)
    if rtol is None:
        tol = S.max(axis=-1, keepdims=True) * max(x.shape[-2:]) * xp.finfo(S.dtype).eps
    else:
        tol = S.max(axis=-1, keepdims=True)*xp.asarray(xortol)[..., xp.newaxis]
    return xp.count_nonzero(S > tol, axis=-1)

def pinv(
    x: Array, /, xp: Namespace, *, rtol: float | Array | None = None, **kwargs: object,
) -> Array:
    if rtol is None:
        rtol = max(x.shape[-2:]) * xp.finfo(x.dtype).eps
    return xp.linalg.pinv(x, rcond=rtol, **kwargs)

def matrix_norm(
    x: Array, /, xp: Namespace, *, keepdims: bool = False,
    ord: Literal[1, 2, -1, -2] | JustFloat | Literal["fro", "nuc"] | None = "fro",
) -> Array:
    return xp.linalg.norm(x, axis=(-2, -1), keepdims=keepdims, ord=ord)

def vector_norm(
    x: Array, /, xp: Namespace,
    *, axis: int | tuple[int, ...] | None = None, keepdims: bool = False,
    ord: JustInt | JustFloat = 2,
) -> Array:
    if axis is None:
        _x = x.ravel()
        _axis = 0
    elif isinstance(axis, tuple):
        normalized_axis = cast(
            "tuple[int, ...]",
            normalize_axis_tuple(axis, x.ndim),
        )
        rest = tuple(i for i in range(x.ndim) if i not in normalized_axis)
        newshape = axis + rest
        _x = xp.transpose(x, newshape).reshape(
            (math.prod([x.shape[i] for i in axis]),], *[x.shape[i] for i in rest]))
        _axis = 0
    else:
        _x = x
        _axis = axis

    res = xp.linalg.norm(_x, axis=_axis, ord=ord)

    if keepdims:
        shape = list(x.shape)
        axes = cast(
            "tuple[int, ...]",
            normalize_axis_tuple(range(x.ndim) if axis is None else axis, x.ndim),
        )
        for i in axes:
            shape[i] = 1
        res = xp.reshape(res, tuple(shape))

    return res

matrix_rank 用奇异值分解实现,与 xp.linalg.matrix_rank 的 rtol 语义不同(标准要求乘以 max(M,N) * eps)。pinv 同样需要调整 rcond。matrix_norm 对最后两维做 norm。vector_norm 的关键技巧是先把多维数组 reshape 为 1-D,避免 xp.linalg.norm 把 2-D 输入当作矩阵处理。

65.11.4 对角线与迹

diagonaltrace 处理"最后两维 vs前两维"的差异:

源码路径:sklearn/externals/array_api_compat/common/_linalg.py - diagonal/trace/svdvals

def diagonal(x: Array, /, xp: Namespace, *, offset: int = 0, **kwargs: object) -> Array:
    return xp.diagonal(x, offset=offset, axis1=-2, axis2=-1, **kwargs)

def trace(
    x: Array, /, xp: Namespace, *, offset: int = 0, dtype: DType | None = None,
    **kwargs: object,
) -> Array:
    return xp.asarray(
        xp.trace(x, offset=offset, dtype=dtype, axis1=-2, axis2=-1, **kwargs)
    )

def svdvals(x: Array, /, xp: Namespace) -> Array | tuple[Array, ...]:
    return xp.linalg.svd(x, compute_uv=False)

xp.diagonalxp.trace 默认作用于前两维,Array API 要求作用于最后两维,所以显式指定 axis1=-2, axis2=-1svdvals 在 NumPy 中通过 svd(compute_uv=False) 实现。

65.11.5 FFT 精度保持

所有 FFT 函数都强制 float32 → complex64float64 → complex128 精度保持:

源码路径:sklearn/externals/array_api_compat/common/_fft.py - fft()

def fft(x: Array, /, xp: Namespace, *, n: int | None = None, axis: int = -1,
        norm: _Norm = "backward") -> Array:
    res = xp.fft.fft(x, n=n, axis=axis, norm=norm)
    if x.dtype in [xp.float32, xp.complex64]:
        return res.astype(xp.complex64)
    return res

def ifft(x, /, xp, *, n=None, axis=-1, norm="backward"):
    res = xp.fft.ifft(x, n=n, axis=axis, norm=norm)
    if x.dtype in [xp.float32, xp.complex64]:
        return res.astype(xp.complex64)
    return res

def fftn(x, /, xp, *, s=None, axes=None, norm="backward"):
    res = xp.fft.fftn(x, s=s, axes=axes, norm=norm)
    if x.dtype in [xp.float32, xp.complex64]:
        return res.astype(xp.complex64)
    return res

def ifftn(x, /, xp, *, s=None, axes=None, norm="backward"):
    res = xp.fft.ifftn(x, s=s, axes=axes, norm=norm)
    if x.dtype in [xp.float32, xp.complex64]:
        return res.astype(xp.complex64)
    return res

def rfft(x, /, xp, *, n=None, axis=-1, norm="backward"):
    res = xp.fft.rfft(x, n=n, axis=axis, norm=norm)
    if x.dtype == xp.float32:
        return res.astype(xp.complex64)
    return res

def irfft(x, /, xp, *, n=None, axis=-1, norm="backward"):
    res = xp.fft.irfft(x, n=n, axis=axis, norm=norm)
    if x.dtype == xp.complex64:
        return res.astype(xp.float32)
    return res

核心问题是 NumPy 的 np.fft.fft 会错误地将 float32 输入上转为 complex128 输出,违反 Array API 的精度保持规则。兼容层在输出后用 .astype(xp.complex64) 强制转回。hfft/ihfft 还需要 float32 → float32 精度保持,比 FFT 更复杂。fftfreq/rfftfreq 不支持 GPU 设备。

源码路径:sklearn/externals/array_api_compat/common/_fft.py - fftfreq() / fftshift()

def fftfreq(n: int, /, xp: Namespace, *, d: float = 1.0,
            dtype: DType | None = None, device: Device | None = None) -> Array:
    if device not in ["cpu", None]:
        raise ValueError(f"Unsupported device {device!r}")
    res = xp.fft.fftfreq(n, d=d)
    if dtype is not None:
        return res.astype(dtype)
    return res

def fftshift(x: Array, /, xp: Namespace, *, axes: int | Sequence[int] | None = None) -> Array:
    return xp.fft.fftshift(x, axes=axes)

def ifftshift(x: Array, /, xp: Namespace, *, axes: int | Sequence[int] | None = None) -> Array:
    return xp.fft.ifftshift(x, axes=axes)

fftfreq 不支持 GPU 设备(即使 CuPy/PyTorch 支持 GPU,但 Array API 标准尚未定义此行为)。fftshift/ifftshift 是简单的包装——NumPy 原生支持良好。

65.12 后端专属适配器

下面用 Mermaid 流程图展示各后端适配器的统一入口模式:

flowchart TB subgraph "后端适配器统一模式" Init["__init__.py:<br/>clone_module() 批量导入"] Aliases["_aliases.py:<br/>get_xp(backend)(通用函数)"] Backend["后端特定实现<br/>(如 _fix_promotion)"] Linalg["linalg.py:<br/>linalg 命名空间包装"] FFT["fft.py:<br/>axes→dim 重命名"] Info["_info.py:<br/>__array_namespace_info__"] Typing["_typing.py:<br/>后端专属类型"] end Init --> Aliases Aliases --> Backend Aliases --> Linalg Aliases --> FFT Init --> Info Init --> Typing

65.12.1 NumPy 后端:copy 语义映射

NumPy 1.x 没有 asarraycopy 参数,2.0 引入了 _CopyMode 枚举:

源码路径:sklearn/externals/array_api_compat/numpy/_aliases.py - asarray/astype/count_nonzero/take_along_axis/vecdot/isdtype/unstack/ceil/floor/trunc

def asarray(
    obj, /, *, dtype=None, device=None, copy=None, **kwargs,
):
    _helpers._check_device(np, device)

    # None 在 NumPy 1.0 不支持,但可以用内部枚举
    # False 在 NumPy 1.0 表示 None (NumPy 2.0 与 Array API 一致)
    if copy is None:
        copy = np._CopyMode.IF_NEEDED
    elif copy is False:
        copy = np._CopyMode.NEVER

    return np.array(obj, copy=copy, dtype=dtype, **kwargs)

def astype(
    x, dtype, /, *, copy=True, device=None,
):
    _helpers._check_device(np, device)
    return x.astype(dtype=dtype, copy=copy)

# 第 65 章 —— count_nonzero 在 axis=None 且 keepdims=False 时返回 Python int
# 第 65 章 —— https://github.com/numpy/numpy/issues/17562
def count_nonzero(x, axis=None, keepdims=False):
    result = cast("Any", np.count_nonzero(x, axis=axis, keepdims=keepdims))
    if axis is None and not keepdims:
        return np.asarray(result)  # 包装为 ndarray 保持类型一致
    return result

# 第 65 章 —— take_along_axis: axis 在 NumPy 中是必需参数
def take_along_axis(x, indices, /, *, axis=-1):
    return np.take_along_axis(x, indices, axis=axis)

# 第 65 章 —— ceil, floor, trunc: NumPy 1.x 对整数返回整数而非浮点
def ceil(x, /):
    if np.__version__ < '2' and np.issubdtype(x.dtype, np.integer):
        return x.copy()    # 整数直接返回副本
    return np.ceil(x)

def floor(x, /):
    if np.__version__ < '2' and np.issubdtype(x.dtype, np.integer):
        return x.copy()
    return np.floor(x)

def trunc(x, /):
    if np.__version__ < '2' and np.issubdtype(x.dtype, np.integer):
        return x.copy()
    return np.trunc(x)

# 第 65 章 —— 如果原生已有 vecdoot/isdtype/unstack,则用原生版本(NumPy 2.0+)
if hasattr(np, "vecdot"):
    vecdot = np.vecdot
else:
    vecdot = get_xp(np)(_aliases.vecdot)

if hasattr(np, "isdtype"):
    isdtype = np.isdtype
else:
    isdtype = get_xp(np)(_aliases.isdtype)

if hasattr(np, "unstack"):
    unstack = np.unstack
else:
    unstack = get_xp(np)(_aliases.unstack)

这段代码展示了 NumPy 后端适配器。asarray 处理 copy 参数的 None/False 到 NumPy 1.x/2.0 内部枚举的映射。astype 添加了 device 参数的检查。count_nonzero 修复了 axis=None and not keepdims 时返回 Python int 而非 ndarray 的问题。ceil/floor/trunc 修复了 NumPy 1.x 对整数类型返回整数的问题。vecdot/isdtype/unstack 则优先用 NumPy 2.0+ 的原生版本。

65.12.2 NumPy linalg 与 fft 模块

NumPy 后端的 linalg 子模块沿用统一的 clone_module + 包装模式:

源码路径:sklearn/externals/array_api_compat/numpy/linalg.py

from .._internal import clone_module, get_xp
from ..common import _linalg

__all__ = clone_module("numpy.linalg", globals())

from ._aliases import matmul, matrix_transpose, tensordot, vecdot
from ..common._linalg import EighResult, QRResult, SlogdetResult, SVDResult
from ..common._linalg import cross, outer, eigh, qr, slogdet, svd, cholesky
from ..common._linalg import matrix_rank, pinv, matrix_norm, svdvals, diagonal, trace

# 第 65 章 —— NumPy 特有的 solve:仅当 x2 为 1-D 时调用 solve1
def solve(x1: Array, x2: Array, /) -> Array:
    # ... 复制 _linalg.solve 与 np.linalg.solve 的逻辑
    # 在 x2.ndim == 1 时调用 solve1,其他情况调用 solve

if hasattr(np.linalg, "vector_norm"):
    vector_norm = np.linalg.vector_norm
else:
    vector_norm = get_xp(np)(_linalg.vector_norm)

NumPy linalg 的 solve 是特别处理的——NumPy 1.x 的 np.linalg.solve 对 1-D x2 有歧义(与 Array API 标准不同),所以这里重新实现了 solve 函数。

源码路径:sklearn/externals/array_api_compat/numpy/fft.py

import numpy as np
from .._internal import clone_module, get_xp
from ..common import _fft

__all__ = clone_module("numpy.fft", globals())

# 第 65 章 —— 所有 FFT 函数都需包装为精度保持版
fft = get_xp(np)(_fft.fft)
ifft = get_xp(np)(_fft.ifft)
fftn = get_xp(np)(_fft.fftn)
ifftn = get_xp(np)(_fft.ifftn)
rfft = get_xp(np)(_fft.rfft)
irfft = get_xp(np)(_fft.irfft)
rfftn = get_xp(np)(_fft.rfftn)
irfftn = get_xp(np)(_fft.irfftn)
hfft = get_xp(np)(_fft.hfft)
ihfft = get_xp(np)(_fft.ihfft)
fftfreq = get_xp(np)(_fft.fftfreq)
rfftfreq = get_xp(np)(_fft.rfftfreq)
fftshift = get_xp(np)(_fft.fftshift)
ifftshift = get_xp(np)(_fft.ifftshift)

NumPy fft 子模块完全依赖 common/_fft.py 提供的精度保持包装,每个函数都用 get_xp(np)(_fft.xxx) 绑定 xp 参数。

65.12.3 NumPy _info 模块

NumPy 的元信息查询模块实现完整的 __array_namespace_info__

源码路径:sklearn/externals/array_api_compat/numpy/_info.py - __array_namespace_info__.capabilities()

class __array_namespace_info__:
    __module__ = 'numpy'

    def capabilities(self):
        return {
            "boolean indexing": True,
            "data-dependent shapes": True,
            "max dimensions": 64,
        }

    def default_device(self):
        return "cpu"

    def default_dtypes(self, *, device: Device | None = None) -> DefaultDTypes:
        if device not in ["cpu", None]:
            raise ValueError('Device not understood. Only "cpu" is allowed, but received:'
                            f' {device}')
        return {
            "real floating": dtype(float64),
            "complex floating": dtype(complex128),
            "integral": dtype(intp),
            "indexing": dtype(intp),
        }

    def dtypes(self, *, device: Device | None = None,
               kind: str | tuple[str, ...] | None = None) -> dict[str, DType]:
        if device not in ["cpu", None]:
            raise ValueError('Device not understood. Only "cpu" is allowed, but received:'
                            f' {device}')
        if kind is None:
            return {
                "bool": dtype(bool),
                "int8": dtype(int8),
                "int16": dtype(int16),
                "int32": dtype(int32),
                "int64": dtype(int64),
                "uint8": dtype(uint8),
                "uint16": dtype(uint16),
                "uint32": dtype(uint32),
                "uint64": dtype(uint64),
                "float32": dtype(float32),
                "float64": dtype(float64),
                "complex64": dtype(complex64),
                "complex128": dtype(complex128),
            }
        # ... 其他 kind 处理
        if isinstance(kind, tuple):
            res: dict[str, DType] = {}
            for k in kind:
                res.update(self.dtypes(kind=k))
            return res
        raise ValueError(f"unsupported kind: {kind!r}")

    def devices(self) -> list[Device]:
        return ["cpu"]

default_dtypesdtypes 都严格检查 device in ["cpu", None],因为 NumPy 仅支持 CPU。dtypes 方法支持 7 种 kind("bool"/"signed integer"/"unsigned integer"/"integral"/"real floating"/"complex floating"/"numeric"),通过 elif 链依次匹配。devices() 永远返回 ["cpu"]

65.12.4 NumPy _typing 模块

NumPy 后端的类型模块定义了 ArrayDevice 类型别名:

源码路径:sklearn/externals/array_api_compat/numpy/_typing.py

from __future__ import annotations
from typing import TYPE_CHECKING, Any, Literal, TypeAlias
import numpy as np

Device: TypeAlias = Literal["cpu"]

if TYPE_CHECKING:
    DType: TypeAlias = np.dtype[
        np.bool_
        | np.integer[Any]
        | np.float32
        | np.float64
        | np.complex64
        | np.complex128
    ]
    Array: TypeAlias = np.ndarray[Any, DType]
else:
    DType: TypeAlias = np.dtype
    Array: TypeAlias = np.ndarray

Device 是字面量 "cpu"(NumPy 唯一的设备)。Array/DType 在类型检查时是精确的 NumPy 类型,运行时退化为宽泛的 np.ndarray/np.dtype(避免运行时检查开销)。

65.12.5 CuPy 后端:设备上下文管理

CuPy 适配器重点处理设备上下文:

源码路径:sklearn/externals/array_api_compat/cupy/_aliases.py - asarray/astype/count_nonzero/take_along_axis/ceil/floor/trunc

def asarray(obj, /, *, dtype=None, device=None, copy=None, **kwargs):
    # 设备上下文:确保后续操作在指定 GPU 上执行
    with cp.cuda.Device(device):
        if copy is None:
            return cp.asarray(obj, dtype=dtype, **kwargs)
        else:
            res = cp.array(obj, dtype=dtype, copy=copy, **kwargs)
            # copy=False 时检查是否真的避免了复制
            if not copy and res is not obj:
                raise ValueError("Unable to avoid copy while creating an array as requested")
            return res

def astype(x, dtype, /, *, copy=True, device=None):
    if device is None:
        return x.astype(dtype=dtype, copy=copy)
    # 先转换 dtype(不复制),再移动到目标设备
    out = _helpers.to_device(x.astype(dtype=dtype, copy=False), device)
    return out.copy() if copy and out is x else out

# 第 65 章 —— cupy.count_nonzero 没有 keepdims
def count_nonzero(x, axis=None, keepdims=False):
    result = cp.count_nonzero(x, axis)
    if keepdims:
        if axis is None:
            return cp.reshape(result, [1]*x.ndim)
        return cp.expand_dims(result, axis)
    return result

# 第 65 章 —— take_along_axis 默认 axis=-1
def take_along_axis(x, indices, /, *, axis=-1):
    return cp.take_along_axis(x, indices, axis=axis)

# 第 65 章 —— ceil/floor/trunc 对整数返回副本
def ceil(x, /):
    if cp.issubdtype(x.dtype, cp.integer):
        return x.copy()
    return cp.ceil(x)

def floor(x, /):
    if cp.issubdtype(x.dtype, cp.integer):
        return x.copy()
    return cp.floor(x)

def trunc(x, /):
    if cp.issubdtype(x.dtype, cp.integer):
        return x.copy()
    return cp.trunc(x)

这段代码展示了 CuPy 适配器的核心特点。asarraywith cp.cuda.Device(device): 上下文管理器确保在指定 GPU 上执行。astype 支持跨设备 dtype 转换。count_nonzero 手动补全 keepdims 功能。ceil/floor/trunc 处理整数类型的特殊情况。

65.12.6 CuPy linalg 与 fft 模块

CuPy linalg 沿用 clone_module + 包装模式,但因 cupy.linalg 不有 __all__,需要手动提取符号列表:

源码路径:sklearn/externals/array_api_compat/cupy/linalg.py

from cupy.linalg import *
_n: dict[str, object] = {}
exec('from cupy.linalg import *', _n)
del _n['__builtins__']
linalg_all = list(_n)
del _n

from ..common import _linalg
from .._internal import get_xp

import cupy as cp

# 第 65 章 —— 这些函数在主命名空间与 linalg 命名空间中都存在
from ._aliases import matmul, matrix_transpose, tensordot, vecdot

cross = get_xp(cp)(_linalg.cross)
outer = get_xp(cp)(_linalg.outer)
EighResult = _linalg.EighResult
# 第 65 章 —— ... 其余与 NumPy linalg 类似

if hasattr(cp.l, 'vector_norm'):
    vector_norm = cp.l
else:
    vector_norm = get_xp(cp)(_linalg.vector_norm)

CuPy linalg 包装逻辑与 NumPy linalg 完全一致,只是后端换成 cp

源码路径:sklearn/externals/array_api_compat/cupy/fft.py

from cupy.fft import *
# 第 65 章 —— cupy.fft 不有 __all__,用 exec 提取
_n: dict[str, object] = {}
exec("from cupy.fft import *", _n)
del _n["__builtins__"]
fft_all = list(_n)
del _n

from ..common import _fft
from .._internal import get_xp
import cupy as cp

fft = get_xp(cp)(_fft.fft)
# 第 65 章 —— ... 其余与 NumPy fft 类似

CuPy fft 的逻辑与 NumPy fft 一致。

65.12.7 CuPy _info 模块

CuPy 的元信息模块与 NumPy 类似,但 default_device 返回 cuda.Device(0)devices 返回实际可用的 GPU 列表:

源码路径:sklearn/externals/array_api_compat/cupy/_info.py - __array_namespace_info__.default_device()

class __array_namespace_info__:
    __module__ = 'cupy'

    def capabilities(self):
        return {
            "boolean indexing": True,
            "data-dependent shapes": True,
            "max dimensions": 64,
        }

    def default_device(self):
        return cuda.Device(0)

    def default_dtypes(self, *, device=None):
        return {
            "real floating": dtype(float64),
            "complex floating": dtype(complex128),
            "integral": dtype(intp),
            "indexing": dtype(intp),
        }

    # dtypes 与 NumPy 类似,但不去验证 device

    def devices(self):
        return [cuda.Device(i) for i in range(cuda.runtime.getDeviceCount())]

CuPy 的 default_device() 返回 cuda.Device(0)(初始化时的默认 GPU)。devices() 返回当前所有可见 GPU 的列表。与 NumPy 不同,CuPy 的 dtypes 不验证 device(因为 CuPy 可以创建任意 GPU 设备的数组)。

65.12.8 CuPy _typing 模块

源码路径:sklearn/externals/array_api_compat/cupy/_typing.py(结构与 NumPy _typing 类似,但 Device 是 cuda.Device

from __future__ import annotations
from typing import TYPE_CHECKING, Any, Literal, TypeAlias
import cupy as cp

Device: TypeAlias = cp.cuda.Device
# 第 65 章 —— Array 与 DType 类似 NumPy

CuPy 的 Device 类型是 cp.cuda.Device(具体的设备对象)。

65.12.9 PyTorch 后端:_fix_promotion 类型提升修正

PyTorch 0-D 张量的类型提升不符合 Array API 标准:

源码路径:sklearn/externals/array_api_compat/torch/_aliases.py - _promotion_table/_fix_promotion/result_type/can_cast

_int_dtypes = {
    torch.uint8, torch.int8, torch.int16, torch.int32, torch.int64,
}
try:
    # torch >=2.3
    _int_dtypes |= {torch.uint16, torch.uint32, torch.uint64}
except AttributeError:
    pass

_array_api_dtypes = {
    torch.bool, *_int_dtypes,
    torch.float32, torch.float64,
    torch.complex64, torch.complex128,
}

# 第 65 章 —— 下面表格定义了 PyTorch 类型提升规则,与 Array API 略有差异
_promotion_table = {
    # ints
    (torch.int8, torch.int16): torch.int16,
    (torch.int8, torch.int32): torch.int32,
    (torch.int8, torch.int64): torch.int64,
    (torch.int16, torch.int32): torch.int32,
    (torch.int16, torch.int64): torch.int64,
    (torch.int32, torch.int64): torch.int64,
    # ints and uints (mixed sign)
    (torch.uint8, torch.int8): torch.int16,
    (torch.uint8, torch.int16): torch.int16,
    (torch.uint8, torch.int32): torch.int32,
    (torch.uint8, torch.int64): torch.int64,
    # floats
    (torch.float32, torch.float64): torch.float64,
    # complexes
    (torch.complex64, torch.complex128): torch.complex128,
    # Mixed float and complex
    (torch.float32, torch.complex64): torch.complex64,
    (torch.float32, torch.complex128): torch.complex128,
    (torch.float64, torch.complex64): torch.complex128,
    (torch.float64, torch.complex128): torch.complex128,
}
_promotion_table.update({(b, a): c for (a, b), c in _promotion_table.items()})
_promotion_table.update({(a, a): a for a in _array_api_dtypes})


def _fix_promotion(x1, x2, only_scalar=True):
    """修正 PyTorch 0-D 张量的类型提升行为,使其符合 Array API 标准。"""
    if not isinstance(x1, torch.Tensor) or not isinstance(x2, torch.Tensor):
        return x1, x2
    if x1.dtype not in _array_api_dtypes or x2.dtype not in _array_api_dtypes:
        return x1, x2
    # 如果参数是 0-D,PyTorch 会下会下转另一个参数
    if not only_scalar or x1.shape == ():
        dtype = result_type(x1, x2)
        x2 = x2.to(dtype)
    if not only_scalar or x2.shape == ():
        dtype = result_type(x1, x2)
        x1 = x1.to(dtype)
    return x1, x2


def result_type(*arrays_and_dtypes):
    """计算多个数组/dtype 的结果 dtype。"""
    num = len(arrays_and_dtypes)
    if num == 0:
        raise ValueError("At least one array or dtype must be provided")
    elif num == 1:
        x = arrays_and_dtypes[0]
        if isinstance(x, torch.dtype):
            return x
        return x.dtype
    if num == 2:
        x, y = arrays_and_dtypes
        return _result_type(x, y)
    else:
        # 多个参数:标量放最后
        scalars, others = [], []
        for x in arrays_and_dtypes:
            if isinstance(x, _py_scalars):
                scalars.append(x)
            else:
                others.append(x)
        return _reduce(_result_type, others + scalars)


def _result_type(x, y):
    if not (isinstance(x, _py_scalars) or isinstance(y, _py_scalars)):
        xdt = x if isinstance(x, torch.dtype) else x.dtype
        ydt = y if isinstance(y, torch.dtype) else y.dtype
        try:
            return _promotion_table[xdt, ydt]
        except KeyError:
            pass
    x = torch.tensor([], dtype=x) if isinstance(x, torch.dtype) else x
    y = torch.tensor([], dtype=y) if isinstance(y, torch.dtype) else y
    return torch.result_type(x, y)


def can_cast(from_, to, /):
    if not isinstance(from_, torch.dtype):
        from_ = from_.dtype
    return torch.can_cast(from_, to)

这段代码定义了 PyTorch 类型提升修正表。_promotion_table 明确定义 int8 + int16 → int16、float32 + complex64 → complex64 等规则,补充了 PyTorch 默认行为与 Array API 不一致的地方。update 调用添加了反向键(如 (int16, int8) → int16)和恒等键(如 (int, int) → int)。

_fix_promotion 的核心问题是:当一个参数是 0-D 张量时,PyTorch 会下转另一个参数(即 (0-D int8, int16) 会被强制转为 (0-D int8, int8))。这与 Array API 标准相反。修正方法是用 result_type 计算正确的提升 dtype,然后 .to(dtype) 显式转换。only_scalar=True 时仅修正 0-D 情况,避免改变高维张量的 dtype。

result_type 是 Array API 的类型提升查询接口。_result_type 内部实现先用 _promotion_table 查表,失败则用 torch.result_type 兜底。

PyTorch 后端适配器还有其他大量函数修正:

源码路径:sklearn/externals/array_api_compat/torch/_aliases.py - max/min/sort/argsort/sum/prod/any/all/mean/std/var/concat/squeeze/flip/roll/diff/count_nonzero/where/reshape/arange/eye/linspace/full/ones/zeros/empty/tril/triu/expand_dims/astype/broadcast_arrays/unique_all/unique_counts/unique_inverse/unique_values/matmul/vecdot/tensordot/isdtype/take/take_along_signal/sign/meshgrid/repeat

# 第 65 章 —— torch.min/max 返回元组,不支持多 axis
def max(x, /, *, axis=None, keepdims=False):
    if axis == ():
        return torch.clone(x)
    return torch.amax(x, axis, keepdims=keepdims)

def min(x, /, *, axis=None, keepdims=False):
    if axis == ():
        return torch.clone(x)
    return torch.amin(x, axis, keepdims=keepdims)

# 第 65 章 —— torch.sort 返回元组
def sort(x, /, *, axis=-1, descending=False, stable=True, **kwargs):
    return torch.sort(x, dim=axis, descending=descending, stable=stable, **kwargs).values

def argsort(x, /, *, axis=-1, descending=False, stable=True, **kwargs):
    return torch.argsort(x, dim=axis, descending=descending, stable=stable, **kwargs)

# 第 65 章 —— torch.sum/prod 不支持多 axis 和 keepdim+ axis=None
def sum(x, /, *, axis=None, dtype, keepdims=False, **kwargs):
    if axis == ():
        return _sum_prod_no_axis(x, dtype)
    if axis is None:
        res = torch.sum(x, dtype=dtype, **kwargs)
        return _axis_none_keepdims(res, x.ndim, keepdims)
    return torch.sum(x, axis, dtype=dtype, keepdims=keepdims, **kwargs)

# 第 65 章 —— torch.any/all 不支持多 axis、uint8 不返回 bool
def any(x, /, *, axis=None, keepdims=False, **kwargs):
    if axis == ():
        return x.to(torch.bool)
    if isinstance(axis, tuple):
        res = _reduce_multiple_axes(torch.any, x, axis, keepdims=keepdims, **kwargs)
        return res.to(torch.bool)
    if axis is None:
        res = torch.any(x, **kwargs)
        res = _axis_none_keepdims(res, x.ndim, keepdims)
        return res.to(torch.bool)
    return torch.any(x, axis, keepdims=keepdims).to(torch.bool)

# 第 65 章 —— torch.flip/roll 接受 dim 而非 axis
def flip(x, /, *, axis=None, **kwargs):
    if axis is None:
        axis = tuple(range(x.ndim))
    return x.flip(axis, **kwargs)

# 第 65 章 —— torch.diff 用 dim 而非 axis
def diff(x, /, *, axis=-1, n=1, prepend=None, append=None):
    return torch.diff(x, dim=axis, n=n, prepend=prepend, append=append)

# 第 65 章 —— torch.count_nonzero 不支持 keepdims
def count_nonzero(x, /, *, axis=None, keepdims=False):
    result = torch.count_nonzero(x, dim=axis)
    if keepdims:
        if isinstance(axis, int):
            return result.unsqueeze(axis)
        elif isinstance(axis, tuple):
            n_axis = [x.ndim + ax if ax < 0 else ax for ax in axis]
            sh = [1 if i in n_axis else x.shape[i] for i in range(x.ndim)]
            return torch.reshape(result, sh)
        return _axis_none_keepdims(result, x.ndim, keepdims)
    return result

# 第 65 章 —— torch.where 需要类型提升修正
def where(condition, x1, x2, /):
    x1, x2 = _fix_promotion(x1, x2)
    return torch.where(condition, x1, x2)

# 第 65 章 —— torch.reshape 不支持 copy 参数
def reshape(x, /, shape, *, copy=None, **kwargs):
    if copy is not None:
        raise NotImplementedError("torch.reshape doesn't yet support the copy keyword")
    return torch.reshape(x, shape, **kwargs)

# 第 65 章 —— torch.arange 不支持空数组
def arange(start, /, stop=None, step=1, *, dtype=None, device=None, **kwargs):
    if stop is None:
        start, stop = 0, start
    if step > 0 and stop <= start or step < 0 and stop >= start:
        # 空数组:手动创建
        if dtype is None:
            if _builtin_all(isinstance(i, int) for i in [start, stop, step]):
                dtype = torch.int64
            else:
                dtype = torch.float32
        return torch.empty(0, dtype=dtype, device=device, **kwargs)
    return torch.arange(start, stop, step, dtype=dtype, device=device, **kwargs)

# 第 65 章 —— torch.eye 不支持 k 偏移和默认 N=M
def eye(n_rows, n_cols=None, /, *, k=0, dtype=None, device=None, **kwargs):
    if n_cols is None:
        n_cols = n_rows
    z = torch.zeros(n_rows, n_cols, dtype=dtype, device=device, **kwargs)
    if abs(k) <= n_rows + n_cols:
        z.diagonal(k).fill_(1)
    return z

# 第 65 章 —— torch.linspace 不支持 endpoint=False
def linspace(start, stop, /, num, *, dtype=None, device=None, endpoint=True, **kwargs):
    if not endpoint:
        return torch.linspace(start, stop, num+1, dtype=dtype, device=device, **kwargs)[:-1]
    return torch.linspace(start, stop, num, dtype=dtype, device=device, **kwargs)

# 第 65 章 —— torch.full 不支持 int size
def full(shape, fill_value, *, dtype=None, device=None, **kwargs):
    if isinstance(shape, int):
        shape = (shape,)
    return torch.full(shape, fill_value, dtype=dtype, device=device, **kwargs)

# 第 65 章 —— unique_all 不支持(PyTorch 缺少 indices 返回)
def unique_all(x):
    raise NotImplementedError("unique_all() not yet implemented for pytorch (see https://github.com/pytorch/pytorch/issues/36748)")

def unique_counts(x):
    values, counts = torch.unique(x, return_counts=True)
    counts[torch.isnan(values)] = 1  # 修复 NaN 计数为 0 的 bug
    return UniqueCountsResult(values, counts)

# 第 65 章 —— matmul 需要类型提升修正
def matmul(x1, x2, /, **kwargs):
    x1, x2 = _fix_promotion(x1, x2, only_scalar=False)
    return torch.matmul(x1, x2, **kwargs)

# 第 65 章 —— vecdot 需要类型提升修正
def vecdot(x1, x2, /, *, axis=-1):
    x1, x2 = _fix_promotion(x1, x2, only_scalar=False)
    return _vecdot(x1, x2, axis=axis)

# 第 65 章 —— take/take_along_axis 需要负索引
def take(x, indices, /, *, axis=None, **kwargs):
    if axis is None:
        if x.ndim != 1:
            raise ValueError("axis must be specified when ndim > 1")
        axis = 0
    return torch.index_select(
        x, axis,
        torch.where(indices < 0, indices + x.shape[axis], indices),
        **kwargs
    )

这段代码展示了 PyTorch 后端适配器的核心特点:

  • 类型提升修正_fix_promotion 处理 0-D 张量特殊情况

  • axis/axis=() 边界:PyTorch 对 axis=()axis=None + keepdimsaxis` 为 tuple 等情况支持不完善

  • 参数名差异dim vs axisdims vs axessize vs shape

  • 缺失功能补全take/take_along_axis 负索引修正、expand_dims 替代品、unique_* 系列的 NaN 计数修复

  • 特殊行为修复uint8any/all 后转 boolsign 的 NaN 传播与复数支持

65.12.10 PyTorch linalg 模块

源码路径:sklearn/externals/array_api_compat/torch/linalg.py

from __future__ import annotations

import torch
import torch.linalg

from .._internal import clone_module

__all__ = clone_module("torch.linalg", globals())

# 第 65 章 —— outer 在 torch 中但不在 linalg 命名空间
from torch import outer
from ._aliases import _fix_promotion, sum
from ._aliases import matmul, matrix_transpose, tensordot
from ._typing import Array, DType
from ..common._typing import JustInt, JustFloat

# 第 65 章 —— torch.linalg.cross 默认 axis 是首个 size=3 的轴,不默认 axis=-1
def cross(x1: Array, x2: Array, /, *, axis: int = -1) -> Array:
    x1, x2 = _fix_promotion(x1, x2, only_scalar=False)
    if not (-min(x1.ndim, x2.ndim) <= axis < max(x1.ndim, x2.ndim)):
        raise ValueError(f"axis {axis} out of bounds for cross product of arrays with shapes {x1.shape} and {x2.shape}")
    if not (x1.shape[axis] == x2.shape[axis] == 3):
        raise ValueError(f"cross product axis must have size 3, got {x1.shape[axis]} and {x2.shape[axis]}")
    x1, x2 = torch.broadcast_tensors(x1, x2)
    return torch.linalg.cross(x1, x2, dim=axis)

# 第 65 章 —— torch.linalg.vecdot 不支持整数 dtype
def vecdot(x1: Array, x2: Array, /, *, axis: int = -1, **kwargs: object) -> Array:
    from ._aliases import isdtype
    x1, x2 = _fix_promotion(x1, x2, only_scalar=False)
    if x1.shape[axis] != x2.shape[axis]:
        raise ValueError("x1 and x2 must have the same size along the given axis")
    if isdtype(x1.dtype, 'integral') or isdtype(x2.dtype, 'integral'):
        if kwargs:
            raise RuntimeError("vecdot kwargs not supported for integral dtypes")
        x1_ = torch.moveaxis(x1, axis, -1)
        x2_ = torch.moveaxis(x2, axis, -1)
        x1_, x2_ = torch.broadcast_tensors(x1_, x2_)
        res = x1_[..., None, :] @ x2_[..., None]
        return res[..., 0, 0]
    return torch.linalg.vecdot(x1, x2, dim=axis, **kwargs)

# 第 65 章 —— torch.linalg.solve 在 x1.ndim - 1 == x2.ndim 且 shape 匹配时
# 第 65 章 —— 会把 x2 当作 batched 1-D solve
def solve(x1: Array, x2: Array, /, **kwargs: object) -> Array:
    x1, x2 = _fix_promotion(x1, x2, only_scalar=False)
    if x2.ndim != 1 and x1.ndim - 1 == x2.ndim and x1.shape[:-1] == x2.shape:
        x2 = x2[None]
    return torch.linalg.solve(x1, x2, **kwargs)

# 第 65 章 —— torch.trace 不支持 offset
def trace(x: Array, /, *, offset: int = 0, dtype: DType | None = None) -> Array:
    return sum(torch.diagonal(x, offset=offset, dim1=-2, dim2=-1), axis=-1, dtype=dtype)

# 第 65 章 —— torch.vector_norm 错误地将 axis=() 当作 axis=None
def vector_norm(x, /, *, axis=None, keepdims=False, ord=2, **kwargs):
    if axis == ():
        out = kwargs.get('out')
        if out is None:
            dtype = None
            if x.dtype == torch.complex64:
                dtype = torch.float32
            elif x.dtype == torch.complex128:
                dtype = torch.float64
            out = torch.zeros_like(x, dtype=dtype)
        if ord == 0:
            out[:] = (x != 0)
        else:
            out[:] = torch.abs(x)
        return out
    return torch.linalg.vector_norm(x, ord=ord, axis=axis, keepdim=keepdims, **kwargs)

这段代码展示了 PyTorch 线性代数适配器。cross 修复了 axis 默认值和不支持广播的问题。vecdot 修复了不支持整数 dtype 的问题(手动用 matmul 实现)。solve 修复了 1-D RHS 的歧义处理。trace 手动用 sum + diagonal 实现以支持 offsetvector_norm 修复了 axis=() 被错误处理为 axis=None 的问题。

65.12.11 PyTorch fft 模块

源码路径:sklearn/externals/array_api_compat/torch/fft.py

import torch
import torch.fft
from ._typing import Array
from .._internal import clone_module

__all__ = clone_module("torch.fft", globals())

# 第 65 章 —— torch 的 fftn/ifftn 等用 dim 而非 axes
def fftn(x, /, *, s=None, axes=None, norm="backward", **kwargs):
    return torch.fft.fftn(x, s=s, dim=axes, norm=norm, **kwargs)

def ifftn(x, /, *, s=None, axes=None, norm="backward", **kwargs):
    return torch.fft.ifftn(x, s=s, dim=axes, norm=norm, **kwargs)

def rfftn(x, /, *, s=None, axes=None, norm="backward", **kwargs):
    return torch.fft.rfftn(x, s=s, dim=axes, norm=norm, **kwargs)

def irfftn(x, /, *, s=None, axes=None, norm="backward", **kwargs):
    return torch.fft.irfftn(x, s=s, dim=axes, norm=norm, **kwargs)

def fftshift(x, /, *, axes=None, **kwargs):
    return torch.fft.fftshift(x, dim=axes, **kwargs)

def ifftshift(x, /, *, axes=None, **kwargs):
    return torch.fft.ifftshift(x, dim=axes, **kwargs)

这段代码展示了 PyTorch FFT 适配器。核心问题是 torch.fft.* 的多维函数用 dim 而非 axes。只需简单的参数重命名即可。

65.12.12 PyTorch _info 模块

PyTorch 的元信息模块重点是 devices() 的实现——需要从错误信息中提取设备名:

源码路径:sklearn/externals/array_api_compat/torch/_info.py - __array_namespace_info__.devices()

class __array_namespace_info__:
    __module__ = 'torch'

    def capabilities(self):
        return {
            "boolean indexing": True,
            "data-dependent shapes": True,
            "max dimensions": 64,
        }

    def default_device(self):
        return torch.device("cpu")

    def default_dtypes(self, *, device=None):
        default_floating = torch.get_default_dtype()
        default_complex = torch.complex64 if default_floating == torch.float32 else torch.complex128
        default_integral = torch.int64
        return {
            "real floating": default_floating,
            "complex floating": default_complex,
            "integral": default_integral,
            "indexing": default_integral,
        }

    def dtypes(self, *, device=None, kind=None):
        # ... 按 kind 返回字典
        # 重要:过滤掉当前 device 不支持的 dtype
        res = self._dtypes(kind)
        for k, v in res.copy().items():
            try:
                torch.empty((0,), dtype=v, device=device)
            except:
                del res[k]
        return res

    @cache
    def devices(self):
        # 从错误信息解析支持的设备名
        try:
            torch.device('notadevice')
            raise AssertionError("unreachable")
        except RuntimeError as e:
            devices_names = e.args[0].split('Expected one of ')[1].split(' device type')[0].split(', ')
        # 逐个检测可用 index
        devices = []
        for device_name in devices_names:
            i = 0
            while True:
                try:
                    a = torch.empty((0,), device=torch.device(device_name, index=i))
                    if a.device in devices:
                        break
                    devices.append(a.device)
                except:
                    break
                i += 1
        return devices

PyTorch 没有提供查询所有可用设备的 API,devices() 通过尝试 torch.device('notadevice') 触发错误,从错误信息中提取设备名(如 cpu, cuda, mps, meta),再逐个测试可用索引。这种"从错误中提取信息"的技巧非常巧妙。default_dtypes 使用 torch.get_default_dtype() 动态获取默认 dtype(不是硬编码 float64)。

65.12.13 PyTorch _typing 模块

PyTorch 的 _typing 模块非常简短,直接从 torch 导入类型:

源码路径:sklearn/externals/array_api_compat/torch/_typing.py

__all__ = ["Array", "Device", "DType"]
from torch import device as Device, dtype as DType, Tensor as Array

Device 类型是 torch.device,Array 类型是 torch.Tensor,DType 类型是 torch.dtype

65.12.14 Dask 后端:惰性约束下的 sort

Dask 是惰性计算框架,sort 沿某轴排序需要先重分块为单分块:

源码路径:sklearn/externals/array_api_compat/dask/array/_aliases.py - astype/arange/clip/sort/argsort/count_nonzero/asarray

# 第 65 章 —— da.astype 不完全支持 copy=True
def astype(x, dtype, /, *, copy=True, device=None):
    _helpers._check_device(da, device)
    if not copy and dtype == x.dtype:
        return x
    x = x.astype(dtype)
    return x.copy() if copy else x

# 第 65 章 —— arange 不支持 stop/step 作为关键字参数
def arange(start, /, stop=None, step=1, *, dtype=None, device=None, **kwargs):
    _helpers._check_device(da, device)
    args: list = [start]
    if stop is not None:
        args.append(stop)
    else:
        args.insert(0, 0)
    args.append(step)
    return da.arange(*args, dtype=dtype, **kwargs)

# 第 65 章 —— dask.array.clip 必须三个参数都提供
def clip(x, /, min=None, max=None):
    def _isscalar(a, /):
        return a is None or isinstance(a, (int, float))
    min_shape = () if _isscalar(min) else min.shape
    max_shape = () if _isscalar(max) else max.shape
    result_shape = np.broadcast_shapes(x.shape, min_shape, max_shape)
    if min is not None:
        min = da.broadcast_to(da.asarray(min), result_shape)
    if max is not None:
        max = da.broadcast_to(da.asarray(max), result_shape)
    if min is None and max is None:
        return da.positive(x)
    if min is None:
        return astype(da.minimum(x, max), x.dtype)
    if max is None:
        return astype(da.maximum(x, min), x.dtype)
    return astype(da.minimum(da.maximum(x, min), max), x.dtype)

# 第 65 章 —— 确保沿 axis 是单分块
def _ensure_single_chunk(x, axis):
    if axis < 0:
        axis += x.ndim
    if x.numblocks[axis] < 2:
        return x, lambda x: x
    x = x.rechunk({i: -1 if i == axis else "auto" for i in range(x.ndim)})
    return x, lambda x: x.rechunk()

# 第 65 章 —— sort:重分块后用 map_blocks
def sort(x, /, *, axis=-1, descending=False, stable=True):
    x, restore = _ensure_single_chunk(x, axis)
    meta_xp = array_namespace(x)
    x = da.map_blocks(
        meta_xp.sort, x, axis=axis, meta=x._meta, dtype=x.dtype,
        descending=descending, stable=stable,
    )
    return restore(x)

# 第 65 章 —— argsort 同样处理
def argsort(x, /, *, axis=-1, descending=False, stable=True):
    x, restore = _ensure_single_chunk(x, axis)
    meta_xp = array_namespace(x)
    dtype = meta_xp.argsort(x._meta).dtype
    meta = meta_xp.astype(x._meta, dtype)
    x = da.map_blocks(
        meta_xp.argsort, x, axis=axis, meta=meta, dtype=dtype,
        descending=descending, stable=stable,
    )
    return restore(x)

# 第 65 章 —— count_nonzero 没有 keepdims
def count_nonzero(x, axis=None, keepdims=False):
    result = da.count_nonzero(x, axis)
    if keepdims:
        if axis is None:
            return da.reshape(result, [1] * x.ndim)
        return da.expand_dims(result, axis)
    return result

# 第 65 章 —— asarray 处理 dask 输入与外部输入的差异
def asarray(obj, /, *, dtype=None, device=None, copy=None, **kwargs):
    _helpers._check_device(da, device)
    if isinstance(obj, da.Array):
        if dtype is not None and dtype != obj.dtype:
            if copy is False:
                raise ValueError("Unable to avoid copy when changing dtype")
            obj = obj.astype(dtype)
        return obj.copy() if copy else obj
    if copy is False:
        raise ValueError("Unable to avoid copy when converting a non-dask object to dask")
    obj = np.array(obj, dtype=dtype, copy=True)
    return da.from_array(obj)

这段代码展示了 Dask 适配器的典型挑战。astype 修复了 copy=True 的语义。arange 修复了关键字参数问题。clip 因为通用掩码法对 dask 不可用(uint64→float64 类型提升),所以改用 da.minimum/da.maximum 组合实现。sort/argsort 必须先重分块为单分块,否则无法跨分块排序。_ensure_single_chunk 返回重分块后的数组和一个 restore 回调(用于在排序后尝试恢复分块)。

Dask 后端的信息接口与 NumPy/CuPy 类似但有 device 的差异:

源码路径:sklearn/externals/array_api_compat/dask/array/_info.py - __array_namespace_info__.devices()

class __array_namespace_info__:
    __module__ = "dask.array"

    def capabilities(self):
        return {
            "boolean indexing": True,
            "data-dependent shapes": True,
            "max dimensions": 64,
        }

    def default_device(self):
        return "cpu"

    def default_dtypes(self, /, *, device=None):
        _check_device(da, device)
        return {
            "real floating": dtype(float64),
            "complex floating": dtype(complex128),
            "integral": dtype(intp),
            "indexing": dtype(intp),
        }

    def dtypes(self, /, *, device=None, kind=None):
        _check_device(da, device)
        # ... 与 NumPy/CuPy 类似的实现
        if isinstance(kind, tuple):
            res: dict[str, DType] = {}
            for k in kind:
                res.update(self.dtypes(kind=k))
            return res
        raise ValueError(f"unsupported kind: {kind!r}")

    def devices(self):
        return ["cpu", _DASK_DEVICE]

Dask 的 __array_namespace_info__ 与 NumPy/CuPy 的实现基本一致,区别在于 devices() 返回 ["cpu", _DASK_DEVICE]DASK_DEVICE 是一个特殊单例对象(_dask_device 类),因为 Dask 数组的真实数据设备不固定(元数据可能是 NumPy 也可能是 CuPy)。

65.12.15 Dask linalg 模块

Dask 的 linalg 模块由于 dask.array.linalg 不支持所有标准函数,需要大量自定义包装:

源码路径:sklearn/externals/array_api_compat/dask/array/linalg.py

from dask.array import matmul, outer, tensordot

__all__ = clone_module("dask.array.linalg", globals())

from ._aliases import matrix_transpose, vecdot

EighResult = _linalg.EighResult
QRResult = _linalg.QRResult
SlogdetResult = _linalg.SlogdetResult
SVDResult = _linalg.SVDResult

# 第 65 章 —— Dask 不支持 mode keyword
def qr(x: Array, mode: Literal["reduced", "complete"] = "reduced", **kwargs) -> QRResult:
    if mode != "reduced":
        raise ValueError("dask arrays only support using mode='reduced'")
    return QRResult(*da.linalg.qr(x, **kwargs))

trace = get_xp(da)(_linalg.trace)
cholesky = get_xp(da)(_linalg.cholesky)
matrix_rank = get_xp(da)(_linalg.matrix_rank)
matrix_norm = get_xp(da)(_linalg.matrix_norm)

# 第 65 章 —— Wrap svd to not pass full_matrices
def svd(x, full_matrices=True, **kwargs):
    if full_matrices:
        raise ValueError("full_matrics=True is not supported by dask.")
    return da.linalg.svd(x, coerce_signs=False, **kwargs)

def svdvals(x):
    # TODO: can't avoid computing U or V for dask
    _, s, _ = svd(x)
    return s

vector_norm = get_xp(da)(_linalg.vector_norm)
diagonal = get_xp(da)(_linalg.diagonal)

Dask linalg 的特点是:仅支持 reduced QR、full_matrices=True 的 SVD 报错、svdvals 无法避免计算 U/V(需要完全 SVD 才能获取奇异值)。

65.13 类型系统与构建工具

65.13.1 跨后端类型契约

common/_typing.py 用 TypedDict 与 Protocol 定义跨后端类型契约:

源码路径:sklearn/externals/array_api_compat/common/_typing.py - Array/Device/DType/Namespace/SupportsArrayNamespace/HasShape/SupportsBufferProtocol/NestedSequence

from __future__ import annotations
from collections.abc import Mapping
from types import ModuleType as Namespace
from typing import TYPE_CHECKING, Literal, Protocol, TypeAlias, TypedDict, TypeVar, final

if TYPE_CHECKING:
    from _typeshed import Incomplete
    # 类型检查时使用 Incomplete 占位
    SupportsBufferProtocol: TypeAlias = Incomplete
    Array: TypeAlias = Incomplete
    Device: TypeAlias = Incomplete
    DType: TypeAlias = Incomplete
else:
    # 运行时用 object 替代(避免循环导入)
    SupportsBufferProtocol = object
    Array = object
    Device = object
    DType = object


# 第 65 章 —— JustInt/JustFloat 协议:精确数值类型约束
@final
class JustInt(Protocol):
    @property
    def __class__(self, /) -> type[int]: ...
    @__class__.setter
    def __class__(self, value: type[int], /) -> None: ...

@final
class JustFloat(Protocol):
    @property
    def __class__(self, /) -> type[float]: ...
    @__class__.setter
    def __class__(self, value: type[float], /) -> None: ...

@final
class JustComplex(Protocol):
    @property
    def __class__(self, /) -> type[complex]: ...
    @__class__.setter
    def __class__(self, value: type[complex], /) -> None: ...

# 第 65 章 —— 嵌套序列协议
class NestedSequence(Protocol[_T_co]):
    def __getitem__(self, key: int, /) -> _T_co | NestedSequence[_T_co]: ...
    def __len__(self, /) -> int: ...

# 第 65 章 —— 数组命名空间协议
class SupportsArrayNamespace(Protocol[_T_co]):
    def __array_namespace__(self, /, *, api_version: str | None) -> _T_co: ...

# 第 65 章 —— 有 shape 属性的协议
class HasShape(Protocol[_T_co]):
    @property
    def shape(self, /) -> _T_co: ...

这段代码定义了跨后端的核心类型别名与协议。Array/Device/DType 在类型检查时是 Incomplete(让类型检查器推断),运行时是 object(避免循环导入)。JustInt/JustFloat/JustComplex 是特殊协议,通过 __class__ 属性欺骗 mypy/pyright 强制让类型检查器相信某个值是特定类型。NestedSequence/SupportsArrayNamespace/HasShape 则是结构性协议,描述数组 API 要求的接口。

TypedDict 部分定义 __array_namespace_info__ 的返回类型:

Capabilities = TypedDict(
    "Capabilities",
    {
        "boolean indexing": bool,
        "data-dependent shapes": bool,
        "max dimensions": int,
    },
)

DefaultDTypes = TypedDict(
    "DefaultDTypes",
    {
        "real floating": DType,
        "complex floating": DType,
        "integral": DType,
        "indexing": DType,
    },
)

# 第 65 章 —— DTypeKind 是 kind 参数的合法值
_DTypeKind: TypeAlias = Literal[
    "bool", "signed integer", "unsigned integer", "integral",
    "real floating", "complex floating", "numeric",
]
DTypeKind: TypeAlias = _DTypeKind | tuple[_DTypeKind, ...]

# 第 65 章 —— 不同 kind 对应的返回 TypedDict
class DTypesBool(TypedDict):
    bool: DType

class DTypesSigned(TypedDict):
    int8: DType
    int16: DType
    int32: DType
    int64: DType

class DTypesUnsigned(TypedDict):
    uint8: DType
    uint16: DType
    uint32: DType
    uint64: DType

class DTypesIntegral(DTypesSigned, DTypesUnsigned):
    pass

class DTypesReal(TypedDict):
    float32: DType
    float64: DType

class DTypesComplex(TypedDict):
    complex64: DType
    complex128: DType

class DTypesNumeric(DTypesIntegral, DTypesReal, DTypesComplex):
    pass

class DTypesAll(DTypesBool, DTypesNumeric):
    pass

DTypesAny: TypeAlias = Mapping[str, DType]

Capabilities 描述后端能力(是否支持布尔索引、数据依赖形状、最大维度数)。DefaultDTypes 描述各类型 kind 的默认 dtype。TypedDict 让 IDE 和 mypy 能正确推断返回字典的键值类型。

65.13.2 get_xp 装饰器

get_xp 是构建兼容层 API 的核心工具:

源码路径:sklearn/externals/array_api_compat/_internal.py - get_xp()

from inspect import signature
from functools import wraps

def get_xp(xp: ModuleType) -> Callable[[Callable[..., _T]], Callable[..., _T]]:
    """
    装饰器工厂:自动将 xp 替换为对应的数组模块。

    使用方式:
    @get_xp(np)
    def func(x, /, xp, kwarg=None):
        return xp.func(x, kwarg=kwarg)

    注意 xp 必须是 keyword argument,且位于所有非 keyword 之后。
    """
    def inner(f: Callable[..., _T], /) -> Callable[..., _T]:
        @wraps(f)
        def wrapped_f(*args: object, **kwargs: object) -> object:
            return f(*args, xp=xp, **kwargs)

        sig = signature(f)
        new_sig = sig.replace(
            parameters=[par for i, par in sig.parameters.items() if i != "xp"]
        )
        if wrapped_f.__doc__ is None:
            wrapped_f.__doc__ = f"""\
Array API compatibility wrapper for {f.__name__}.

See the corresponding documentation in NumPy/CuPy and/or the array API
specification for more details.
"""
        wrapped_f.__signature__ = new_sig
        return wrapped_f

    return inner

这段代码定义了 get_xp 装饰器工厂。get_xp(np) 返回一个装饰器,该装饰器包装原始函数,自动注入 xp=np 参数。wrapped_f.__signature__ = new_sig 修改包装函数的签名,移除 xp 参数(这样 IDE 自动补全不会显示它)。@wraps(f) 保留原函数的元数据(__name____doc__ 等)。

使用示例:

@get_xp(np)
def arange(start, /, xp, dtype=None, **kwargs):
    return xp.arange(start, dtype=dtype, **kwargs)

调用 arange(5) 等价于 np.arange(5),但用户不需要关心命名空间来源。

65.13.3 clone_module 工具

clone_module 用于批量导入原生库符号:

源码路径:sklearn/externals/array_api_compat/_internal.py - clone_module()

import importlib

def clone_module(mod_name: str, globals_: dict[str, object]) -> list[str]:
    """从模块导入所有符号到 globals()。返回 __all__。"""
    mod = importlib.import_module(mod_name)
    objs = {}
    exec(f"from {mod.__name__} import *", objs)

    for n in dir(mod):
        if not n.startswith("_") and hasattr(mod, n):
            objs[n] = getattr(mod, n)

    globals_.update(objs)
    return list(objs)


__all__ = ["get_xp", "clone_module"]

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

这段代码实现模块克隆。exec(f"from {mod.__name__} import *") 触发 Python 的 import * 机制,导入 __all__ 列表中的符号(如果模块定义了)。dir(mod) 遍历所有公共属性(包括 __all__ 之外的)。两种方法结合可以覆盖大多数库的导出布局。返回的列表用作 __all____dir__ 自定义 IDE 自动补全列表。

65.13.4 后端包初始化模式

所有后端包的 __init__.py 都遵循相同的模式:先 clone_module 批量导入原生符号,然后覆盖/添加 Array API 特有的符号:

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

# 第 65 章 —— ruff: noqa: PLC0414
from typing import Final

from .._internal import clone_module

# 第 65 章 —— 必须显式加载(clone_module 需要它)
import numpy.typing  # noqa: F401

# 第 65 章 —— 批量导入 numpy 的所有公共符号
__all__ = clone_module("numpy", globals())

# 第 65 章 —— 接下来可能覆盖上面导入的名字
from . import _aliases
from ._aliases import *  # noqa: F403
from ._info import __array_namespace_info__  # noqa: F401

# 第 65 章 —— 必须用 __import__ 动态导入 linalg 和 fft
__import__(__package__ + ".linalg")
__import__(__package__ + ".fft")

from .linalg import matrix_transpose, vecdot  # noqa: F401

__array_api_version__: Final = "2024.12"

__all__ = sorted(
    set(__all__)
    | set(_aliases.__all__)
    | {"__array_api_version__", "__array_namespace_info__", "linalg", "fft"}
)

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

CuPy/Torch/Dask 后端的初始化模式类似:

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

from typing import Final
from cupy import *  # noqa: F403

from cupy import abs, max, min, round  # noqa: F401

from ._aliases import *  # noqa: F403
from ._info import __array_namespace_info__  # noqa: F401

__import__(__package__ + '.linalg')
__import__(__package__ + '.fft')

__array_api_version__: Final = '2024.12'

__all__ = sorted(
    {name for name in globals() if not name.startswith("__")}
    - {"Final", "_aliases", "_info", "_typing"}
    | {"__array_api_version__", "__array_namespace_info__", "linalg", "fft"}
)

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

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

from typing import Final

from .._internal import clone_module

__all__ = clone_module("torch", globals())

from . import _aliases
from ._aliases import *  # noqa: F403
from ._info import __array_namespace_info__  # noqa: F401

__import__(__package__ + '.linalg')
__import__(__package__ + '.fft')

__array_api_version__: Final = '2024.12'

__all__ = sorted(
    set(__all__)
    | set(_aliases.__all__)
    | {"__array_api_version__", "__array_namespace_info__", "linalg", "fft"}
)

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

Dask 后端与其他类似:

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

from typing import Final

from ..._internal import clone_module

__all__ = clone_module("dask.array", globals())

from . import _aliases
from ._aliases import *  # noqa: F403
from ._info import __array_namespace_info__  # noqa: F401

__array_api_version__: Final = "2024.12"
del Final

__import__(__package__ + '.linalg')
__import__(__package__ + '.fft')

__all__ = sorted(
    set(__all__)
    | set(_aliases.__all__)
    | {"__array_api_version__", "__array_namespace_info__", "linalg", "fft"}
)

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

四种后端的初始化模式高度一致:1) clone_module 批量导入;2) _aliases 可能覆盖;3) _info 添加 __array_namespace_info__;4) __import__ 动态加载 linalg/fft 子模块(因为相对导入不会覆盖主命名空间的 linalg/fft);5) 设置 __array_api_version__ = "2024.12";6) 重新计算 __all____dir__

Dask 还包含一个空的 dask/__init__.py

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


这个空文件唯一的作用是标记 dask 是一个 Python 包——没有它 Python 不会识别 array_api_compat.dask 为子包。

65.14 设计中的取舍

问:为什么 scikit-learn 选择 vendor array-api-compat 而不是作为运行时依赖?

答:核心原因是是版本稳定与与离线支持。PEP 440 版本解析、ARFF 格式读写、NumPy 风格 docstring 解析——这些是 scikit-learn 内部工具链的基础设施,一旦上游库(如 packaging)的版本变化导致行为变化,scikit-learn 的发布就会受影响。vendor 让 scikit-learn 精确控制每个工具的版本与行为,且 wheel 包自包含,用户安装时无需额外网络请求。这种设计的代价是维护负担——上游库的安全修复需要手动同步。但对于 scikit-learn 来说,这些内部工具(packaging、numpydoc)已经相当稳定,vendor 的维护成本可控。

问:Array API 兼容层的 trade-off 是什么?

答:兼容层的设计在零开销识别完全 API 对齐之间做了精细的平衡。_issubclass_fast 通过 sys.modules 缓存实现零开销类型检测,但要求所有后端的类名固定(如 "numpy""cupy")。这种约定简化了实现,但也意味着如果用户传入了自定义子类(如 my.numpy.MyArray),兼容层可能无法识别。

完全 API 对齐的代价是每个函数都要写专门的修正——PyTorch 的类型提升、CuPy 的设备上下文、Dask 的惰性约束、NumPy 的 copy 语义。这种"逐函数修补"的策略相比"重写整个 API"复杂度更高,但避免了完全重写实现带来的性能损失。scikit-learn 选择复用通用别名层(common/_aliases)作为默认实现,仅在必要时做后端特化,这是一种务实的折中。

问:use_compat 三态设计的意义是什么?

答:另一个重要的取舍是 use_compat 三态设计。None(默认)尽量返回原生库,但当原生库 API 与 Array API 不一致时仍返回 compat 包装;True 强制返回 compat;False 强制返回原生。这种设计让 scikit-learn 可以在新版本中平滑过渡——先用 compat 包装测试,确认无误后再让原生库接管。

问:为什么选择 vendor SciPy 的 _laplacian.py 而不是直接依赖 SciPy 1.12+?

答:vendor 出于向前兼容与依赖控制的考虑。scikit-learn 的最低支持版本可能远低于 SciPy 1.12,强制升级 SciPy 会影响大量现有用户。vendor 让新功能(如 sparse array 支持的图拉普拉斯)可以立即在 scikit-learn 中可用,而不必等待所有用户升级 SciPy。这种设计虽然增加了一点维护负担(需要手动同步上游修复),但提供了更好的向后兼容性。

问:vendor 模式相比 pip install 的最大优势是什么?

答:vendor 的最大优势是可重现性。一旦 vendored 版本的代码被打包进 scikit-learn 的 wheel,用户安装的版本就完全确定——不会因 PyPI 上传了不兼容的新版本而导致已安装的 scikit-learn 出现意外行为。对于科学计算库来说,这种确定性至关重要,因为算法结果的轻微差异可能导致论文结论改变。

65.15 动手练习

  1. 阅读外部依赖打包与 ARFF 解析器

    阅读 sklearn/externals/_array_api_compat_vendor.pysklearn/externals/_arff.py

    • 理解 vendor 脚本如何将 array-api-compat 源码复制到 sklearn/externals/array_api_compat 目录

    • 分析 ARFF 解析器如何处理稀疏数据格式 {0 1.0, 2 3.0} 与多标签属性定义

    • 对比 ARFF 解析器与 pandas.read_csv 处理缺失值、注释行的差异

    回答问题:

    • 为什么 scikit-learn 选择内置 array-api-compat 而不是作为运行时依赖?

    • ARFF 解析器中 _parse_sparse_data 如何将稀疏字符串转换为 Coo 矩阵的索引?

    • 文档字符串中 @paramParameters 章节的解析优先级如何?

  2. 探究 PEP 440 版本解析与比较逻辑

    阅读 sklearn/externals/_packaging/version.py_structures.py

    • 追踪 Version.__init__ 如何解析 '1.0a1.post2.dev3' 等复杂版本字符串

    • 理解 Version.__lt__ 如何实现预发布版 < 正式版 < 后发布版的排序规则

    • 分析 NormalizedVersionLegacyVersion 的兼容策略

    回答问题:

    • Version 类中 _version_regex 正则如何捕获 epoch、release、pre、post、dev 五大组件?

    • 为什么 '1.0.dev0' < '1.0a0' < '1.0' < '1.0.post0'?

    • LegacyVersion 如何处理不符合 PEP 440 的传统版本号(如 '1.0.r123')?

  3. 分析 Array API 兼容层核心分发器的后端识别机制

    阅读 sklearn/externals/array_api_compat/common/_helpers.pyarray_namespace_cls_to_namespace 函数:

    • 追踪 array_namespace(np.array([1]), cp.array([1])) 调用时的报错路径

    • 理解 use_compat=None/True/False 三种模式对 NumPy 后端的不同影响

    • 分析 _ClsToXPInfo.SCALARMAYBE_JAX_ZERO_GRADIENT 如何处理 Python 标量与 JAX 零梯度数组

    回答问题:

    • 为什么 _cls_to_namespaceissubclass(cls, int|float|complex|None) 判断必须在 np.generic 之后?

    • array_namespace 如何保证多输入数组来自同一后端?违反时抛出什么异常?

  4. 对比 clip 函数的跨后端实现差异

    对比阅读三个后端的 clip 实现:

    • sklearn/externals/array_api_compat/common/_aliases.py 的通用实现

    • sklearn/externals/array_api_compat/torch/_aliases.pyclip = get_xp(torch)(_aliases.clip)(复用通用)

    • sklearn/externals/array_api_compat/dask/array/_aliases.py 中的独立实现

    回答问题:

    • 通用实现如何通过 out[()] = x + 掩码赋值实现 dtype 保持?为什么要处理 Python 整数溢出截断?

    • Dask 为何不能使用通用掩码实现?它的替代方案是什么?有什么局限?

    • PyTorch 为何选择复用通用实现而非原生 torch.clamp

  5. 研究设备抽象的跨后端统一

    阅读 sklearn/externals/array_api_compat/common/_helpers.pydevice()to_device() 函数:

    • 对比 NumPy、Dask、CuPy、PyTorch、JAX、Sparse 六个后端的 device 返回值差异

    • 分析 to_device 中 CuPy 的 _cupy_to_device 如何利用 cp.cuda.Device 上下文管理器与 stream 参数

    • 理解 PyTorch _torch_to_device 为何不支持 stream 参数

    回答问题:

    • Dask 为何引入 _DASK_DEVICE 单例对象而非字符串?它如何区分 CPU/GPU 元数据?

    • JAX to_devicejax.jit 上下文中为何可能失效?代码中有什么 workaround?

    • Sparse 数组的 device 递归查找逻辑(x.data 回退)处理了哪种存储格式例外?

  6. 实现跨后端的 unique_all 函数

    阅读 sklearn/externals/array_api_compat/common/_aliases.pyunique_all 与各后端的实现差异:

    • 通用实现如何基于 xp.unique 返回 NamedTuple 并修正 inverse_indices shape

    • PyTorch 为何抛出 NotImplementedError(缺失 indices 返回)

    • NumPy/CuPy 如何通过条件导出使用原生 unique_all(NumPy 2.0+)

    回答问题:

    • 通用实现中 inverse_indices.reshape(x.shape) 修正了什么 NumPy 行为?

    • PyTorch 实现 unique_all 的主要障碍是什么?有无替代方案?

  7. 排查 Dask 后端 sort/argsort 的内存风险

    阅读 sklearn/externals/array_api_compat/dask/array/_aliases.pysortargsort 实现:

    • 理解 _ensure_single_chunk 如何将分块重组作为单分块

    • 分析 map_blocks 调用元后端实现的开销与 restore 闭包的分块恢复风险

    • 对比 Dask 与 NumPy/CuPy 实现的根本差异(惰性 vs 急切)

    回答问题:

    • 为什么 Dask 必须重分块为单分块才能排序?这对大规模数据意味着什么?

    • restore 闭包为何可能导致过度分块?如何在生产环境中规避?

65.16 本章小结

这一章中我们学习了 scikit-learn 的外部工具与依赖管理体系,了解了一个大型 Python项目库如何实现"自给自足"的工具箱。我们首先理解了为什么了 scikit-learn 选择 vendor 第三方 库而非作为运行时依赖——版本稳定、离线支持、避免依赖冲突;其次,我们深入 ARFF 解析器的状态机驱动实现,看它如何在 Weka 与 scikit-learn 间传递数据,包括稀疏格式、多标签属性、缺失值处理等细节。接着,我们学习了 NumPy 风格 docstring 解析器如何将非结构化文本提炼为结构化章节,以及 PEP 440 版本解析如何用单一正则捕获五大组件、用 Infinity 哨兵实现 dev < alpha < release < post 的排序规则以及 LegacyVersion 兑底。然后,我们探索了 SciPy 稀疏图拉普拉斯计算的独立封装,理解了就地修改、归一化技巧与 LinearOperator 输出。最后,我们深入 Array API 兼容层——这是 scikit-learn 最庞大的外部依赖——从核心分发器 array_namespace 的零开销类型检测,到 is_*_namespace 系列命名空间判断函数,到设备抽象与 to_device 的跨后端统一,到 is_*_array 与惰性/可写性判断,到通用别名层的 clip/unique_all/cumulative_sum 等语义对齐,再到后端专属适配器对 CuPy/PyTorch/Dask 的定制化修正(类型提升、设备上下文、惰性约束等),最后到类型契约与构建工具(get_xp 装饰器、clone_module 模块克隆、dir 自定义)。

本章我们一起学习了以下概念,以下表格总结了本章涉及的核心概念及其作用:

| 概念 | 解释 |

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

| sklearn/externals/_array_api_compat_vendor.py | 外部依赖打包脚本,将 array-api-compat 等第三方库源码内置到 sklearn/externals/array_api_compat 目录 |

| sklearn/externals/_arff.py | ARFF 格式完整读写实现,支持密集、稀疏数据、多标签、注释、关系声明,兼容 Weka 生态 |

| sklearn/externals/_numpydoc/docscrape.py | NumPy 风格文档字符串解析器,提取 Parameters、Returns、See Also 等章节为结构化对象 |

| sklearn/externals/_packaging/version.py | PEP 440 标准版本解析与比较,支持预发布、后发布、开发版等复杂版本号的排序与规范化 |

| sklearn/externals/_scipy/sparse/csgraph/_laplacian.py | 稀疏图拉普拉斯矩阵计算,支持归一化、对称化,返回 ndarray/sparse matrix/LinearOperator |

| array_api_compat.init | 顶层入口,声明版本与重新导出 common 子包 |

| array_api_compat.common.init | 从 _helpers 重新导出所有公开符号 |

| array_api_compat.common._helpers.array_namespace | 核心分发器,识别输入数组后端并返回统一兼容命名空间,支持 api_version 与 use_compat 控制 |

| array_api_compat.common._helpers._cls_to_namespace | 类型到命名空间的映射表,含 NumPy/CuPy/PyTorch/Dask/JAX 分支与 use_compat 三态逻辑 |

| array_api_compat.common.helpers.is*_namespace | 命名空间类型判断函数族(8 个),通过模块名匹配判断 xp 是否为特定后端的命名空间 |

| array_api_compat.common.helpers.isarray | 数组对象类型判断函数族,与 is_namespace 互补 |

| array_api_compat.common._helpers._issubclass_fast | 零开销类型检测,基于 sys.modules 缓存避免重复导入与属性查找 |

| array_api_compat.common._helpers.device/to_device | 设备抽象统一:NumPy/Dask 固定 cpu,CuPy/PyTorch/JAX 支持真实设备与 stream 复制 |

| array_api_compat.common._helpers.is_lazy_array / is_writeable_array | 惰性与可写性判断,避免触发全图计算 |

| array_api_compat.common._helpers.size | 跨后端元素总数,处理 None 与 NaN |

| array_api_compat.common._aliases | 统一接口层:创建函数注入 device、clip 语义对齐、cumulative_sum/include_initial、unique 系列返回 NamedTuple、std/var 参数重命名、isdtype 类型判断 |

| array_api_compat.common._linalg | 线性代数标准化:SVD/Eigh/QR/Slogdet 返回命名元组、cholesky upper 选项、matrix_rank/pinv 统一 rtol、vector_norm 多轴支持 |

| array_api_compat.common._fft | FFT 标准化:统一 norm 参数、强制 float32→complex64 精度保持、fftfreq 仅支持 cpu device |

| array_api_compat.numpy | NumPy 后端适配:asarray copy 语义(_CopyMode)、astype device 校验、count_nonzero 标量修复、整数 ceil/floor/trunc 行为修正 |

| array_api_compat.numpy._info | NumPy 后端元信息:capabilities/default_device/dtypes/devices |

| array_api_compat.numpy._typing | NumPy 后端类型定义:Device Literal["cpu"]、Array np.ndarray |

| array_api_compat.numpy.linalg / numpy.fft | NumPy 后端的 linalg 与 fft 子模块,包含特有的 solve 函数 |

| array_api_compat.cupy | CuPy 后端适配:asarray 设备上下文管理、astype 跨设备复制、count_nonzero keepdims 补全、整数 ceil/floor/trunc 返回副本 |

| array_api_compat.cupy._info | CuPy 后端元信息:default_device 返回 cuda.Device(0)、devices 返回 GPU 列表 |

| array_api_compat.cupy._typing | CuPy 后端类型定义:Device 为 cp.cuda.Device |

| array_api_compat.torch._aliases | PyTorch 适配核心:_fix_promotion 解决 0-D 类型提升、dim→axis 重命名、返回元组解析、缺失 keepdims 补全、complex sign 语义修正 |

| array_api_compat.torch.linalg | PyTorch 线性代数修正:cross 默认 axis=-1、vecdot 整数支持、solve 1D RHS 歧义消除、trace offset 支持、vector_norm axis=() 特判 |

| array_api_compat.torch.fft | PyTorch FFT 修正:fftn/ifftn/rfftn/irfftn/fftshift/ifftshift 的 axes→dim 重命名 |

| array_api_compat.torch._info | PyTorch 后端元信息:从错误信息中提取 devices 列表 |

| array_api_compat.torch._typing | PyTorch 后端类型定义:直接从 torch 导入 device/dtype/Tensor |

| array_api_compat.dask.array._aliases | Dask 适配:惰性计算约束下的 sort/argsort 强制单分块、clip 无掩码法(用 min/max 组合)、asarray 非 dask 输入强制 copy |

| array_api_compat.dask.array.linalg | Dask 线性代数受限:qr 仅 reduced、svd 禁用 full_matrices、svdvals 无法避免计算 U/V |

| array_api_compat.dask.array._info | Dask 后端元信息:默认 cpu,devices 返回 ["cpu", _DASK_DEVICE] |

| array_namespace_info | 各后端实现的元信息查询接口:capabilities/default_device/dtypes/devices |

| _internal.get_xp / clone_module / dir | 构建工具:get_xp 自动绑定 xp 参数装饰器、clone_module 批量导入原生符号、dir 自定义 IDE 自动补全 |

| common._typing | 跨后端类型契约:Array/Device/DType/Capabilities/DefaultDTypes 等 TypedDict 与 Protocol 定义 |

| 后端 init.py | 统一初始化模式:clone_module + _aliases 覆盖 + _info 注入 + import 加载 linalg/fft + array_api_version + all + dir |

下一章中,我们将学习 Array API Extra 扩展工具集 —— 超越标准的"增强包",理解 lazy_apply 惰性求值、at 类函数式更新、委托分派机制等扩展能力如何为 scikit-learn 的跨后端计算增添灵活性。

第 66 章 —— Array API Extra 扩展工具集 —— 超越标准的“增强包”

66.1 学习目标

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

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

在正式开始之前,先明确本章节的学习目标——这有助于在阅读时保持方向感。

  • 理解核心架构:掌握 array-api-compatarray-api-extra 的层次关系以及命名空间自动识别机制。

  • 统一别名与包装:了解如何通过 _aliases.py 将不同后端的函数签名统一为标准 API。

  • 惰性求值:深入 lazy_apply 在 Dask 与 JAX 中的实现细节以及何时触发计算。

  • 只读更新:学会使用 at 进行函数式“就地”更新,特别是针对 JAX 的不可变数组。

  • 函数委托:掌握后端原生实现与标准实现之间的智能路由策略。

  • 扩展函数:熟悉 covkronsinc 等超出标准的实用工具。

  • 跨后端测试:了解 xp_assert_*lazy_xp_function 等测试设施如何保持行为一致。

生活类比:把 array-api-compat 想成一座“国际翻译中心”,而 array-api-extra 则是这座中心的“增能套件”。后端(NumPy、CuPy、PyTorch、JAX、Dask …)是不同国家的语言,array_namespace() 就像护照读取器,一眼识别语言并切换翻译模式。atlazy_apply_delegation 分别对应“快捷键盘手”“延迟执行的调度员”“智能路由员”。


66.2 源码地图(带简要说明)

| 路径 | 说明 |

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

| sklearn/externals/array_api_extra/__init__.py | 对外暴露的公共 API 与版本号。 |

| sklearn/externals/array_api_extra/_lib/_at.py | 只读数组更新 (at) 的实现细节。 |

| sklearn/externals/array_api_extra/_lib/_lazy.py | 惰性求值 (lazy_apply) 的核心逻辑。 |

| sklearn/externals/array_api_extra/_lib/_funcs.py | 基础扩展函数:apply_wherebroadcast_shapescovkronsinc 等。 |

| sklearn/externals/array_api_extra/_delegation.py | 函数委托:优先调用后端原生实现。 |

| sklearn/externals/array_api_extra/_lib/_testing.py | 跨后端断言工具(xp_assert_*)。 |

| sklearn/externals/array_api_extra/testing.py | 惰性后端测试装饰器与调度器。 |

| sklearn/externals/array_api_extra/_lib/_backends.py | 测试使用的后端枚举与 pytest 参数化。 |

| sklearn/externals/array_api_extra/_lib/_utils/* | 辅助工具:兼容层、类型系统、元数据等。 |

提示:每个源码文件后面都会给出 完整的源码路径声明(如 sklearn/externals/array_api_extra/_lib/_lazy.py),并配以 逐行注释功能解释,请务必仔细阅读。

生活类比:就像翻译中心的每个增能套件都有独特的使用场景,array-api-extra 的每个模块也各司其职——比如 at 是快捷键,让你在只读数组上也能像在可变数组上一样“点石成金”;lazy_apply 则像一个懂得分批调度的助理,在大数据或 GPU 上不急于立即执行,而是等到真正需要结果时才启动计算。


66.3 设计中的取舍

66.3.1 设计取舍概览

| 取舍点 | 选项 | 为什么这样选 |

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

| 性能 vs. 通用性 | 优先使用后端原生实现(如 xp.isclose) | 能够获得底层高度优化的性能,尤其在 GPU 上提升显著。 |

| 性能 vs. 通用性 | 当后端不支持时退回标准实现 | 保证功能可用性,避免因缺失实现导致 API 不完整。 |

| 惰性求值的安全性 | 对 Dask 强制单块计算 | 防止跨 worker 内存爆炸,牺牲分布式潜力以保证单机可执行。 |

| 惰性求值的安全性 | 在 JAX jit 中要求完整形状 | 防止 TracerBoolConversionError,确保编译阶段安全。 |

| 更新操作的副作用 | copy=None 允许就地修改 | 提升性能,但文档中强制立即重新赋值以避免意外共享。 |

| 测试隔离 | lazy_xp_function + patch_lazy_xp_functions | 通过 monkey‑patch 实现后端特定行为检测;需要显式标记为 thread_unsafe。 |

一问一答(取舍 1)

:为何在 lazy_apply 中对 Dask 强制单块重分块?

:单块可以保证所有计算在同一进程的内存中完成,防止跨 worker 的数据搬迁导致 OOM;这是在保证安全的前提下的性能折中。

生活类比:就像翻译中心的套餐设计——有时为了快速响应(性能),我们优先使用当地母语者(后端原生实现);但为了确保 nessun 被落下(通用性),当某地方言暂时不支持时,我们退回到通用翻译手册(标准实现)。


66.4 惰性求值与 JIT 包装 —— 延迟计算的“时间管理者”

66.4.1 代码实作(完整路径 + 逐行注释)

# 第 66 章 —— File: sklearn/externals/array_api_extra/_lib/_lazy.py
"""Public API Functions."""

from __future__ import annotations

import math
from collections.abc import Callable, Sequence
from functools import partial, wraps
from types import ModuleType
from typing import TYPE_CHECKING, Any, ParamSpec, TypeAlias, cast, overload

from ._funcs import broadcast_shapes
from ._utils import _compat
from ._utils._compat import (
    array_namespace,
    is_dask_namespace,
    is_jax_namespace,
)
from ._utils._helpers import is_python_scalar
from ._utils._typing import Array, DType

# 第 66 章 —— ----------------------------------------------------------------------
# 第 66 章 —— `lazy_apply` 参数签名(支持单输出 / 多输出)
# 第 66 章 —— ----------------------------------------------------------------------
@overload
def lazy_apply(
    func: Callable[P, Array | ArrayLike],
    *args: Array | complex | None,
    shape: tuple[int | None, ...] | None = None,
    dtype: DType | None = None,
    as_numpy: bool = False,
    xp: ModuleType | None = None,
    **kwargs: P.kwargs,
) -> Array: ...

@overload
def lazy_apply(
    func: Callable[P, Sequence[Array | ArrayLike]],
    *args: Array | complex | None,
    shape: Sequence[tuple[int | None, ...]],
    dtype: Sequence[DType] | None = None,
    as_numpy: bool = False,
    xp: ModuleType | None = None,
    **kwargs: P.kwargs,
) -> tuple[Array, ...]: ...

def lazy_apply(
    func: Callable[P, Array | ArrayLike | Sequence[Array | ArrayLike]],
    *args: Array | complex | None,
    shape: tuple[int | None, ...] | Sequence[tuple[int | None, ...]] | None = None,
    dtype: DType | Sequence[DType] | None = None,
    as_numpy: bool = False,
    xp: ModuleType | None = None,
    **kwargs: P.kwargs,
) -> Array | tuple[Array, ...]:
    """
    Lazily apply an eager function.
    """
    # 1️⃣ 过滤掉 `None` 参数,只保留实际输入
    args_not_none = [arg for arg in args if arg is not None]
    # 2️⃣ 将标量从数组列表中剔除,得到真正的数组参数
    array_args = [arg for arg in args_not_none if not is_python_scalar(arg)]
    if not array_args:
        msg = "Must have at least one argument array"
        raise ValueError(msg)

    # 自动推断后端命名空间(如果未显式提供)
    if xp is None:
        xp = array_namespace(*args)  # ← `array_namespace` 负责后端识别

    # ------------------------------------------------------------------
    # 形状 / dtype 解析与校验
    # ------------------------------------------------------------------
    shapes: list[tuple[int | None, ...]]
    dtypes: list[DType]
    multi_output = False

    if shape is None:
        # 默认广播所有输入数组的形状
        shapes = [broadcast_shapes(*(arg.shape for arg in array_args))]
    elif all(isinstance(s, int | None) for s in shape):
        shapes = [cast(tuple[int | None, ...], shape)]
    else:
        shapes = list(shape)
        multi_output = True

    if dtype is None:
        # 通过 `xp.result_type` 推断返回 dtype
        dtypes = [xp.result_type(*args_not_none)] * len(shapes)
    elif multi_output:
        if not isinstance(dtype, Sequence):
            msg = "Got multiple shapes but only one dtype"
            raise ValueError(msg)
        dtypes = list(dtype)
    else:
        if isinstance(dtype, Sequence):
            msg = "Got single shape but multiple dtypes"
            raise ValueError(msg)
        dtypes = [dtype]

    if len(shapes) != len(dtypes):
        msg = f"Got {len(shapes)} shapes and {len(dtypes)} dtypes"
        raise ValueError(msg)

    # ------------------------------------------------------------------
    # 2️⃣ 后端特化分支
    # ------------------------------------------------------------------
    if is_dask_namespace(xp):
        # ---- Dask 分支 ----
        import dask
        # 使用 Dask 的 meta‑namespace(即底层实际数组的命名空间)
        metas: list[Array] = [arg._meta for arg in array_args]
        meta_xp = array_namespace(*metas)

        # 使用 `dask.delayed` 将函数包装为延迟任务
        wrapped = dask.delayed(
            _lazy_apply_wrapper(func, as_numpy, multi_output, meta_xp),
            pure=True,
        )
        delayed_out = wrapped(*args, **kwargs)

        # 将延迟对象包装回 Dask 数组,使用 `math.nan` 表示未知维度
        out = tuple(
            xp.from_delayed(
                delayed_out[i],
                shape=tuple(math.nan if s is None else s for s in shape),
                dtype=dtype,
                meta=metas[0],
            )
            for i, (shape, dtype) in enumerate(zip(shapes, dtypes, strict=True))
        )

    elif is_jax_namespace(xp) and _is_jax_jit_enabled(xp):
        # ---- JAX JIT 分支 ----
        import jax

        # JIT 中必须知道完整的输出形状
        if any(None in shape for shape in shapes):
            msg = "Output shape must be fully known when running inside jax.jit"
            raise ValueError(msg)

        # 屏蔽 kwargs 防止它们被强制转为 JAX 数组
        wrapped = _lazy_apply_wrapper(
            partial(func, **kwargs), as_numpy, multi_output, xp
        )

        # 使用 `jax.pure_callback` 把惰性函数桥接到 eager 执行
        out = cast(
            tuple[Array, ...],
            jax.pure_callback(
                wrapped,
                tuple(
                    jax.ShapeDtypeStruct(shape, dtype)
                    for shape, dtype in zip(shapes, dtypes, strict=True)
                ),
                *args,
            ),
        )
    else:
        # ---- Eager 后端(NumPy、CuPy、PyTorch 等)----
        wrapped = _lazy_apply_wrapper(func, as_numpy, multi_output, xp)
        out = wrapped(*args, **kwargs)

    # 如果是多输出则返回元组,否则返回单个数组
    return out if multi_output else out[0]

66.4.1.1 代码作用说明

  • 自动后端识别array_namespace(*args) 根据输入参数的类型返回对应的命名空间(NumPy、CuPy、PyTorch、JAX、Dask 等)。

  • 形状/ dtype 统一:即使后端返回 math.nan(Dask)或 None(Array API),最终都会统一为 None,保持跨后端一致性。

  • 后端分支

    • Dask:使用 dask.delayed 延迟执行,并在 from_delayed 时把未知维度标记为 math.nan

    • JAX JIT:通过 jax.pure_callbackjit 环境下安全调用普通函数;若形状未知则提前报错。

    • Eager:直接执行函数,无额外包装。

为什么需要惰性求值?

在大规模数据或 GPU 计算中,立即执行会导致不必要的内存占用和计算开销。lazy_apply 让计算图在真正需要结果时才 materialize,兼容 Dask 分块和 JAX JIT 编译两大惰性后端。

生活类比:想象你是一位厨师(函数),而食材是分布在不同仓库(后端)的食材。lazy_apply 像一个智能助理——它不会立刻去所有仓库取食材(避免不必要的搬运),而是等到你真的要开始烹饪(需要结果)时,才根据食材所在的仓库(后端类型)选择最快的取货方式:如果是冷冻仓库(Dask),它会先把所有食材集中到一个装卸平台(单块);如果是高压锅仓库(JAX JIT),它会先确认食谱需要的确切份量(完整形状)才开始烹饪;否则,它直接在就近的厨房(Eager 后端)动手。


66.4.2 检测 JAX JIT 上下文的辅助函数

# 第 66 章 —— File: sklearn/externals/array_api_extra/_lib/_lazy.py
def _is_jax_jit_enabled(xp: ModuleType) -> bool:
    """Return True if this function is being called inside ``jax.jit``."""
    import jax  # pylint: disable=import-outside-toplevel

    # 在 JIT 中,`xp.asarray(False)` 返回一个 tracer。尝试将 tracer 转为 bool
    # 会抛出 `TracerBoolConversionError`,我们捕获它即可判断是否处于 JIT 环境。
    x = xp.asarray(False)
    try:
        return bool(x)
    except jax.errors.TracerBoolConversionError:
        return True

解释jax.jit 会把所有值包装为 tracer,这些 tracer 无法直接转为布尔值。此函数利用异常来判断当前是否在 JIT 编译阶段,从而在 lazy_apply 中决定是否强制要求完整形状。

生活类比:就像在高压锅(JAX JIT)里烹饪时,你不能用普通的手(tracer)去尝汤是否咸淡——手会溶掉!所以得用特殊的探针(异常捕获)来判断是否真的在高压状态下。


66.4.3 包装器:_lazy_apply_wrapper

# 第 66 章 —— File: sklearn/externals/array_api_extra/_lib/_lazy.py
def _lazy_apply_wrapper(
    func: Callable[..., Array | ArrayLike | Sequence[Array | ArrayLike]],
    as_numpy: bool,
    multi_output: bool,
    xp: ModuleType,
) -> Callable[..., tuple[Array, ...]]:
    """
    Helper of `lazy_apply`.
    """
    @wraps(func)
    def wrapper(*args: Array | complex | None, **kwargs: Any) -> tuple[Array, ...]:
        args_list = []
        device = None
        for arg in args:
            if arg is not None and not is_python_scalar(arg):
                if device is None:
                    device = _compat.device(arg)        # 记录设备(CPU / GPU)
                if as_numpy:
                    # 当 `as_numpy=True` 时,把输入强制转为 NumPy,以兼容只接受 NumPy 的函数
                    import numpy as np
                    arg = cast(Array, np.asarray(arg))
                args_list.append(arg)
        assert device is not None

        # 执行原始函数
        out = func(*args_list, **kwargs)

        # 根据返回值类型统一为元组
        if multi_output:
            assert isinstance(out, Sequence)
            return tuple(xp.asarray(o, device=device) for o in out)
        return (xp.asarray(out, device=device),)

    return wrapper

关键点

  • as_numpy:在需要调用仅接受 NumPy 的 Cython/Numba 实现时使用。

  • 设备保持:返回的数组会被放回原始后端的同一设备,避免不必要的跨设备拷贝。

  • 多输出统一:无论函数返回单个数组还是序列,外层均返回 tuple[Array, …],便于后续统一处理。

生活类比:这个包装器就像一个多功能厨房助手——它会先检查食材是否需要预处理(比如解冻成室温食材,对应 as_numpy),记住食材来自哪个储存区(设备保持),然后按菜谱执行烹饪,最后无论做出一道菜还是一套餐,都统一装盘(多输出统一为元组)方便端上餐桌。


66.5 只读数组更新操作符 —— 函数式“点石成金”

66.5.1 at 类的声明与入口(完整路径)

# 第 66 章 —— File: sklearn/externals/array_api_extra/_lib/_at.py
class at:
    """
    Update operations for read‑only arrays.
    """
    _x: Array
    _idx: SetIndex | Undef
    __slots__: ClassVar[tuple[str, ...]] = ("_idx", "_x")

    def __init__(self, x: Array, idx: SetIndex | Undef = _undef, /) -> None:
        self._x = x
        self._idx = idx

    def __getitem__(self, idx: SetIndex, /) -> Self:
        """支持 ``at(x)[slice]`` 语法。"""
        if self._idx is not _undef:
            raise ValueError("Index has already been set")
        return type(self)(self._x, idx)
  • at(x)[idx].set(v)链式调用,与 JAX x.at[idx].set(v) 同义。

  • Undef 用作“未设置索引”的哨兵,确保只能调用一次 __getitem__

生活类比at 就像一把多功能厨房刀——无论你是切菜(set)、加盐(add)还是反过来削皮(subtract),它都提供一个统一的握把(at 对象),让你在只读菜板(如 JAX 数组)上也能安全操作,而不会弄伤自己或损坏菜板。

66.5.2 核心内部实现 _op

# 第 66 章 —— File: sklearn/externals/array_api_extra/_lib/_at.py
def _op(
    self,
    at_op: _AtOp,
    in_place_op: Callable[[Array, Array | complex], Array] | None,
    out_of_place_op: Callable[[Array, Array], Array] | None,
    y: Array | complex,
    /,
    copy: bool | None,
    xp: ModuleType | None,
) -> Array:
    """
    Implement all update operations.
    """
    from ._funcs import apply_where  # 只在特殊布尔掩码情况下使用

    x, idx = self._x, self._idx
    xp = array_namespace(x, y) if xp is None else xp

    # ---------- 参数合法性检查 ----------
    if isinstance(idx, Undef):
        raise ValueError(
            "Index has not been set.\n"
            "Usage: either\n"
            "    at(x, idx).set(value)\n"
            "or\n"
            "    at(x)[idx].set(value)\n"
            "(same for all other methods)."
        )
    if copy not in (True, False, None):
        raise ValueError(f"copy must be True, False, or None; got {copy!r}")

    # ---------- 写权限判定 ----------
    writeable = None if copy else is_writeable_array(x)

    # ---------- 特殊布尔掩码处理(Dask / JAX) ----------
    if (
        (is_dask_array(idx) or is_jax_array(idx))
        and idx.dtype == xp.bool
        and idx.shape == x.shape
    ):
        y_xp = xp.asarray(y, dtype=x.dtype, device=_compat.device(x))
        if y_xp.ndim == 0:
            if out_of_place_op:
                out = apply_where(
                    idx, (x, y_xp), out_of_place_op, fill_value=x, xp=xp
                )
                out = xp.astype(out, x.dtype, copy=False)
            else:
                out = xp.where(idx, y_xp, x)

            if copy is False:
                x[()] = out
                return x
            return out

    # ---------- JAX 原生 at[] ----------
    if is_jax_array(x):
        func = cast(
            Callable[[Array | complex], Array],
            getattr(x.at[idx], at_op.value),
        )
        out = func(y)
        return xp.astype(out, x.dtype, copy=False)

    # ---------- 对不可写数组的复制路径 ----------
    if copy or (copy is None and not writeable):
        if is_jax_array(x):
            x = xp.asarray(x, copy=True)   # JAX 必须复制
        else:
            x = xp.asarray(x, copy=True)
        writeable = None

    # ---------- 再次确认写权限 ----------
    if writeable is None:
        writeable = is_writeable_array(x)
    if not writeable:
        raise ValueError(f"Can't update read‑only array {x}")

    # ---------- PyTorch dtype 兼容补丁 ----------
    if is_torch_array(y):
        y = xp.astype(y, x.dtype, copy=False)

    # ---------- 实际更新 ----------
    if in_place_op:
        x[idx] = in_place_op(x[idx], y)
    else:
        x[idx] = y
    return x

66.5.2.1 关键实现要点

| 步骤 | 目的 |

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

| 布尔掩码特殊路径 | 在 Dask / JAX 中,仅支持标量 y 的布尔掩码更新,使用 apply_where 以函数式方式避免就地写入。 |

| JAX 原生 at[] | 若输入是 JAX 数组,直接调用 x.at[idx].<op>(y),保持函数式语义并利用 JIT。 |

| 不可写数组复制 (copy=True / copy=None 且不可写) | 对只读数组(如只读 NumPy、JAX)先 copy=True 再执行更新,确保不会在原数组上产生副作用。 |

| PyTorch dtype 修正 | 兼容 PyTorch 对 __setitem__ 的 dtype 限制(只能接受 int64),提前转换。 |

| 统一返回 | 始终返回更新后的数组(可能是原数组的就地修改或复制体)。 |

使用建议永远在惰性后端上使用 x = at(x)[...]... 并立即 重新赋值,避免后续对已经被内部修改的对象产生意外共享。

生活类比:用 at 更新只读数组就像在玻璃菜板上切菜——你不能直接在玻璃上刻痕(就地修改),所以要么先换一块菜板(复制),要么用特殊技巧(如撒盐而不划痕,对应布尔掩码下的 apply_where)。记得用完立刻换下菜板(重新赋值),否则下次可能不知不觉在旧菜板上继续切菜。

66.5.3 常用公开方法(thin wrappers)

def set(self, y, /, copy=None, xp=None) -> Array:
    """x[idx] = y"""
    return self._op(_AtOp.SET, None, None, y, copy=copy, xp=xp)

def add(self, y, /, copy=None, xp=None) -> Array:
    """x[idx] += y"""
    return self._op(_AtOp.ADD, operator.iadd, operator.add, y, copy=copy, xp=xp)

def subtract(self, y, /, copy=None, xp=None) -> Array:
    """x[idx] -= y"""
    return self._op(_AtOp.SUBTRACT, operator.isub, operator.sub, y, copy=copy, xp=xp)

# 第 66 章 —— … 其余方法(multiply、divide、power、min、max)同理

每个方法仅把对应的枚举值和算子传给 _op,保持实现的 单一入口,易于维护。

生活类比:这些方法就像刀架上的不同刀具——切片刀(set)、削皮刀(add)、磨刀石(subtract)等等,虽然形状不同,但都插在同一个刀座(_op)上,换起来毫不费力。


66.6 函数委托机制 —— 后端优化的“智能路由器”

66.6.1 isclose 委托实现(完整路径)

# 第 66 章 —— File: sklearn/externals/array_api_extra/_delegation.py
def isclose(
    a: Array | complex,
    b: Array | complex,
    *,
    rtol: float = 1e-05,
    atol: float = 1e-08,
    equal_nan: bool = False,
    xp: ModuleType | None = None,
) -> Array:
    """
    Return a boolean array where two arrays are element‑wise equal within a tolerance.
    """
    xp = array_namespace(a, b) if xp is None else xp

    # ✅ 优先使用后端原生实现(如果可用)
    if (
        is_numpy_namespace(xp)
        or is_cupy_namespace(xp)
        or is_dask_namespace(xp)
        or is_jax_namespace(xp)
    ):
        return xp.isclose(a, b, rtol=rtol, atol=atol, equal_nan=equal_nan)

    # 🔧 PyTorch 需要先确保输入是 Array API 对象(2024.12 支持)
    if is_torch_namespace(xp):
        a, b = asarrays(a, b, xp=xp)
        return xp.isclose(a, b, rtol=rtol, atol=atol, equal_nan=equal_nan)

    # ⏬ 若后端不提供原生实现,回退到标准实现
    return _funcs.isclose(a, b, rtol=rtol, atol=atol, equal_nan=equal_nan, xp=xp)

66.6.1.1 设计要点

  • 快速路径:对已实现 isclose 的后端(NumPy、CuPy、Dask、JAX)直接调用,获得最佳性能。

  • 特例:PyTorch 缺少符合最新 Array API 规范的实现,需先使用 asarrays 将标量提升为数组。

  • 回退:仅当后端不支持时才使用 _funcs.isclose(纯 Python 实现),保证功能完整性。

一问一答

:如果后端是 torch,为何仍要走 asarrays

torch.isclose 需要 torch.Tensor,而 asarrays 能把标量包装为 Tensor 并统一 dtype,从而兼容 Array API 2024.12 的新特性。

其他委托函数nan_to_numone_hotpad)遵循相同模式:先尝试后端原生实现,若不存在则回退到 _funcs.py 中的通用实现。

生活类比:想象你在翻译中心遇到一个稀有方言(如 PyTorch 的旧版 API)——中心不会直接放弃翻译,而是先找当地向导(asarrays)把游客的话译成标准普通话(Array API 对象),再用中心的通用翻译手册(标准实现)完成翻译;而对于常见语言(如 NumPy),则直接调用当地速记员(后端原生实现)最快完成任务。


66.7 扩展数组操作函数 —— 标准之外的“数学百宝箱”

下面针对几个关键函数提供 完整路径 + 逐行注释,随后给出 功能概要,帮助快速定位实现细节。

66.7.1 1️⃣ apply_where

# 第 66 章 —— File: sklearn/externals/array_api_extra/_lib/_funcs.py
def apply_where(
    cond: Array,
    args: Array | tuple[Array, ...],
    f1: Callable[..., Array],
    f2: Callable[..., Array] | None = None,
    /,
    *,
    fill_value: Array | complex | None = None,
    xp: ModuleType | None = None,
) -> Array:
    """
    Run one of two elementwise functions depending on a condition.
    """
    # ---------- 参数合法性 ----------
    if (f2 is None) == (fill_value is None):
        raise TypeError("Exactly one of `fill_value` or `f2` must be given.")

    # 标准化 `args` 为列列表
    args_ = list(args) if isinstance(args, tuple) else [args]

    # 自动获取命名空间
    xp = array_namespace(cond, fill_value, *args_) if xp is None else xp

    # ---------- 广播 ---
    if isinstance(fill_value, int | float | complex | NoneType):
        cond, *args_ = xp.broadcast_arrays(cond, *args_)
    else:
        cond, fill_value, *args_ = xp.broadcast_arrays(cond, fill_value, *args_)

    # ---------- Dask 专属路径 ----------
    if is_dask_namespace(xp):
        meta_xp = meta_namespace(cond, fill_value, *args_, xp=xp)
        # `map_blocks` 会把函数分别作用于每个块
        return xp.map_blocks(_apply_where, cond, f1, f2, fill_value, *args_, xp=meta_xp)

    # ---------- Eager 路径 ----------
    return _apply_where(cond, f1, f2, fill_value, *args_, xp=xp)

作用:仅在满足 cond 时求值 f1(或 fill_value),否则求值 f2;在 Dask 中使用 map_blocks 保持惰性计算。

生活类比:这就像一个智能分配器——根据条件(比如天气)决定是去室内健身房(f1)还是户外跑步(f2fill_value)。在分布式系统(Dask)中,它不会等所有人都到齐才决定,而是让每个健身房或跑步路径(数据块)自行根据当地天气做决定,最后再汇总结果。


66.7.2 2️⃣ broadcast_shapes

# 第 66 章 —— File: sklearn/externals/array_api_extra/_lib/_funcs.py
def broadcast_shapes(*shapes: tuple[float | None, ...]) -> tuple[int | None, ...]:
    """
    Compute the shape of the broadcasted arrays.
    """
    if not shapes:
        return ()

    ndim = max(len(shape) for shape in shapes)
    out: list[int | None] = []
    for axis in range(-ndim, 0):
        # 收集该轴上所有尺寸(可能为 int、None、math.nan)
        sizes = {shape[axis] for shape in shapes if axis >= -len(shape)}
        # 若出现 `None` 或 `math.nan`,该轴尺寸未知
        none_size = None in sizes or math.nan in sizes
        # 去除 “1” 与 “未知” 只保留真正的尺寸
        sizes -= {1, None, math.nan}
        if len(sizes) > 1:
            raise ValueError(
                "shape mismatch: objects cannot be broadcast to a single shape: "
                f"{shapes}."
            )
        # 输出使用 `None` 统一表示未知尺寸
        out.append(None if none_size else cast(int, sizes.pop()) if sizes else 1)

    return tuple(out)

要点:兼容 Array API (None) 与 Dask (math.nan) 两种未知维度表示,输出始终使用 None 符合标准。

生活类比:这就像计算拼图的最终尺寸——有些拼图块可能边缘磨损导致尺寸不确定(用 Nonemath.nan 表示),但只要没有真正的冲突(比如一块要求 3cm,另一块坚持 5cm 在同一位置),我们就能算出一个统一的框架尺寸,所有不确定的地方用 None 标记出来。


66.7.3 3️⃣ cov(协方差矩阵)

# 第 66 章 —— File: sklearn/externals/array_api_extra/_lib/_funcs.py
def cov(m: Array, /, *, xp: ModuleType | None = None) -> Array:
    """
    Estimate a covariance matrix.
    """
    if xp is None:
        xp = array_namespace(m)

    # 复制并提升至 float64(若原始为整数则转为 float64)
    m = xp.asarray(m, copy=True)
    dtype = xp.float64 if xp.isdtype(m.dtype, "integral") else xp.result_type(m, xp.float64)
    m = atleast_nd(m, ndim=2, xp=xp)     # 保证是二维
    m = xp.astype(m, dtype)

    # 均值沿每一行(变量)计算
    avg = _helpers.mean(m, axis=1, xp=xp)

    # 计算自由度(样本数 - 1)
    fact = m.shape[1] - 1
    if fact <= 0:
        warnings.warn("Degrees of freedom <= 0 for slice", RuntimeWarning, stacklevel=2)
        fact = 0

    # 去均值
    m -= avg[:, None]
    m_T = m.T
    # 若复数,使用共轭转置
    if xp.isdtype(m_T.dtype, "complex floating"):
        m_T = xp.conj(m_T)

    # 矩阵乘法得到协方差,随后除以自由度
    c = m @ m_T
    c /= fact

    # 去掉多余的单维度
    axes = tuple(axis for axis, length in enumerate(c.shape) if length == 1)
    return xp.squeeze(c, axis=axes)

细节:对 复数 使用共轭转置,确保得到真实的协方差矩阵;自由度为零时会发出警告并防止除零错误。

生活类比:想象你在研究多个股票(变量)随时间的价格(观察)如何共同变动。cov 就像计算它们之间的“同步指数”——先把所有价格转换为相同的单位(如百分比变化,对应 dtype 提升),然后扣除每只股票的平均价格(去均值),最后看它们偏离平均值的程度如何相互关联(矩阵乘法)。如果某只股票只有一个数据点(自由度为零),那就没法判断它是否真的在变动,所以发个警告并跳过。


66.7.4 4️⃣ kron(克罗内克积)

# 第 66 章 —— File: sklearn/externals/array_api_extra/_lib/_funcs.py
def kron(a: Array | complex, b: Array | complex, /, *, xp: ModuleType | None = None) -> Array:
    """
    Kronecker product of two arrays.
    """
    if xp is None:
        xp = array_namespace(a, b)
    a, b = asarrays(a, b, xp=xp)

    # 对低维数组进行前置 1 维度填充,以保持维度相等
    a = cast(Array, xp.broadcast_to(a, (1,) * (b.ndim - a.ndim) + a.shape))

    # 通过插入空维度并使用广播实现块乘
    a_arr = expand_dims(a, axis=tuple(range(b.ndim - a.ndim)), xp=xp)
    b_arr = expand_dims(b, axis=tuple(range(a.ndim - b.ndim)), xp=xp)

    a_arr = expand_dims(a_arr, axis=tuple(range(1, max(a.ndim, b.ndim) * 2, 2)), xp=xp)
    b_arr = expand_dims(b_arr, axis=tuple(range(0, max(a.ndim, b.ndim) * 2, 2)), xp=xp)

    result = xp.multiply(a_arr, b_arr)
    res_shape = tuple(a_s * b_s for a_s, b_s in zip(a.shape, b.shape, strict=True))
    return xp.reshape(result, res_shape)

思路:先把两个数组统一到相同维数(前置 1),再交错插入新维度实现块乘,最后 reshape 为 Kronecker 积的目标形状。

生活类比:就像用图腾柱(Kronecker 积)象征两个部落的联盟——假设部落 A 有 2 个氏族,部落 B 有 3 个图腾,那么联盟就有 2×3=6 种组合。为保证每个氏族和图腾都能“对齐”,我们先给人数少的部落的图腾柱加些虚拟的底座(前置 1),然后把每个氏族的名字和每个图腾的花纹交错排列(插入空维度+广播),最后把这张巨图重新裁成标准尺寸(reshape)。


66.7.5 5️⃣ sinc(归一化 sinc)

# 第 66 章 —— File: sklearn/externals/array_api_extra/_lib/_funcs.py
def sinc(x: Array, /, *, xp: ModuleType | None = None) -> Array:
    """
    Return the normalized sinc function.
    """
    if xp is None:
        xp = array_namespace(x)

    if not xp.isdtype(x.dtype, "real floating"):
        raise ValueError("`x` must have a real floating data type.")

    # 防止除以 0:在 `x == 0` 位置使用 eps 代替
    y = xp.pi * xp.where(
        xp.astype(x, xp.bool),
        x,
        xp.asarray(xp.finfo(x.dtype).eps, dtype=x.dtype, device=_compat.device(x)),
    )
    return xp.sin(y) / y

技巧:在 x==0 时使用机器 epsilon 代替 0,避免除以 0 的 NaN。

生活类比:sinc 函数就像测量一个完美音频信号在零点附近的行为——直接算 sin(πx)/(πx) 在 x=0 时会除以零,就像用量尺量零长度的物体一样荒谬。所以我们用一个极其微小但非零的刻度(机器 epsilon)来替代零点,这样既能避免无穷大,又能在足够小时近似真实值(极限为 1)。


66.8 测试工具与后端枚举 —— 跨后端的“质量检测站”

66.8.1 断言工具 xp_assert_close

# 第 66 章 —— File: sklearn/externals/array_api_extra/_lib/_testing.py
def xp_assert_close(
    actual: Array,
    desired: Array,
    *,
    rtol: float | None = None,
    atol: float = 0,
    err_msg: str = "",
    check_dtype: bool = True,
    check_shape: bool = True,
    check_scalar: bool = False,
) -> None:
    """
    Array‑API compatible version of `np.testing.assert_allclose`.
    """
    # ① 检查命名空间、形状、dtype 是否匹配
    xp = _check_ns_shape_dtype(actual, desired, check_dtype, check_shape, check_scalar)

    # ② 对于不可 materialize(如 meta‑tensor)直接跳过
    if not _is_materializable(actual):
        return

    # ③ 如果未指定 rtol,根据 dtype 自动推导一个合理的容忍度
    if rtol is None:
        if xp.isdtype(actual.dtype, ("real floating", "complex floating")):
            rtol = xp.finfo(actual.dtype).eps ** 0.5 * 4
        else:
            rtol = 1e-7

    # ④ 把两边统一转换为 NumPy(统一后端)进行比较
    actual_np = as_numpy_array(actual, xp=xp)
    desired_np = as_numpy_array(desired, xp=xp)
    np.testing.assert_allclose(actual_np, desired_np, rtol=rtol, atol=atol, err_msg=err_msg)
  • 统一转 NumPyas_numpy_array 会根据后端类型(CuPy、PyTorch、JAX、Sparse)选择合适的转化路径,规避 GPU‑CPU transfer guard 与稀疏阵列的 densify。

  • 自动容忍度:对实数/复数 dtype 使用 eps**0.5 * 4,比默认 1e-7 更贴合数值误差范围。

生活类比:就像在国际翻译中心审校译作——不管原文是用哪种语言写的(后端类型),审校员都会先把译文统一转换成一种通用语言(NumPy)来比较;如果译文还是某种只有作者才能读的草稿(meta‑tensor),那就直接跳过审校;容忍度则根据语言的精细度调整——比如德语的长复合词可能需要更宽松的容忍度,而日语的敬语则需要更精准。


66.8.2 lazy_xp_functionpatch_lazy_xp_functions

# 第 66 章 —— File: sklearn/externals/array_api_extra/testing.py
def lazy_xp_function(
    func: Callable[..., Any],
    *,
    allow_dask_compute: bool | int = False,
    jax_jit: bool = True,
    static_argnums: Deprecated = DEPRECATED,
    static_argnames: Deprecated = DEPRECATED,
) -> None:
    """
    Tag a function to be tested on lazy backends.
    """
    # 记录元数据
    tags = {"allow_dask_compute": allow_dask_compute, "jax_jit": jax_jit}
    try:
        func._lazy_xp_function = tags   # type: ignore[attr-defined]
    except AttributeError:               # 对于 Cython ufunc
        _ufuncs_tags[func] = tags
def patch_lazy_xp_functions(request, *, xp):
    """
    Test lazy execution of functions tagged with :func:`lazy_xp_function`.
    """
    # …(遍历模块 & 标记的函数)…

    if is_dask_namespace(xp):
        wrapped = _dask_wrap(func, n)   # 限制 compute 调用次数
    elif is_jax_namespace(xp):
        if tags["jax_jit"]:
            wrapped = jax_autojit(func)   # 用 jax_autojit 包装
    # 用 monkey‑patch 替换原函数

工作流程

  1. 标记 lazy_xp_function → 在函数对象上挂载元数据。
  1. 测试时 patch_lazy_xp_functions 读取元数据 → 根据后端(Dask / JAX)进行自动包装。
  1. Dask:使用 CountingDaskScheduler 限制 compute 调用次数,以检测是否不小心 materialize。
  1. JAX:使用 jax_autojit 实现“智能 JIT”,把标量参数视为 static,保持函数式。

生活类比:想象翻译中心要测试新翻译员(函数)是否真正掌握了语言而不仅是死记硬背。lazy_xp_function 就像在翻译员的档案上贴上标签:“此人需在实战中验证”;“patch_lazy_xp_functions”则是考试安排——如果考试用的是即时口译(Dask),就限制翻译员只能查一定次数的词典(防止作弊);如果是同声传译(JAX JIT),则提供一个智能耳机,让非专业术语(标量)自动过滤,只专注于核心内容。


66.8.3 后端枚举 Backend

# 第 66 章 —— File: sklearn/externals/array_api_extra/_lib/_backends.py
class Backend(Enum):
    # …(枚举定义略)…

    def pytest_param(self) -> Any:
        """
        Backend as a pytest parameter
        """
        id_ = self.name.lower().replace("_gpu", ":gpu").replace("_readonly", ":readonly")
        marks = []
        if self.like(Backend.ARRAY_API_STRICT):
            marks.append(pytest.mark.skipif(NUMPY_VERSION < (1, 26), reason="..."))
        if self.like(Backend.DASK, Backend.JAX):
            marks.append(pytest.mark.thread_unsafe)   # 因为 lazy_xp_function 进行 monkey‑patch
        return pytest.param(self, id=id_, marks=marks)
  • thread_unsafe:标记 Dask 与 JAX 为非线程安全,以配合并行测试框架(pytest‑run‑parallel)。

  • NUMPY_VERSION:在低于 1.26 时跳过 array_api_strict 的测试。

生活类比:就像翻译中心的语言分类表——每种语言(后端)都有自己的编号和备注。有些语言在低版本工具下测试不稳定(如 NumPy <1.26 对 array_api_strict 的支持),所以会自动跳过;有些语言在测试时需要特殊环境(如 Dask 和 JAX 的并发测试不安全),所以会标记上“请勿并行”。


66.9 小结表(前置说明句)

在本节我们回顾了本章节涵盖的核心概念与实现细节,帮助你快速定位并使用 array-api-extra 提供的功能。

| 概念 | 解释 |

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

| array_namespace() | 自动识别输入数组所属后端的入口函数。 |

| is_*_namespace() 系列 | 7 种后端类型检测函数,支撑后端感知的分派逻辑。 |

| _aliases.py(未展开) | 统一 50+ 标准函数的别名,处理 devicecopy 等差异。 |

| _linalg.py_fft.py | 标准化线性代数、傅里叶变换接口,确保跨后端数值一致。 |

| 类型提升表修正 | 通过 _fix_promotion_table/_wrap_func 修正后端之间的 dtype 提升差异。 |

| Array/DType/Device 协议 | 在 _typing.py 中定义协议,支撑静态类型检查。 |

| lazy_apply | 为 Dask/JAX 等惰性后端提供延迟函数执行能力。 |

| at | 类似 JAX 的 .at[] 更新语法,支持只读数组的函数式更新。 |

| _delegation.py | 根据后端智能路由到原生实现或标准实现。 |

| _funcs.py | 提供 apply_wherebroadcast_shapescovkronsinc 等增强函数。 |

| _testing.py / testing.py | 跨后端断言、惰性后端测试装饰器、后端枚举,确保多后端行为一致性。 |

生活类比:就像翻译中心的运营手册——它规定了如何识别语言(array_namespace)、如何快速调用当地专家(后端原生实现)、如何处理罕见方言(标准实现备用)、如何测试翻译员的真实水平(lazy_xp_function),以及如何确保所有译文风格统一(类型协议和 _aliases)。


66.10 下一章预告

在下一章,我们将转向 文档系统架构——探讨 scikit‑learn 如何使用 Sphinx、API 引用生成、自定义扩展与前端脚本,构建一个可扩展、可维护的大型项目文档体系。

66.11 生活类比

想象 array-api-compat 是一座跨语言通用的“通用翻译官中枢”Array API 标准 = 统一的外交礼仪规范(握手方式、会议流程、文件格式) 各后端 (NumPy/CuPy/PyTorch/JAX/Dask...) = 不同国家的政府部门(有的用美元、有的用欧元、有的用比特币、有的用信用积分) array_namespace() 识别机制 = 智能护照识别系统:一眼识别来访者国籍(后端类型),自动切换对应翻译模式 通用别名层 (_aliases.py) = 标准化外交备忘录模板:统一各国不同格式的公文(函数签名)为标准格式,处理货币单位(device参数)、复印件权限(copy语义)等差异 线性代数/FFT 封装 (_linalg.py/_fft.py) = 专业技术协议翻译组:将核心数学运算(SVD、QR、FFT等)翻译成各国通用的技术标准,确保计算结果精度一致 后端专属适配器 = 各国专属联络官:深谙本国法规(PyTorch类型提升规则、Dask惰性计算、JAX JIT限制),处理标准模板覆盖不到的特殊情况 类型系统与协议 (_typing.py) = 国际通用的法律条文定义:定义什么是有效签名(Array协议)、什么是合法货币(DType协议)、什么是主权领土(Device协议),支撑静态检查与运行时验证 而 array-api-extra 则像是这座翻译官的增强套件 —— 惰性求值 (lazy_apply) = 延迟生产线的智能调度:不立即消耗资源,等到真正需要时才执行计算(适用于 Dask、JAX 等惰性后端) 更新操作 (at) = 精准零件更换机器:在不可变产品线上实现“就地修改”效果,如同在封装货箱上打标签而不拆箱 函数委托 (_delegation.py) = 本地专家转介系统:当标准翻译无法处理某些方言时,自动将任务转交给后端原生实现(如 JAX 的 one_hot、PyTorch 的 pad) 扩展函数 (_funcs.py) = 通用工具箱扩展:额外提供如 Kronecker 积、sinc 函数、协方差矩阵等高级工具,弥补标准库的不足 测试工具 (_testing.py/testing.py) = 跨国标准检验局:提供统一的断言工具、惰性后端测试装饰器,确保各后端行为一致性

66.12 动手练习

66.12.1 阅读后端识别与命名空间机制

阅读 sklearn/externals/array_api_compat/common/__init__.py_helpers.py

  1. array_namespace() 如何通过输入数组的 __array_namespace__ 方法或类型特征识别后端?

  2. is_dask_namespace() 等检测函数的判断依据是什么?如何避免误判?

  3. _get_namespace() 内部如何处理多数组输入时的命名空间一致性检查?

66.12.2 深入通用别名与函数包装机制

阅读 sklearn/externals/array_api_compat/common/_aliases.py

  1. _wrap_func() 如何统一处理 devicecopy 等标准参数与后端特有参数的差异?

  2. astype() 等函数在不同后端的包装逻辑有何不同?以 NumPy 与 PyTorch 为例对比。

  3. 类型提升表 _fix_promotion_table() 如何修正跨后端的类型推导不一致问题?

66.12.3 探究线性代数与 FFT 封装的一致性保障

阅读 sklearn/externals/array_api_compat/common/_linalg.py_fft.py

  1. svd() 等线性代数函数如何通过 _linalg_func() 包装器统一返回值格式(如 Vh vs V.T)?

  2. fft() 系列函数如何处理不同后端对归一化参数(norm)的支持差异?

  3. 对比 NumPy 与 CuPy 后端的 linalg.py 实现,核心计算调用有何异同?

66.12.4 剖析后端专属适配器的差异化处理

阅读 sklearn/externals/array_api_compat/torch/_aliases.pydask/array/_aliases.py

  1. PyTorch 适配器中 _fix_torch_promotion() 如何解决 PyTorch 类型提升规则与标准的冲突?

  2. Dask 适配器中 copy() 语义为何默认为惰性?astype() 如何处理分块数组的类型转换?

  3. torch/_info.pycapabilities() 如何针对 meta device 禁用布尔索引与数据依赖形状?

66.12.5 探索 array-api-extra 的惰性求值与更新操作

阅读 sklearn/externals/array_api_extra/_lib/_lazy.py

  1. lazy_apply() 如何区分 eager 和 lazy 后端(如 Dask、JAX)的执行策略?

  2. 在 JAX 的 jax.jit 上下文中,为什么需要 _is_jax_jit_enabled() 检查?

  3. _lazy_apply_wrapper 如何处理 as_numpy 参数和多输出情况?

66.12.6 理解 array-api-extra 的更新操作 (at)

阅读 sklearn/externals/array_api_extra/_lib/_at.py

  1. at 类如何实现类似 JAX 的 .at[].set() 链式调用?

  2. 对于只读数组(如 JAX 在 jit 中),at.add() 如何回退到函数式更新?

  3. 为什么文档警告说在懒惰后端上应避免重用输入数组?请解释其中的副本语义与引用风险。

66.12.7 分析函数委托与扩展函数的实现

阅读 sklearn/externals/array_api_extra/_delegation.py_lib/_funcs.py

  1. isclose() 在不同后端(NumPy、CuPy、PyTorch 等)上的委托逻辑有何不同?

  2. _funcs.py 中的 broadcast_shapes() 如何处理 Dask 的 NaN 形状与 Array API 的 None 形状?

  3. cov() 函数在处理复数输入时,是如何分离实部和虚部进行计算的?

66.12.8 实践 array-api-extra 的测试工具与后端枚举

阅读 sklearn/externals/array_api_extra/_lib/_testing.pytesting.py

  1. xp_assert_close() 如何统一处理不同后端(包括 GPU 上的 PyTorch/JAX)数组到 NumPy 的转换?

  2. lazy_xp_function 装饰器如何配合 patch_lazy_xp_functions 实现 Dask/JAX 的测试隔离?

  3. Backend 枚举的 pytest_param() 如何为不同后端自动添加 thread_unsafe 等标记?

66.13 架构与数据流图

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

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

66.14 源码地图(带简要说明)

sklearn/externals/array_api_compat/common/__init__.py
├── array_namespace()                    # 核心后端识别入口
├── is_numpy_namespace()                 # NumPy 后端检测
├── is_cupy_namespace()                  # CuPy 后端检测
├── is_torch_namespace()                 # PyTorch 后端检测
├── is_jax_namespace()                   # JAX 后端检测
├── is_dask_namespace()                  # Dask 后端检测
├── is_array_api_strict_namespace()      # array-api-strict 后端检测
├── is_pydata_sparse_namespace()         # PyData Sparse 后端检测
├── is_writeable_array()                 # 可写数组检测
├── is_cupy_array()                      # CuPy 数组类型检测
├── is_dask_array()                      # Dask 数组类型检测
├── is_jax_array()                       # JAX 数组类型检测
├── is_numpy_array()                     # NumPy 数组类型检测
├── is_pydata_sparse_array()             # PyData Sparse 数组类型检测
├── is_torch_array()                     # PyTorch 数组类型检测
├── is_lazy_array()                      # 惰性数组检测
├── is_array_api_obj()                   # 标准数组对象检测
├── device()                             # 设备获取
├── size()                               # 数组大小获取
├── to_device()                          # 设备迁移
└── array_namespace.__module__           # 命名空间模块属性
sklearn/externals/array_api_compat/common/_helpers.py
├── _get_namespace()                     # 内部命名空间解析
├── _check_api_version()                 # API 版本兼容性检查
└── _parse_api_version()                 # 版本字符串解析
sklearn/externals/array_api_compat/__init__.py
├── array_namespace()                    # 公共导出:后端识别主函数
├── is_numpy_namespace()                 # 公共导出
├── is_cupy_namespace()                  # 公共导出
├── is_torch_namespace()                 # 公共导出
├── is_jax_namespace()                   # 公共导出
├── is_dask_namespace()                  # 公共导出
└── ...                                  # 其余检测函数公共导出
sklearn/externals/array_api_compat/_internal.py
├── _NamespaceInfo.__init__()            # 命名空间信息基类
├── _NamespaceInfo.capabilities()        # 能力查询接口
├── _NamespaceInfo.default_dtypes()      # 默认 dtype 查询
└── _get_namespace_info()                # 获取命名空间信息实例
sklearn/externals/array_api_compat/common/_aliases.py
├── _fix_promotion_table()               # 类型提升表修正
├── _wrap_func()                         # 函数包装器:处理 device/copy 等参数差异
├── asarray()                            # 统一数组创建接口
├── empty()                              # 统一空数组创建
├── zeros()                              # 统一零数组创建
├── ones()                               # 统一一数组创建
├── full()                               # 统一填充数组创建
├── arange()                             # 统一等差数列创建
├── linspace()                           # 统一线性间距创建
├── eye()                                # 统一单位矩阵创建
├── from_dlpack()                        # 统一 DLPack 互操作
├── reshape()                            # 统一形状变换
├── broadcast_to()                       # 统一广播
├── concat()                             # 统一拼接
├── stack()                              # 统一堆叠
├── moveaxis()                           # 统一轴移动
├── transpose()                          # 统一转置
├── flip()                               # 统一翻转
├── sort()                               # 统一排序
├── argsort()                            # 统一索引排序
├── unique_values()                      # 统一唯一值
├── unique_inverse()                     # 统一唯一值逆索引
├── unique_counts()                      # 统一唯一值计数
├── abs()                                # 统一绝对值
├── sqrt()                               # 统一平方根
├── sin()                                # 统一正弦
├── cos()                                # 统一余弦
├── tan()                                # 统一正切
├── asin()                               # 统一反正弦
├── acos()                               # 统一反余弦
├── atan()                               # 统一反正切
├── exp()                                # 统一指数
├── log()                                # 统一对数
├── sum()                                # 统一求和
├── mean()                               # 统一均值
├── var()                                # 统一方差
├── std()                                # 统一标准差
├── matmul()                             # 统一矩阵乘法
├── tensordot()                          # 统一张量点积
└── ...                                  # 更多统一别名函数
sklearn/externals/array_api_compat/common/_linalg.py
├── matmul()                             # 矩阵乘法统一接口
├── tensordot()                          # 张量点积统一接口
├── vecdot()                             # 向量点积统一接口
├── svd()                                # 奇异值分解统一接口
├── qr()                                 # QR 分解统一接口
├── cholesky()                           # Cholesky 分解统一接口
├── eigh()                               # 对称/厄米特矩阵特征值分解
├── eigvalsh()                           # 对称/厄米特矩阵特征值
├── solve()                              # 线性方程组求解
├── inv()                                # 矩阵求逆
├── det()                                # 行列式计算
├── matrix_norm()                        # 矩阵范数
├── vector_norm()                        # 向量范数
├── diagonal()                           # 对角线提取
├── trace()                              # 迹计算
└── _linalg_func()                       # 内部线性代数函数包装器
sklearn/externals/array_api_compat/common/_fft.py
├── fft()                                # 快速傅里叶变换统一接口
├── ifft()                               # 逆快速傅里叶变换
├── rfft()                               # 实数快速傅里叶变换
├── irfft()                              # 实数逆快速傅里叶变换
├── fft2()                               # 2D 快速傅里叶变换
├── ifft2()                              # 2D 逆快速傅里叶变换
├── rfft2()                              # 2D 实数快速傅里叶变换
├── irfft2()                             # 2D 实数逆快速傅里叶变换
├── fftn()                               # N维快速傅里叶变换
├── ifftn()                              # N维逆快速傅里叶变换
├── rfftn()                              # N维实数快速傅里叶变换
├── irfftn()                             # N维实数逆快速傅里叶变换
└── _fft_func()                          # 内部 FFT 函数包装器
sklearn/externals/array_api_compat/numpy/__init__.py
├── array_namespace()                    # NumPy 命名空间获取
├── _info                               # NumPy 命名空间信息实例
├── _aliases.*                           # 导入通用别名
├── _linalg.*                            # 导入线性代数接口
└── _fft.*                               # 导入 FFT 接口
sklearn/externals/array_api_compat/numpy/_aliases.py
├── _numpy_promotion_table()             # NumPy 类型提升表
├── _wrap_numpy_func()                   # NumPy 函数包装
├── copy()                               # NumPy copy 语义处理
├── astype()                             # NumPy 类型转换
└── ...                                  # NumPy 特有别名调整
sklearn/externals/array_api_compat/numpy/_info.py
├── _NumPyNamespaceInfo.__init__()       # NumPy 命名空间信息类
├── _NumPyNamespaceInfo.capabilities()   # NumPy 能力集
├── _NumPyNamespaceInfo.default_dtypes() # NumPy 默认 dtype
└── _numpy_namespace_info                # NumPy 信息单例
sklearn/externals/array_api_compat/numpy/_typing.py
├── Array 协议定义                       # NumPy 数组类型协议
├── DType 协议定义                       # NumPy dtype 协议
├── Device 协议定义                      # NumPy 设备协议
└── ...                                  # 索引类型别名
sklearn/externals/array_api_compat/numpy/linalg.py
├── matmul()                             # NumPy 矩阵乘法实现
├── svd()                                # NumPy SVD 实现
├── qr()                                 # NumPy QR 实现
├── cholesky()                           # NumPy Cholesky 实现
├── eigh()                               # NumPy eigh 实现
├── eigvalsh()                           # NumPy eigvalsh 实现
├── solve()                              # NumPy solve 实现
├── inv()                                # NumPy inv 实现
├── det()                                # NumPy det 实现
├── matrix_norm()                        # NumPy 矩阵范数实现
├── vector_norm()                        # NumPy 向量范数实现
├── diagonal()                           # NumPy diagonal 实现
└── trace()                              # NumPy trace 实现
sklearn/externals/array_api_compat/numpy/fft.py
├── fft()                                # NumPy FFT 实现
├── ifft()                               # NumPy IFFT 实现
├── rfft()                               # NumPy RFFT 实现
├── irfft()                              # NumPy IRFFT 实现
├── fft2()                               # NumPy FFT2 实现
├── ifft2()                              # NumPy IFFT2 实现
├── rfft2()                              # NumPy RFFT2 实现
├── irfft2()                             # NumPy IRFFT2 实现
├── fftn()                               # NumPy FFTN 实现
├── ifftn()                              # NumPy IFFTN 实现
├── rfftn()                              # NumPy RFFTN 实现
└── irfftn()                             # NumPy IRFFTN 实现
sklearn/externals/array_api_compat/cupy/__init__.py
├── array_namespace()                    # CuPy 命名空间获取
├── _info                               # CuPy 命名空间信息
├── _aliases.*                           # 导入通用别名
├── _linalg.*                            # 导入线性代数接口
└── _fft.*                               # 导入 FFT 接口
sklearn/externals/array_api_compat/cupy/_aliases.py
├── _cupy_promotion_table()              # CuPy 类型提升表
├── _wrap_cupy_func()                    # CuPy 函数包装
├── copy()                               # CuPy copy 语义
├── astype()                             # CuPy 类型转换
└── ...                                  # CuPy 特有别名调整
sklearn/externals/array_api_compat/cupy/_info.py
├── _CuPyNamespaceInfo.__init__()        # CuPy 命名空间信息类
├── _CuPyNamespaceInfo.capabilities()    # CuPy 能力集
├── _CuPyNamespaceInfo.default_dtypes()  # CuPy 默认 dtype
└── _cupy_namespace_info                 # CuPy 信息单例
sklearn/externals/array_api_compat/cupy/_typing.py
├── Array 协议定义                       # CuPy 数组类型协议
├── DType 协议定义                       # CuPy dtype 协议
├── Device 协议定义                      # CuPy 设备协议
└── ...                                  # 索引类型别名
sklearn/externals/array_api_compat/cupy/linalg.py
├── matmul()                             # CuPy 矩阵乘法实现
├── svd()                                # CuPy SVD 实现
├── qr()                                 # CuPy QR 实现
├── cholesky()                           # CuPy Cholesky 实现
├── eigh()                               # CuPy eigh 实现
├── eigvalsh()                           # CuPy eigvalsh 实现
├── solve()                              # CuPy solve 实现
├── inv()                                # CuPy inv 实现
├── det()                                # CuPy det 实现
├── matrix_norm()                        # CuPy 矩阵范数实现
├── vector_norm()                        # CuPy 向量范数实现
├── diagonal()                           # CuPy diagonal 实现
└── trace()                              # CuPy trace 实现
sklearn/externals/array_api_compat/cupy/fft.py
├── fft()                                # CuPy FFT 实现
├── ifft()                               # CuPy IFFT 实实现
├── rfft()                               # CuPy RFFT 实现
├── irfft()                              # CuPy IRFFT 实现
├── fft2()                               # CuPy FFT2 实现
├── ifft2()                              # CuPy IFFT2 实现
├── rfft2()                              # CuPy RFFT2 实现
├── irfft2()                             # CuPy IRFFT2 实现
├── fftn()                               # CuPy FFTN 实现
├── ifftn()                              # CuPy IFFTN 实现
├── rfftn()                              # CuPy RFFTN 实现
└── irfftn()                             # CuPy IRFFTN 实现
sklearn/externals/array_api_compat/torch/__init__.py
├── array_namespace()                    # PyTorch 命名空间获取
├── _info                               # PyTorch 命名空间信息
├── _aliases.*                           # 导入通用别名
├── _linalg.*                            # 导入线性代数接口
└── _fft.*                               # 导入 FFT 接口
sklearn/externals/array_api_compat/torch/_aliases.py
├── _torch_promotion_table()             # PyTorch 类型提升表
├── _wrap_torch_func()                   # PyTorch 函数包装
├── copy()                               # PyTorch copy 语义
├── astype()                             # PyTorch 类型转换
├── _fix_torch_promotion()               # PyTorch 类型提升修正
└── ...                                  # PyTorch 特有别名调整
sklearn/externals/array_api_compat/torch/_info.py
├── _TorchNamespaceInfo.__init__()       # PyTorch 命名空间信息类
├── _TorchNamespaceInfo.capabilities()   # PyTorch 能力集
├── _TorchNamespaceInfo.default_dtypes() # PyTorch 默认 dtype
└── _torch_namespace_info                # PyTorch 信息单例
sklearn/externals/array_api_compat/torch/_typing.py
├── Array 协议定义                       # PyTorch 张量类型协议
├── DType 协议定义                       # PyTorch dtype 协议
├── Device 协议定义                      # PyTorch 设备协议
└── ...                                  # 索引类型别名
sklearn/externals/array_api_compat/torch/linalg.py
├── matmul()                             # PyTorch 矩阵乘法实现
├── svd()                                # PyTorch SVD 实现
├── qr()                                 # PyTorch QR 实现
├── cholesky()                           # PyTorch Cholesky 实现
├── eigh()                               # PyTorch eigh 实现
├── eigvalsh()                           # PyTorch eigvalsh 实现
├── solve()                              # PyTorch solve 实现
├── inv()                                # PyTorch inv 实现
├── det()                                # PyTorch det 实现
├── matrix_norm()                        # PyTorch 矩阵范数实现
├── vector_norm()                        # PyTorch 向量范数实现
├── diagonal()                           # PyTorch diagonal 实现
└── trace()                              # PyTorch trace 实现
sklearn/externals/array_api_compat/torch/fft.py
├── fft()                                # PyTorch FFT 实现
├── ifft()                               # PyTorch IFFT 实现
├── rfft()                               # PyTorch RFFT 实现
├── irfft()                              # PyTorch IRFFT 实现
├── fft2()                               # PyTorch FFT2 实现
├── ifft2()                              # PyTorch IFFT2 实现
├── rfft2()                              # PyTorch RFFT2 实现
├── irfft2()                             # PyTorch IRFFT2 实现
├── fftn()                               # PyTorch FFTN 实现
├── ifftn()                              # PyTorch IFFTN 实现
├── rfftn()                              # PyTorch RFFTN 实现
└── irfftn()                             # PyTorch IRFFTN 实现
sklearn/externals/array_api_compat/dask/__init__.py
├── array_namespace()                    # Dask 命名空间获取
├── _info                               # Dask 命名空间信息
├── _aliases.*                           # 导入通用别名
├── _linalg.*                            # 导入线性代数接口
└── _fft.*                               # 导入 FFT 接口
sklearn/externals/array_api_compat/dask/array/__init__.py
├── array_namespace()                    # Dask Array 命名空间获取
├── _info                               # Dask Array 命名空间信息
├── _aliases.*                           # 导入通用别名
├── _linalg.*                            # 导入线性代数接口
└── _fft.*                               # 导入 FFT 接口
sklearn/externals/array_api_compat/dask/array/_aliases.py
├── _dask_promotion_table()              # Dask 类型提升表
├── _wrap_dask_func()                    # Dask 函数包装
├── copy()                               # Dask copy 语义(惰性)
├── astype()                             # Dask 类型转换
└── ...                                  # Dask 特有别名调整
sklearn/externals/array_api_compat/dask/array/_info.py
├── _DaskNamespaceInfo.__init__()        # Dask 命名空间信息类
├── _DaskNamespaceInfo.capabilities()    # Dask 能力集
├── _DaskNamespaceInfo.default_dtypes()  # Dask 默认 dtype
└── _dask_namespace_info                 # Dask 信息单例
sklearn/externals/array_api_compat/dask/array/linalg.py
├── matmul()                             # Dask 矩阵乘法实现
├── svd()                                # Dask SVD 实现
├── qr()                                 # Dask QR 实现
├── cholesky()                           # Dask Cholesky 实现
├── eigh()                               # Dask eigh 实现
├── eigvalsh()                           # Dask eigvalsh 实现
├── solve()                              # Dask solve 实现
├── inv()                                # Dask inv 实现
├── det()                                # Dask det 实现
├── matrix_norm()                        # Dask 矩阵范数实现
├── vector_norm()                        # Dask 向量范数实现
├── diagonal()                           # Dask diagonal 实现
└── trace()                              # Dask trace 实现
sklearn/externals/array_api_compat/dask/array/fft.py
├── fft()                                # Dask FFT 实现
├── ifft()                               # Dask IFFT 实现
├── rfft()                               # Dask RFFT 实现
├── irfft()                              # Dask IRFFT 实现
├── fft2()                               # Dask FFT2 实现
├── ifft2()                              # Dask IFFT2 实现
├── rfft2()                              # Dask RFFT2 实现
├── irfft2()                             # Dask IRFFT2 实现
├── fftn()                               # Dask FFTN 实现
├── ifftn()                              # Dask IFFTN 实现
├── rfftn()                              # Dask RFFTN 实现
└── irfftn()                             # Dask IRFFTN 实现
sklearn/externals/array_api_compat/common/_typing.py
├── Array 协议定义                       # 通用数组类型协议
├── DType 协议定义                       # 通用 dtype 协议
├── Device 协议定义                      # 通用设备协议
├── GetIndex 类型别名                    # 读取索引类型
└── SetIndex 类型别名                    # 写入索引类型
sklearn/externals/array_api_extra/__init__.py
├── __version__                          # 版本号
├── apply_where                          # 条件化函数应用
├── at                                   # 更新操作(类似 JAX 的 .at[].set())
├── atleast_nd                           # 扩展维度至至少 ndim
├── broadcast_shapes                     # 广播形状计算
├── cov                                  # 协方差矩阵估计
├── create_diagonal                      # 构建对角矩阵
├── default_dtype                        # 获取默认 dtype
├── expand_dims                          # 扩展数组维度
├── isclose                              # 容忍度相等比较
├── kron                                 # 克罗内克积
├── lazy_apply                           # 惰性函数应用(支持 Dask/JAX)
├── nan_to_num                           # 替换 NaN 和无穷大
├── nunique                              # 唯一元素计数
├── one_hot                              # 独热编码
├── pad                                  # 数组填充
├── setdiff1d                            # 集合差集
├── sinc                                 # 归一化 sinc 函数
└── ...                                  # 更多扩展函数
sklearn/externals/array_api_extra/_lib/__init__.py
└── __main__                             # 包初始化入口
sklearn/externals/array_api_extra/_lib/_at.py
├── at.__init__                          # 初始化 at 对象
├── at.__getitem__                       # 支持 at[x] 语法
├── at.set                               # 设置值:x[idx] = y
├── at.add                               # 加法更新:x[idx] += y
├── at.subtract                          # 减法更新:x[idx] -= y
├── at.multiply                          # 乘法更新:x[idx] *= y
├── at.divide                            # 除法更新:x[idx] /= y
├── at.power                             # 幂运算:x[idx] **= y
├── at.min                               # 最小值更新:x[idx] = min(x[idx], y)
├── at.max                               # 最大值更新:x[idx] = max(x[idx], y)
└── at._op                               # 内部操作实现
sklearn/externals/array_api_extra/_lib/_lazy.py
├── lazy_apply                           # 惰性函数应用主入口
├── _is_jax_jit_enabled                  # 检测是否在 jax.jit 内部
├── _lazy_apply_wrapper                  # 包装器:处理后端差异
├── jax_autojit                          # 自动 JIT 包装(处理非数组参数)
└── __main__                             # 模块测试入口
sklearn/externals/array_api_extra/_lib/_funcs.py
├── apply_where                          # 条件化函数应用
├── atleast_nd                           # 维度扩展
├── broadcast_shapes                     # 广播形状计算
├── cov                                  # 协方差矩阵
├── create_diagonal                      # 对角矩阵构造
├── expand_dims                          # 维度扩展
├── kron                                 # 克罗内克积
├── nunique                              # 唯一元素计数
├── pad                                  # 数组填充
├── setdiff1d                            # 集合差集
├── sinc                                 # 归一化 sinc 函数
└── ...                                  # 更多基础函数
sklearn/externals/array_api_extra/_lib/_testing.py
├── xp_assert_equal                      # 断言数组相等
├── xp_assert_close                      # 断言数组接近相等
├── xp_assert_less                       # 断言数组元素小于
├── as_numpy_array                       # 转换为 NumPy 数组(规避转移守护)
├── _check_ns_shape_dtype                # 检查命名空间、形状、 dtype
├── _is_materializable                   # 判断是否可 materialize
├── xfail                                # 标记预期失败(允许后续执行)
└── ...                                  # 更多测试工具
sklearn/externals/array_api_extra/testing.py
├── lazy_xp_function                     # 标记函数以在懒惰后端上测试
├── patch_lazy_xp_functions              # 应用懒惰后端测试补丁
├── CountingDaskScheduler                # 计算 Dask 调用次数的调度器
├── _dask_wrap                           # 内部 Dask 包装器
└── ...                                  # 更多测试辅助
sklearn/externals/array_api_extra/_delegation.py
├── isclose                              # 容忍度相等(委托至后端)
├── nan_to_num                           # 替换 NaN/Inf(委托至后端)
├── one_hot                              # 独热编码(委托至后端)
└── pad                                  # 数组填充(委托至后端)
sklearn/externals/array_api_extra/_lib/_utils/_compat.py
├── array_namespace                      # 从 array_api_compat 获取命名空间
├── device                               # 设备获取
├── is_array_api_obj                     # 标准数组对象判断
├── is_array_api_strict_namespace        # array-api-strict 检测
├── is_cupy_array                        # CuPy 数组类型检测
├── is_cupy_namespace                    # CuPy 命名空间检测
├── is_dask_array                        # Dask 数组类型检测
├── is_dask_namespace                    # Dask 命名空间检测
├── is_jax_array                         # JAX 数组类型检测
├── is_jax_namespace                     # JAX 命名空间检测
├── is_lazy_array                        # 惰性数组检测
├── is_numpy_array                       # NumPy 数组类型检测
├── is_numpy_namespace                   # NumPy 命名空间检测
├── is_pydata_sparse_array               # PyData Sparse 数组类型检测
├── is_pydata_sparse_namespace           # PyData Sparse 命名空间检测
├── is_torch_array                       # PyTorch 数组类型检测
├── is_torch_namespace                   # PyTorch 命名空间检测
├── is_writeable_array                   # 可写数组检测
├── size                                 # 数组大小获取
└── to_device                            # 设备迁移
sklearn/externals/array_api_extra/_lib/_utils/_helpers.py
├── in1d                                 # 元素是否在另一数组中
├── mean                                 # 复数均值计算
├── is_python_scalar                     # 判断是否为 Python 标量
├── asarrays                             # 确保输入为数组(标量转数组)
├── eager_shape                          # 获取非惰性形状
├── meta_namespace                       # 获取 Dask 分片的命名空间
├── capabilities                         # 获取后端能力(修正特殊情况)
├── pickle_flatten                       # 提取对象以便 pickle
├── pickle_unflatten                     # 反序列化对象
├── _AutoJITWrapper.__init__             # 初始化包装器
├── _AutoJITWrapper._register            # 注册 JAX PyTree 节点
├── _lazy_apply_wrapper                  # 包装器:处理 as_numpy、multi_output 等
└── ...                                  # 更多辅助函数
sklearn/externals/array_api_extra/_lib/_utils/_typing.py
├── Array                                # 数组协议定义
├── DType                                # dtype 协议定义
├── Device                               # 设备协议定义
├── GetIndex                             # 读取索引类型别名
└── SetIndex                             # 写入索引类型别名
sklearn/externals/array_api_extra/_lib/_utils/_typing.pyi
├── Array                                # 数组协议定义
├── DType                                # dtype 协议定义
├── Device                               # 设备协议定义
├── GetIndex                             # 读取索引类型别名
└── SetIndex                             # 写入索引类型别名
sklearn/externals/_array_api_compat_vendor.py
└── __main__                             # 供应商钩子入口
sklearn/externals/array_api_extra/_lib/_utils/_compat.pyi
├── array_namespace                      # 命名空间获取存根
├── device                               # 设备获取存根
├── is_array_api_obj                     # 标准数组对象判断存根
├── is_array_api_strict_namespace        # array-api-strict 检测存根
├── is_cupy_namespace                    # CuPy 命名空间检测存根
├── is_dask_namespace                    # Dask 命名空间检测存根
├── is_jax_namespace                     # JAX 命名空间检测存根
├── is_numpy_namespace                   # NumPy 命名空间检测存根
├── is_pydata_sparse_namespace           # PyData Sparse 命名空间检测存根
├── is_torch_namespace                   # PyTorch 命名空间检测存根
├── is_cupy_array                        # CuPy 数组类型检测存根
├── is_dask_array                        # Dask 数组类型检测存根
├── is_jax_array                         # JAX 数组类型检测存根
├── is_numpy_array                       # NumPy 数组类型检测存根
├── is_pydata_sparse_array               # PyData Sparse 数组类型检测存根
├── is_torch_array                       # PyTorch 数组类型检测存根
├── is_lazy_array                        # 惰性数组检测存根
├── is_writeable_array                   # 可写数组检测存根
├── size                                 # 数组大小获取存根
└── to_device                            # 设备迁移存根
sklearn/externals/array_api_extra/_lib/_backends.py
├── Backend                              # 后端枚举类
├── Backend.modname                      # 模块名属性
├── Backend.like                         # 后端相似性检查
├── Backend.pytest_param                 # pytest 参数化支持
└── NUMPY_VERSION                        # NumPy 版本元组
sklearn/externals/array_api_extra/_lib/_utils/__init__.py
└── __main__                             # 工具包初始化入口
posted @ 2026-09-04 04:07  绝不原创的飞龙  阅读(6)  评论(0)    收藏  举报