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_rank、pinv、matrix_norm、vector_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 对角线与迹
diagonal 和 trace 处理"最后两维 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.diagonal 与 xp.trace 默认作用于前两维,Array API 要求作用于最后两维,所以显式指定 axis1=-2, axis2=-1。svdvals 在 NumPy 中通过 svd(compute_uv=False) 实现。
65.11.5 FFT 精度保持
所有 FFT 函数都强制 float32 → complex64、float64 → 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 流程图展示各后端适配器的统一入口模式:
65.12.1 NumPy 后端:copy 语义映射
NumPy 1.x 没有 asarray 的 copy 参数,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_dtypes 与 dtypes 都严格检查 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 后端的类型模块定义了 Array 与 Device 类型别名:
源码路径: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 适配器的核心特点。asarray 用 with 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 + keepdims、axis` 为 tuple 等情况支持不完善 -
参数名差异:
dimvsaxis、dimsvsaxes、sizevsshape等 -
缺失功能补全:
take/take_along_axis负索引修正、expand_dims替代品、unique_*系列的 NaN 计数修复 -
特殊行为修复:
uint8在any/all后转bool、sign的 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 实现以支持 offset。vector_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 动手练习
-
阅读外部依赖打包与 ARFF 解析器
阅读
sklearn/externals/_array_api_compat_vendor.py与sklearn/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 矩阵的索引? -
文档字符串中
@param与Parameters章节的解析优先级如何?
-
-
探究 PEP 440 版本解析与比较逻辑
阅读
sklearn/externals/_packaging/version.py与_structures.py:-
追踪
Version.__init__如何解析 '1.0a1.post2.dev3' 等复杂版本字符串 -
理解
Version.__lt__如何实现预发布版 < 正式版 < 后发布版的排序规则 -
分析
NormalizedVersion与LegacyVersion的兼容策略
回答问题:
-
Version 类中
_version_regex正则如何捕获 epoch、release、pre、post、dev 五大组件? -
为什么 '1.0.dev0' < '1.0a0' < '1.0' < '1.0.post0'?
-
LegacyVersion 如何处理不符合 PEP 440 的传统版本号(如 '1.0.r123')?
-
-
分析 Array API 兼容层核心分发器的后端识别机制
阅读
sklearn/externals/array_api_compat/common/_helpers.py中array_namespace与_cls_to_namespace函数:-
追踪
array_namespace(np.array([1]), cp.array([1]))调用时的报错路径 -
理解
use_compat=None/True/False三种模式对 NumPy 后端的不同影响 -
分析
_ClsToXPInfo.SCALAR与MAYBE_JAX_ZERO_GRADIENT如何处理 Python 标量与 JAX 零梯度数组
回答问题:
-
为什么
_cls_to_namespace中issubclass(cls, int|float|complex|None)判断必须在np.generic之后? -
array_namespace如何保证多输入数组来自同一后端?违反时抛出什么异常?
-
-
对比 clip 函数的跨后端实现差异
对比阅读三个后端的
clip实现:-
sklearn/externals/array_api_compat/common/_aliases.py的通用实现 -
sklearn/externals/array_api_compat/torch/_aliases.py中clip = get_xp(torch)(_aliases.clip)(复用通用) -
sklearn/externals/array_api_compat/dask/array/_aliases.py中的独立实现
回答问题:
-
通用实现如何通过
out[()] = x+ 掩码赋值实现 dtype 保持?为什么要处理 Python 整数溢出截断? -
Dask 为何不能使用通用掩码实现?它的替代方案是什么?有什么局限?
-
PyTorch 为何选择复用通用实现而非原生
torch.clamp?
-
-
研究设备抽象的跨后端统一
阅读
sklearn/externals/array_api_compat/common/_helpers.py中device()与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_device在jax.jit上下文中为何可能失效?代码中有什么 workaround? -
Sparse 数组的
device递归查找逻辑(x.data回退)处理了哪种存储格式例外?
-
-
实现跨后端的 unique_all 函数
阅读
sklearn/externals/array_api_compat/common/_aliases.py中unique_all与各后端的实现差异:-
通用实现如何基于
xp.unique返回 NamedTuple 并修正inverse_indicesshape -
PyTorch 为何抛出 NotImplementedError(缺失 indices 返回)
-
NumPy/CuPy 如何通过条件导出使用原生
unique_all(NumPy 2.0+)
回答问题:
-
通用实现中
inverse_indices.reshape(x.shape)修正了什么 NumPy 行为? -
PyTorch 实现
unique_all的主要障碍是什么?有无替代方案?
-
-
排查 Dask 后端 sort/argsort 的内存风险
阅读
sklearn/externals/array_api_compat/dask/array/_aliases.py中sort与argsort实现:-
理解
_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-compat与array-api-extra的层次关系以及命名空间自动识别机制。 -
统一别名与包装:了解如何通过
_aliases.py将不同后端的函数签名统一为标准 API。 -
惰性求值:深入
lazy_apply在 Dask 与 JAX 中的实现细节以及何时触发计算。 -
只读更新:学会使用
at进行函数式“就地”更新,特别是针对 JAX 的不可变数组。 -
函数委托:掌握后端原生实现与标准实现之间的智能路由策略。
-
扩展函数:熟悉
cov、kron、sinc等超出标准的实用工具。 -
跨后端测试:了解
xp_assert_*、lazy_xp_function等测试设施如何保持行为一致。
生活类比:把
array-api-compat想成一座“国际翻译中心”,而array-api-extra则是这座中心的“增能套件”。后端(NumPy、CuPy、PyTorch、JAX、Dask …)是不同国家的语言,array_namespace()就像护照读取器,一眼识别语言并切换翻译模式。at、lazy_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_where、broadcast_shapes、cov、kron、sinc 等。 |
| 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_callback在jit环境下安全调用普通函数;若形状未知则提前报错。 -
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)→ 链式调用,与 JAXx.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_num、one_hot、pad)遵循相同模式:先尝试后端原生实现,若不存在则回退到 _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)还是户外跑步(f2或fill_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符合标准。
生活类比:这就像计算拼图的最终尺寸——有些拼图块可能边缘磨损导致尺寸不确定(用
None或math.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)
-
统一转 NumPy:
as_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_function 与 patch_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 替换原函数
工作流程:
- 标记
lazy_xp_function→ 在函数对象上挂载元数据。
- 测试时
patch_lazy_xp_functions读取元数据 → 根据后端(Dask / JAX)进行自动包装。
- Dask:使用
CountingDaskScheduler限制compute调用次数,以检测是否不小心 materialize。
- 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+ 标准函数的别名,处理 device、copy 等差异。 |
| _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_where、broadcast_shapes、cov、kron、sinc 等增强函数。 |
| _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:
-
array_namespace()如何通过输入数组的__array_namespace__方法或类型特征识别后端? -
is_dask_namespace()等检测函数的判断依据是什么?如何避免误判? -
_get_namespace()内部如何处理多数组输入时的命名空间一致性检查?
66.12.2 深入通用别名与函数包装机制
阅读 sklearn/externals/array_api_compat/common/_aliases.py:
-
_wrap_func()如何统一处理device、copy等标准参数与后端特有参数的差异? -
astype()等函数在不同后端的包装逻辑有何不同?以 NumPy 与 PyTorch 为例对比。 -
类型提升表
_fix_promotion_table()如何修正跨后端的类型推导不一致问题?
66.12.3 探究线性代数与 FFT 封装的一致性保障
阅读 sklearn/externals/array_api_compat/common/_linalg.py 与 _fft.py:
-
svd()等线性代数函数如何通过_linalg_func()包装器统一返回值格式(如 Vh vs V.T)? -
fft()系列函数如何处理不同后端对归一化参数(norm)的支持差异? -
对比 NumPy 与 CuPy 后端的
linalg.py实现,核心计算调用有何异同?
66.12.4 剖析后端专属适配器的差异化处理
阅读 sklearn/externals/array_api_compat/torch/_aliases.py 与 dask/array/_aliases.py:
-
PyTorch 适配器中
_fix_torch_promotion()如何解决 PyTorch 类型提升规则与标准的冲突? -
Dask 适配器中
copy()语义为何默认为惰性?astype()如何处理分块数组的类型转换? -
torch/_info.py中capabilities()如何针对 meta device 禁用布尔索引与数据依赖形状?
66.12.5 探索 array-api-extra 的惰性求值与更新操作
阅读 sklearn/externals/array_api_extra/_lib/_lazy.py:
-
lazy_apply()如何区分 eager 和 lazy 后端(如 Dask、JAX)的执行策略? -
在 JAX 的
jax.jit上下文中,为什么需要_is_jax_jit_enabled()检查? -
_lazy_apply_wrapper如何处理as_numpy参数和多输出情况?
66.12.6 理解 array-api-extra 的更新操作 (at)
阅读 sklearn/externals/array_api_extra/_lib/_at.py:
-
at类如何实现类似 JAX 的.at[].set()链式调用? -
对于只读数组(如 JAX 在 jit 中),
at.add()如何回退到函数式更新? -
为什么文档警告说在懒惰后端上应避免重用输入数组?请解释其中的副本语义与引用风险。
66.12.7 分析函数委托与扩展函数的实现
阅读 sklearn/externals/array_api_extra/_delegation.py 与 _lib/_funcs.py:
-
isclose()在不同后端(NumPy、CuPy、PyTorch 等)上的委托逻辑有何不同? -
_funcs.py中的broadcast_shapes()如何处理 Dask 的 NaN 形状与 Array API 的 None 形状? -
cov()函数在处理复数输入时,是如何分离实部和虚部进行计算的?
66.12.8 实践 array-api-extra 的测试工具与后端枚举
阅读 sklearn/externals/array_api_extra/_lib/_testing.py 与 testing.py:
-
xp_assert_close()如何统一处理不同后端(包括 GPU 上的 PyTorch/JAX)数组到 NumPy 的转换? -
lazy_xp_function装饰器如何配合patch_lazy_xp_functions实现 Dask/JAX 的测试隔离? -
Backend枚举的pytest_param()如何为不同后端自动添加thread_unsafe等标记?
66.13 架构与数据流图
上述图分别展示模块依赖、调用时序、数据流和架构分层。
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__ # 工具包初始化入口

浙公网安备 33010602011771号