Source code for dpnp.scipy.sparse.linalg._iterative

# *****************************************************************************
# Copyright (c) 2026, Intel Corporation
# All rights reserved.
#
# Redistribution and use in source and binary forms, with or without
# modification, are permitted provided that the following conditions are met:
# - Redistributions of source code must retain the above copyright notice,
#   this list of conditions and the following disclaimer.
# - Redistributions in binary form must reproduce the above copyright notice,
#   this list of conditions and the following disclaimer in the documentation
#   and/or other materials provided with the distribution.
# - Neither the name of the copyright holder nor the names of its contributors
#   may be used to endorse or promote products derived from this software
#   without specific prior written permission.
#
# THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
# AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
# IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE
# ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE
# LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR
# CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF
# SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS
# INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN
# CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE)
# ARISING IN ANY WAY OUT OF THE USE OF THIS SOFTWARE, EVEN IF ADVISED OF
# THE POSSIBILITY OF SUCH DAMAGE.
# *****************************************************************************

"""Iterative sparse linear solvers for dpnp -- pure GPU/SYCL implementation.

All O(n)-or-larger work (SpMV, Arnoldi/Lanczos basis vectors, the solution
updates) stays on the device (USM/oneMKL). Only a few restart-sized quantities
touch the host: the per-iteration convergence scalars that drive Python-side
branching, and the GMRES Hessenberg least-squares solve (see ``gmres``).

Solver coverage
---------------
cg     : Conjugate Gradient (Hermitian positive definite)
gmres  : Restarted GMRES (general non-symmetric)
minres : MINRES (symmetric possibly indefinite)

SpMV fast-path
--------------
When a CSR dpnp sparse matrix is passed as A or M, _make_fast_matvec()
reuses the matrix's own cached oneMKL matrix_handle (built lazily by
``csr_matrix._ensure_spmv_handle`` via _sparse_gemv_init: set_csr_data +
optimize_gemv, the expensive sparsity-analysis phase). Each matvec then
calls only _sparse_gemv_compute (the cheap sparse::gemv kernel), so
optimize_gemv runs once per operator, not once per iteration.

Only the forward matvec is implemented for sparse operators; rmatvec /
adjoint (A^H) is not supported (the solvers never require it).

Supported dtypes for the oneMKL SpMV fast-path:

* values : float32, float64, complex64, complex128
* indices : int32, int64

Complex dtypes require oneMKL sparse BLAS support (available since
oneMKL 2023.x); if the dispatch table slot is nullptr (types_matrix.hpp
does not register the pair) _make_fast_matvec returns None and the
operator is routed through the generic LinearOperator path.
"""

# Single-letter and CamelCase names mirror the SciPy/CuPy API.
# duplicate-code fires against _csr.py, which shares the cached-SpMV
# and __del__ patterns required by oneMKL.
# pylint: disable=invalid-name,duplicate-code

from __future__ import annotations

import math
import warnings
from typing import Callable

import dpctl.utils as dpu
import numpy

import dpnp

# pylint cannot introspect the compiled extension.
# pylint: disable-next=no-name-in-module
import dpnp.backend.extensions.blas._blas_impl as bi
import dpnp.tensor as dpt
from dpnp.exceptions import ExecutionPlacementError

from ..._lib._sparse import issparse
from ._interface import (
    IdentityOperator,
    LinearOperator,
    _adjoint_unsupported,
    aslinearoperator,
)

_SUPPORTED_DTYPES = frozenset("fdFD")


def _np_dtype(dp_dtype) -> numpy.dtype:
    """Normalise any dtype-like to numpy.dtype."""
    return numpy.dtype(dp_dtype)


def _check_dtype(dtype, name: str) -> None:
    if _np_dtype(dtype).char not in _SUPPORTED_DTYPES:
        raise TypeError(
            f"{name} has unsupported dtype {dtype}; "
            "only float32, float64, complex64, complex128 are accepted."
        )


