/
githubmirror
/
scikit-learn
Обзор
Документация
Войти
/
githubmirror
/
scikit-learn
Код
Запросы
0
Пакеты
0
Релизы
0
Аналитика
Безопасность
main
sklearn/externals/array_api_compat/numpy/linalg.py
208 строк
6 KB
Lucas Colley
MNT: bump to array API 2025.12, array-api-compat 1.15, array-api-extra 0.10.3 (#34231)
11 июн 2026, 17:53
Не верифицирован
11 июн 2026, 17:53
2b91769
Код
Авторство
О чём код?
# pyright: reportAttributeAccessIssue=false # pyright: reportUnknownArgumentType=false # pyright: reportUnknownMemberType=false # pyright: reportUnknownVariableType=false from __future__ import annotations import numpy as np from .._internal import clone_module, get_xp from ..common import _linalg __all__ = clone_module("numpy.linalg", globals()) # These functions are in both the main and linalg namespaces from ._aliases import matmul, matrix_transpose, tensordot, vecdot # noqa: F401 from ._typing import Array cross = get_xp(np)(_linalg.cross) outer = get_xp(np)(_linalg.outer) EighResult = _linalg.EighResult EigResult = _linalg.EigResult QRResult = _linalg.QRResult SlogdetResult = _linalg.SlogdetResult SVDResult = _linalg.SVDResult eigh = get_xp(np)(_linalg.eigh) qr = get_xp(np)(_linalg.qr) slogdet = get_xp(np)(_linalg.slogdet) svd = get_xp(np)(_linalg.svd) cholesky = get_xp(np)(_linalg.cholesky) matrix_rank = get_xp(np)(_linalg.matrix_rank) pinv = get_xp(np)(_linalg.pinv) matrix_norm = get_xp(np)(_linalg.matrix_norm) svdvals = get_xp(np)(_linalg.svdvals) diagonal = get_xp(np)(_linalg.diagonal) trace = get_xp(np)(_linalg.trace) # Note: unlike np.linalg.solve, the array API solve() only accepts x2 as a # vector when it is exactly 1-dimensional. All other cases treat x2 as a stack # of matrices. The np.linalg.solve behavior of allowing stacks of both # matrices and vectors is ambiguous c.f. # https://github.com/numpy/numpy/issues/15349 and # https://github.com/data-apis/array-api/issues/285. # To workaround this, the below is the code from np.linalg.solve except # only calling solve1 in the exactly 1D case. # This code is here instead of in common because it is numpy specific. Also # note that CuPy's solve() does not currently support broadcasting (see # https://github.com/cupy/cupy/blob/main/cupy/cublas.py#L43). def solve(x1: Array, x2: Array, /) -> Array: try: from numpy.linalg._linalg import ( # type: ignore[attr-defined] _assert_stacked_2d, _assert_stacked_square, _commonType, _makearray, _raise_linalgerror_singular, isComplexType, ) except ImportError: from numpy.linalg.linalg import ( # type: ignore[attr-defined] _assert_stacked_2d, _assert_stacked_square, _commonType, _makearray, _raise_linalgerror_singular, isComplexType, ) from numpy.linalg import _umath_linalg x1, _ = _makearray(x1) _assert_stacked_2d(x1) _assert_stacked_square(x1) x2, wrap = _makearray(x2) t, result_t = _commonType(x1, x2) # This part is different from np.linalg.solve gufunc: np.ufunc if x2.ndim == 1: gufunc = _umath_linalg.solve1 else: gufunc = _umath_linalg.solve # This does nothing currently but is left in because it will be relevant # when complex dtype support is added to the spec in 2022. signature = "DD->D" if isComplexType(t) else "dd->d" with np.errstate( call=_raise_linalgerror_singular, invalid="call", over="ignore", divide="ignore", under="ignore", ): r: Array = gufunc(x1, x2, signature=signature) return wrap(r.astype(result_t, copy=False)) # Unlike numpy.linalg.eig, Array API version always returns complex results def eig(x: Array, /) -> tuple[Array, Array]: try: from numpy.linalg._linalg import ( # type: ignore[attr-defined] _assert_stacked_square, _assert_finite, _commonType, _makearray, _raise_linalgerror_eigenvalues_nonconvergence, isComplexType, _complexType, ) except ImportError: from numpy.linalg.linalg import ( # type: ignore[attr-defined] _assert_stacked_square, _assert_finite, _commonType, _makearray, _raise_linalgerror_eigenvalues_nonconvergence, isComplexType, _complexType, ) from numpy.linalg import _umath_linalg x, wrap = _makearray(x) _assert_stacked_square(x) _assert_finite(x) t, result_t = _commonType(x) signature = 'D->DD' if isComplexType(t) else 'd->DD' with np.errstate(call=_raise_linalgerror_eigenvalues_nonconvergence, invalid='call', over='ignore', divide='ignore', under='ignore'): w, vt = _umath_linalg.eig(x, signature=signature) result_t = _complexType(result_t) vt = vt.astype(result_t, copy=False) return EigResult(w.astype(result_t, copy=False), wrap(vt)) def eigvals(x: Array, /) -> Array: try: from numpy.linalg._linalg import ( # type: ignore[attr-defined] _assert_stacked_square, _assert_finite, _commonType, _makearray, _raise_linalgerror_eigenvalues_nonconvergence, isComplexType, _complexType, ) except ImportError: from numpy.linalg.linalg import ( # type: ignore[attr-defined] _assert_stacked_square, _assert_finite, _commonType, _makearray, _raise_linalgerror_eigenvalues_nonconvergence, isComplexType, _complexType, ) from numpy.linalg import _umath_linalg x, wrap = _makearray(x) _assert_stacked_square(x) _assert_finite(x) t, result_t = _commonType(x) signature = 'D->D' if isComplexType(t) else 'd->D' with np.errstate(call=_raise_linalgerror_eigenvalues_nonconvergence, invalid='call', over='ignore', divide='ignore', under='ignore'): w = _umath_linalg.eigvals(x, signature=signature) result_t = _complexType(result_t) return w.astype(result_t, copy=False) # These functions are completely new here. If the library already has them # (i.e., numpy 2.0), use the library version instead of our wrapper. if hasattr(np.linalg, "vector_norm"): vector_norm = np.linalg.vector_norm else: vector_norm = get_xp(np)(_linalg.vector_norm) _all = [ "LinAlgError", "cond", "det", "eig", "eigvals", "eigvalsh", "inv", "lstsq", "matrix_power", "multi_dot", "norm", "solve", "tensorinv", "tensorsolve", "vector_norm", ] __all__ = sorted(set(__all__) | set(_linalg.__all__) | set(_all)) def __dir__() -> list[str]: return __all__