class _CachedSpMVPair:
    """Forward SpMV closure around a csr_matrix.

    The forward handle is owned by the ``csr_matrix`` itself (built via
    ``csr_matrix._ensure_spmv_handle()``) and therefore shared with any
    other call site -- including a user-issued ``A.dot(x)`` outside the
    solver. rmatvec (adjoint) is not supported.
    """

    __slots__ = ("_A", "_si")

    def __init__(self, A, si):
        self._A = A
        self._si = si

    def matvec(self, x):
        """Apply the forward operator A @ x via the csr's cached handle."""
        # Validated by _make_fast_matvec, so it cannot return None here.
        # pylint: disable-next=protected-access
        _si, handle, val_type_id, exec_q = self._A._ensure_spmv_handle()
        # oneMKL reads x as a bare unit-stride pointer; a strided view
        # (a column of a C-contiguous 2-D array) must be packed first.
        if not x.flags.c_contiguous:
            x = dpnp.ascontiguousarray(x)
        y = dpnp.empty_like(self._A.data, shape=self._A.shape[0])
        _manager = dpu.SequentialOrderManager[exec_q]
        # pylint: disable-next=protected-access
        ht_ev, comp_ev = _si._sparse_gemv_compute(
            exec_q,
            handle,
            val_type_id,
            0,  # trans=N
            1.0,  # alpha
            dpnp.get_usm_ndarray(x),
            0.0,  # beta
            dpnp.get_usm_ndarray(y),
            int(self._A.shape[0]),
            int(self._A.shape[1]),
            _manager.submitted_events,
        )
        _manager.add_event_pair(ht_ev, comp_ev)
        return y

    # pylint: disable-next=unused-argument
    def rmatvec(self, x):
        """Adjoint SpMV is not supported (see ``LinearOperator``)."""
        _adjoint_unsupported("rmatvec")


def _make_fast_matvec(A):
    """Return a _CachedSpMVPair if A is a CSR matrix with oneMKL support,
    or None if A is not an eligible sparse matrix.

    Returns None when A is not a dpnp CSR sparse matrix, when its
    (value, index) dtype combination is not registered with the oneMKL
    dispatch table, or when ``A.nnz == 0`` (oneMKL cannot build a handle
    for an empty matrix). In every None case the caller keeps the generic
    ``MatrixLinearOperator`` wrapper, whose matvec routes through
    ``csr_matrix.dot`` and returns the correct zero vector for nnz == 0.
    """
    if not (issparse(A) and getattr(A, "format", None) == "csr"):
        return None

    # Returns a handle cached on A, or None when no handle can be built.
    if not hasattr(A, "_ensure_spmv_handle"):
        return None
    # pylint: disable-next=protected-access
    handle_info = A._ensure_spmv_handle()
    if handle_info is None:
        return None

    _si, _handle, _val_type_id, _exec_q = handle_info
    return _CachedSpMVPair(A, _si)


def _make_system(A, M, x0, b):
    """
    Make a linear system ``Ax = b``.

    Parameters
    ----------
    A : {dpnp.ndarray, usm_ndarray, csr_matrix, LinearOperator}
        Sparse or dense matrix of the linear system.
    M : {None, dpnp.ndarray, usm_ndarray, csr_matrix, LinearOperator}
        Preconditioner.
    x0 : {None, dpnp.ndarray, usm_ndarray}
        Initial guess for the iterative method.
    b : {dpnp.ndarray, usm_ndarray}
        Right hand side.

    Returns
    -------
    out : tuple
        A tuple ``(A, M, x, b, dtype)`` holding the operator, the
        preconditioner, the initial guess, the right hand side and the
        working data type.
    """
    if not dpnp.is_supported_array_type(b):
        raise TypeError(
            f"b must be a dpnp.ndarray or usm_ndarray, got {type(b).__name__}"
        )
    if x0 is not None and not dpnp.is_supported_array_type(x0):
        raise TypeError(
            f"x0 must be a dpnp.ndarray, usm_ndarray or None, "
            f"got {type(x0).__name__}"
        )

    A_op = aslinearoperator(A)
    if A_op.shape[0] != A_op.shape[1]:
        raise ValueError("A must be a square operator")
    n = A_op.shape[0]

    b = dpnp.reshape(b, -1)
    if b.shape[0] != n:
        raise ValueError(
            f"b length {b.shape[0]} does not match operator dimension {n}"
        )

    # Dtype promotion: prefer A.dtype; fall back via b.dtype. The
    # fallback must respect the device's fp64 aspect -- float64 /
    # complex128 are not supported on every GPU and hardcoding them
    # would raise downstream instead of degrading gracefully like
    # dpnp.default_float_type()
    if (
        A_op.dtype is not None
        and _np_dtype(A_op.dtype).char in _SUPPORTED_DTYPES
    ):
        dtype = A_op.dtype
    elif dpnp.issubdtype(b.dtype, dpnp.complexfloating):
        has_fp64 = b.sycl_device.has_aspect_fp64
        dtype = dpnp.complex128 if has_fp64 else dpnp.complex64
    else:
        dtype = dpnp.default_float_type(sycl_queue=b.sycl_queue)

    b = dpnp.astype(b, dtype, copy=False)
    _check_dtype(b.dtype, "b")

    if x0 is None:
        x = dpnp.zeros(n, dtype=dtype, sycl_queue=b.sycl_queue)
    else:
        x = dpnp.reshape(dpnp.astype(x0, dtype, copy=True), -1)
        if x.shape[0] != n:
            raise ValueError(f"x0 length {x.shape[0]} != n={n}")

    if M is None:
        M_op = IdentityOperator((n, n), dtype=dtype)
    else:
        M_op = aslinearoperator(M)
        if M_op.shape != A_op.shape:
            raise ValueError(
                f"preconditioner shape {M_op.shape} != "
                f"operator shape {A_op.shape}"
            )

        fast_mv_M = _make_fast_matvec(M)
        if fast_mv_M is not None:
            _orig_M = M_op

            class _FastMOp(LinearOperator):
                def __init__(self):
                    super().__init__(_orig_M.dtype, _orig_M.shape)

                def _matvec(self, x):
                    return fast_mv_M.matvec(x)

            M_op = _FastMOp()

    # Inject fast CSR SpMV for A if available.
    fast_mv = _make_fast_matvec(A)
    if fast_mv is not None:
        _orig = A_op

        class _FastOp(LinearOperator):
            def __init__(self):
                super().__init__(_orig.dtype, _orig.shape)

            def _matvec(self, x):
                return fast_mv.matvec(x)

        A_op = _FastOp()

    return A_op, M_op, x, b, dtype


def _check_maxiter(maxiter, default: int) -> int:
    """Resolve `maxiter`, rejecting negative values.

    A negative budget used to propagate straight into the return value
    (``cg(maxiter=-1)`` returned ``info == -1``), contradicting the
    documented contract that ``info < 0`` is never produced. ``0`` is
    kept legal and means "do no iterations", matching SciPy.
    """
    if maxiter is None:
        return default
    maxiter = int(maxiter)
    if maxiter < 0:
        raise ValueError(
            f"maxiter={maxiter} is invalid; must be a non-negative integer."
        )
    return maxiter


def _get_atol(b_norm: float, atol, rtol: float) -> float:
    """Absolute stopping tolerance: max(atol, rtol*||b||), mirroring SciPy."""
    if atol == "legacy" or atol is None:
        atol = 0.0
    atol = float(atol)
    if atol < 0:
        raise ValueError(
            f"atol={atol!r} is invalid; must be a real, non-negative number."
        )
    return max(atol, float(rtol) * float(b_norm))


# pylint: disable-next=too-many-locals,too-many-statements
[docs] def cg( A, b, x0: dpnp.ndarray | None = None, *, rtol: float = 1e-5, tol: float | None = None, maxiter: int | None = None, M=None, callback: Callable | None = None, atol=None, ) -> tuple[dpnp.ndarray, int]: """ Use Conjugate Gradient iteration to solve ``Ax = b`` for a Hermitian positive-definite ``A``. For full documentation refer to :obj:`scipy.sparse.linalg.cg`. Parameters ---------- A : {dpnp.ndarray, usm_ndarray, LinearOperator, csr_matrix} The Hermitian positive-definite operator of shape ``(N, N)``. b : {dpnp.ndarray, usm_ndarray} Right-hand side of the linear system, shape ``(N,)`` or ``(N, 1)``. x0 : {None, dpnp.ndarray, usm_ndarray}, optional Initial guess for the solution. Default: ``None`` (zeros). rtol : float, optional Relative convergence tolerance. Default: ``1e-5``. tol : {None, float}, optional Deprecated alias for `rtol`. Default: ``None``. maxiter : {None, int}, optional Maximum number of iterations. Default: ``10 * N``. M : {None, dpnp.ndarray, usm_ndarray, LinearOperator}, optional Symmetric positive-definite preconditioner. Default: ``None``. callback : {None, callable}, optional Called as ``callback(xk)`` after each iteration. Default: ``None``. atol : {None, float}, optional Absolute convergence tolerance. Default: ``None``. Returns ------- x : dpnp.ndarray The converged solution. info : int ``info`` follows the SciPy / CuPy contract: * ``info == 0`` : converged successfully * ``info > 0`` : did not converge; value is the iteration count at which the solver stopped (equals ``maxiter`` when the iteration budget was exhausted, or the iteration index when a numerical breakdown short-circuited the loop). * ``info < 0`` : reserved for illegal-input errors; not produced by this implementation (illegal inputs raise ``ValueError`` instead). Previous versions of this routine returned ``-1`` for an ``rz``/``pAp`` breakdown, which violated the SciPy contract and broke user code that branched on ``info > 0``. """ if tol is not None: warnings.warn( "'tol' is deprecated in favor of 'rtol' and will be removed in " "a future release; use 'rtol' instead.", DeprecationWarning, stacklevel=2, ) rtol = tol A_op, M_op, x, b, dtype = _make_system(A, M, x0, b) n = b.shape[0] bnrm = dpnp.linalg.norm(b) bnrm_host = float(bnrm) if bnrm_host == 0.0: return dpnp.zeros_like(b), 0 atol_eff_host = _get_atol(bnrm_host, atol=atol, rtol=rtol) maxiter = _check_maxiter(maxiter, n * 10) rhotol = float(numpy.finfo(_np_dtype(dtype)).eps ** 2) r = b - A_op.matvec(x) if x0 is not None else b.copy() z = M_op.matvec(r) p = z.copy() rz = dpnp.real(dpnp.vdot(r, z)) if float(dpnp.abs(rz)) < rhotol: # rz == 0 has two causes: r == 0 (already converged) or a # preconditioner breakdown, M @ r == 0 with r != 0 (singular M). # Only the first is convergence; reporting info == 0 for the # second would claim a solution that was never computed. if float(dpnp.linalg.norm(r)) <= atol_eff_host: return x, 0 return x, 1 info = maxiter for k in range(maxiter): rnorm = dpnp.linalg.norm(r) rnorm_host = float(rnorm) if rnorm_host <= atol_eff_host: info = 0 break if not math.isfinite(rnorm_host): # pAp or rz collapsed last iteration, poisoning r with # inf/NaN. Report info > 0 per the SciPy contract. info = k + 1 break Ap = A_op.matvec(p) pAp = dpnp.real(dpnp.vdot(p, Ap)) alpha = rz / pAp x = x + alpha * p r = r - alpha * Ap if callback is not None: callback(x) z = M_op.matvec(r) rz_new = dpnp.real(dpnp.vdot(r, z)) beta = rz_new / rz p = z + beta * p rz = rz_new else: info = maxiter return x, int(info)
# pylint: disable-next=too-many-locals,too-many-statements,too-many-branches
[docs] def gmres( A, b, x0: dpnp.ndarray | None = None, *, rtol: float = 1e-5, atol: float = 0.0, restart: int | None = None, maxiter: int | None = None, M=None, callback: Callable | None = None, callback_type: str | None = None, ) -> tuple[dpnp.ndarray, int]: """ Use Generalized Minimal RESidual iteration to solve ``Ax = b``. For full documentation refer to :obj:`scipy.sparse.linalg.gmres`. Parameters ---------- A : {dpnp.ndarray, usm_ndarray, LinearOperator, csr_matrix} The real or complex operator of the linear system, shape ``(N, N)``. b : {dpnp.ndarray, usm_ndarray} Right-hand side of the linear system, shape ``(N,)`` or ``(N, 1)``. x0 : {None, dpnp.ndarray, usm_ndarray}, optional Starting guess for the solution. Default: ``None`` (zeros). rtol, atol : float, optional Convergence tolerance: ``||r|| <= max(atol, rtol*||b||)``. Defaults: ``rtol=1e-5``, ``atol=0.0``. restart : {None, int}, optional Number of iterations between restarts. Larger values increase the per-iteration cost but may aid convergence. Default: ``20``. maxiter : {None, int}, optional Maximum number of iterations. Default: ``10 * N``. M : {None, dpnp.ndarray, usm_ndarray, LinearOperator}, optional Preconditioner approximating the inverse of `A`. Default: ``None``. callback : {None, callable}, optional Called on every restart as ``callback(arg)``, where `arg` is selected by `callback_type`. Default: ``None``. callback_type : {None, 'x', 'pr_norm'}, optional ``'x'`` passes the current solution vector; ``'pr_norm'`` passes the relative (preconditioned) residual norm. Default: ``'pr_norm'`` when a callback is supplied. Returns ------- x : dpnp.ndarray The approximate solution (``M @ x`` in the right-preconditioned formulation, matching CuPy's return value). info : int ``0`` if converged; the iteration count if `maxiter` was reached. See Also -------- :obj:`scipy.sparse.linalg.gmres` :obj:`cupyx.scipy.sparse.linalg.gmres` """ A_op, M_op, x, b, dtype = _make_system(A, M, x0, b) matvec = A_op.matvec psolve = M_op.matvec n = A_op.shape[0] if n == 0: return dpnp.empty_like(b), 0 b_norm = float(dpnp.linalg.norm(b)) if b_norm == 0.0: return dpnp.zeros_like(b), 0 atol = max(float(atol), rtol * b_norm) maxiter = _check_maxiter(maxiter, n * 10) if restart is None: restart = 20 # A non-positive restart leaves the Arnoldi basis empty, which used # to surface as an IndexError from V[:, 0] (restart == 0) or an # allocation error (restart < 0) deep inside the loop. if int(restart) < 1: raise ValueError( f"restart={restart!r} is invalid; must be a positive integer." ) restart = min(int(restart), n) if callback_type is None: callback_type = "pr_norm" if callback_type not in ("x", "pr_norm"): raise ValueError(f"Unknown callback_type: {callback_type!r}") if callback is None: callback_type = None queue = b.sycl_queue # V and H are F-ordered so column slices are contiguous USM views, # required by the bi._gemv binding in _make_compute_hu. V = dpnp.empty((n, restart), dtype=dtype, sycl_queue=queue, order="F") H = dpnp.zeros( (restart + 1, restart), dtype=dtype, sycl_queue=queue, order="F" ) compute_hu = _make_compute_hu(V, H) np_dtype = _np_dtype(dtype) e_host = numpy.zeros(restart + 1, dtype=np_dtype) eps = numpy.finfo(np_dtype).eps iters = 0 r_norm_host = math.inf while True: mx = psolve(x) r = b - matvec(mx) r_norm = dpnp.linalg.norm(r) r_norm_host = float(r_norm) if callback_type == "x": callback(mx) elif callback_type == "pr_norm" and iters > 0: callback(r_norm_host / b_norm) if r_norm_host <= atol or iters >= maxiter: break # Start the Arnoldi basis from the normalised residual. v = r / r_norm V[:, 0] = v # Reset H so stale entries from the previous restart cannot # leak into the least-squares system. H[:] = 0 # RHS of the Hessenberg system is r_norm * e_1. e_host[0] = r_norm_host if iters > 0: e_host[1:] = 0 # Arnoldi iteration last_j = restart - 1 for j in range(restart): z = psolve(v) u = matvec(z) u = compute_hu(u, j) h_norm = dpnp.linalg.norm(u) H[j + 1, j] = h_norm if j < last_j: v = u / dpnp.where(h_norm == 0, dpnp.ones_like(h_norm), h_norm) V[:, j + 1] = v # Hessenberg LS + breakdown check run on host: H is tiny # ((restart+1) x restart, ~21x20), so a device SVD would be # launch/alloc-overhead bound. Only y (restart-length) is # copied back to the device. H_host = dpnp.asnumpy(H) eps_break = eps * r_norm_host subdiag = numpy.abs(numpy.diagonal(H_host, offset=-1)) breakdown = numpy.nonzero(subdiag <= eps_break)[0] built = int(breakdown[0]) + 1 if breakdown.size else restart H_sub = H_host[: built + 1, :built] y_host, *_ = numpy.linalg.lstsq(H_sub, e_host[: built + 1], rcond=None) y = dpnp.asarray(y_host, sycl_queue=queue) x = x + dpnp.dot(V[:, :built], y) iters += built info = 0 if iters >= maxiter and r_norm_host > atol: info = iters return mx, info
# pylint: disable-next=too-many-locals,too-many-branches,too-many-statements
[docs] def minres( A, b, x0: dpnp.ndarray | None = None, *, rtol: float = 1e-5, shift: float = 0.0, maxiter: int | None = None, M=None, callback: Callable | None = None, show: bool = False, check: bool = False, ) -> tuple[dpnp.ndarray, int]: """ Use MINimum RESidual iteration to solve ``Ax = b``. Solves the symmetric (possibly indefinite) system ``Ax = b`` or, if `shift` is nonzero, ``(A - shift*I)x = b``. All computation stays on the SYCL device; only scalar recurrence coefficients and norms are transferred to the host for branching. For full documentation refer to :obj:`scipy.sparse.linalg.minres`. Parameters ---------- A : {dpnp.ndarray, usm_ndarray, LinearOperator, csr_matrix} The real symmetric or complex Hermitian operator, shape ``(N, N)``. b : {dpnp.ndarray, usm_ndarray} Right-hand side, shape ``(N,)`` or ``(N, 1)``. x0 : {None, dpnp.ndarray, usm_ndarray}, optional Starting guess for the solution. Default: ``None`` (zeros). shift : float, optional If nonzero, solve ``(A - shift*I)x = b``. Default: ``0.0``. rtol : float, optional Relative convergence tolerance. Default: ``1e-5``. maxiter : {None, int}, optional Maximum number of iterations. Default: ``5 * N``. M : {None, dpnp.ndarray, usm_ndarray, LinearOperator}, optional Preconditioner approximating the inverse of `A`. Default: ``None``. callback : {None, callable}, optional Called as ``callback(xk)`` after each iteration. Default: ``None``. show : bool, optional If ``True``, print a convergence summary each iteration. Default: ``False``. check : bool, optional If ``True``, verify that `A` and `M` are symmetric before iterating (costs extra matvecs). Default: ``False``. Returns ------- x : dpnp.ndarray The converged (or best) solution. info : int ``0`` if converged, ``maxiter`` if the iteration limit was reached. Notes ----- This is a direct translation of the Paige--Saunders MINRES algorithm as implemented in SciPy, adapted for dpnp device arrays with the oneMKL SpMV cached-handle fast-path. See Also -------- scipy.sparse.linalg.minres cupyx.scipy.sparse.linalg.minres """ A_op, M_op, x, b, dtype = _make_system(A, M, x0, b) matvec = A_op.matvec psolve = M_op.matvec n = A_op.shape[0] maxiter = _check_maxiter(maxiter, 5 * n) istop = 0 itn = 0 eps = dpnp.finfo(dtype).eps # Set up y and v for the first Lanczos vector v1: # y = beta1 * P' * v1, where P = M**(-1); v is P' * v1. Ax = matvec(x) r1 = b - Ax y = psolve(r1) beta1 = float(dpnp.vdot(r1, y).real) if beta1 < 0: raise ValueError("indefinite preconditioner") if beta1 == 0: return (x, 0) beta1 = math.sqrt(beta1) if check: # See if A is symmetric. w_chk = matvec(y) r2_chk = matvec(w_chk) s = float(dpnp.vdot(w_chk, w_chk).real) t = float(dpnp.vdot(y, r2_chk).real) if abs(s - t) > (s + eps) * eps ** (1.0 / 3.0): raise ValueError("non-symmetric matrix") # See if M is symmetric. r2_chk = psolve(y) s = float(dpnp.vdot(y, y).real) t = float(dpnp.vdot(r1, r2_chk).real) if abs(s - t) > (s + eps) * eps ** (1.0 / 3.0): raise ValueError("non-symmetric preconditioner") # Initialise remaining quantities (all host-side scalars). oldb = 0 beta = beta1 dbar = 0 epsln = 0 qrnorm = beta1 phibar = beta1 rhs1 = beta1 rhs2 = 0 tnorm2 = 0 gmax = 0 gmin = dpnp.finfo(dtype).max cs = -1 sn = 0 queue = b.sycl_queue w = dpnp.zeros(n, dtype=dtype, sycl_queue=queue) w2 = dpnp.zeros(n, dtype=dtype, sycl_queue=queue) r2 = r1 # Main Lanczos loop. while itn < maxiter: itn += 1 s = 1.0 / beta v = s * y # on device y = matvec(v) y = y - shift * v if itn >= 2: y = y - (beta / oldb) * r1 # alpha = <v, y> alpha = float(dpnp.vdot(v, y).real) y = y - (alpha / beta) * r2 r1 = r2 r2 = y y = psolve(r2) oldb = beta # beta = sqrt(<r2, y>) beta = float(dpnp.vdot(r2, y).real) if beta < 0: raise ValueError("non-symmetric matrix") beta = math.sqrt(beta) tnorm2 += alpha**2 + oldb**2 + beta**2 if itn == 1: if beta / beta1 <= 10 * eps: istop = -1 # Terminate later # Apply previous rotation Q_{k-1} to get # [delta_k epsln_{k+1}] = [cs sn] [dbar_k 0 ] # [gbar_k dbar_{k+1} ] [sn -cs] [alpha_k beta_{k+1}] oldeps = epsln delta = cs * dbar + sn * alpha gbar = sn * dbar - cs * alpha epsln = sn * beta dbar = -cs * beta root = math.hypot(gbar, dbar) # Compute the next plane rotation Q_k. gamma = math.hypot(gbar, beta) gamma = max(gamma, eps) cs = gbar / gamma sn = beta / gamma phi = cs * phibar phibar = sn * phibar # Update x -- all on device. denom = 1.0 / gamma w1 = w2 w2 = w w = (v - oldeps * w1 - delta * w2) * denom x = x + phi * w gmax = max(gmax, gamma) gmin = min(gmin, gamma) z = rhs1 / gamma rhs1 = rhs2 - delta * z rhs2 = -epsln * z # Estimate norms and test for convergence. Anorm = math.sqrt(tnorm2) ynorm = float(dpnp.linalg.norm(x)) epsa = Anorm * eps epsx = Anorm * ynorm * eps epsr = Anorm * ynorm * rtol diag = gbar if diag == 0: diag = epsa qrnorm = phibar rnorm = qrnorm if ynorm == 0 or Anorm == 0: test1 = math.inf else: test1 = rnorm / (Anorm * ynorm) # ||r|| / (||A|| ||x||) if Anorm == 0: test2 = math.inf else: test2 = root / Anorm # ||Ar|| / (||A|| ||r||) # Estimate cond(A). Acond = gmax / gmin # Stopping criteria (SciPy's istop codes). if istop == 0: t1 = 1 + test1 t2 = 1 + test2 if t2 <= 1: istop = 2 if t1 <= 1: istop = 1 if itn >= maxiter: istop = 6 if Acond >= 0.1 / eps: istop = 4 if epsx >= beta1: istop = 3 if test2 <= rtol: istop = 2 if test1 <= rtol: istop = 1 if show: prnt = ( n <= 40 or itn <= 10 or itn >= maxiter - 10 or itn % 10 == 0 or qrnorm <= 10 * epsx or qrnorm <= 10 * epsr or Acond <= 1e-2 / eps or istop != 0 ) if prnt: x1 = float(dpnp.real(x[0])) print( f"{itn:6g} {x1:12.5e} {test1:10.3e}" f" {test2:10.3e}" f" {Anorm:8.1e} {Acond:8.1e}" f" {gbar / Anorm if Anorm else 0:8.1e}" ) if itn % 10 == 0: print() if callback is not None: callback(x) if istop != 0: break if istop == 6: info = maxiter else: info = 0 return (x, info)
def _make_compute_hu(V, H): """Factory for the GMRES Arnoldi inner step on Intel GPU. Returns a closure ``compute_hu(u, j) -> u`` that performs classical Gram-Schmidt orthogonalisation of ``u`` against the first ``j+1`` columns of ``V`` and writes the projection coefficients into column ``j`` of ``H``: h = V[:, :j+1].conj().T @ u H[:j+1, j] = h u = u - V[:, :j+1] @ h Each pass is a single oneMKL ``gemv`` via the ``bi._gemv`` binding: * Pass 1 (project) -- ``gemv(trans_op=T or C, alpha=1, beta=0)`` writing straight into ``H[:j+1, j]``, no temporary buffer. ``trans_op`` is T for real and C for complex matrices, so oneMKL produces ``V^H u`` directly. * Pass 2 (subtract) -- ``gemv(trans_op=N, alpha=-1, beta=1)`` with in-place output ``u``, fusing the AXPY update. Parameters ---------- V : dpnp.ndarray Krylov basis of shape ``(n, restart)``, must be F-contiguous. H : dpnp.ndarray Hessenberg matrix of shape ``(restart+1, restart)``, must be F-contiguous so column slices ``H[:k, j]`` are unit-stride contiguous USM views the C binding can write into. Returns ------- closure : callable ``compute_hu(u, j) -> u`` -- updates ``H[:j+1, j]`` in place and returns the orthogonalised ``u``. """ if V.ndim != 2 or not V.flags.f_contiguous: raise ValueError( "_make_compute_hu: V must be a 2-D column-major (F-order) " "dpnp array" ) if H.ndim != 2 or not H.flags.f_contiguous: raise ValueError( "_make_compute_hu: H must be a 2-D column-major (F-order) " "dpnp array so column slices are unit-stride USM views" ) exec_q = dpt.get_execution_queue((V.sycl_queue, H.sycl_queue)) if exec_q is None: raise ExecutionPlacementError( "_make_compute_hu: V and H must share the same SYCL queue" ) dtype = V.dtype is_cpx = dpnp.issubdtype(dtype, dpnp.complexfloating) # Project with V^T for real matrices (== V^H) and with V^H for # complex ones; conj(V^T @ u) == V^H @ u only holds for real u. pass1_trans_op = 2 if is_cpx else 1 # 2 = conjtrans, 1 = transpose def compute_hu(u, j): Vj = V[:, : j + 1] h_slice = H[: j + 1, j] # length-(j+1) F-contig column slice Vj_usm = dpnp.get_usm_ndarray(Vj) u_usm = dpnp.get_usm_ndarray(u) h_usm = dpnp.get_usm_ndarray(h_slice) _manager = dpu.SequentialOrderManager[exec_q] # Pass 1: H[:j+1, j] = op(Vj) @ u (alpha=1, beta=0) # pylint: disable-next=protected-access ht1, ev1 = bi._gemv( exec_q, Vj_usm, u_usm, h_usm, trans_op=pass1_trans_op, alpha=1.0, beta=0.0, depends=_manager.submitted_events, ) _manager.add_event_pair(ht1, ev1) # Pass 2: u = -Vj @ H[:j+1, j] + 1 * u (alpha=-1, beta=1) # pylint: disable-next=protected-access ht2, ev2 = bi._gemv( exec_q, Vj_usm, h_usm, u_usm, trans_op=0, # N: no transpose alpha=-1.0, beta=1.0, depends=_manager.submitted_events, ) _manager.add_event_pair(ht2, ev2) return u return compute_hu