// This file is part of Eigen, a lightweight C++ template library // for linear algebra. // // Copyright (C) 2026 Rasmus Munk Larsen // // This Source Code Form is subject to the terms of the Mozilla // Public License v. 2.0. If a copy of the MPL was not distributed // with this file, You can obtain one at http://mozilla.org/MPL/2.0/. // Solver expression types for DeviceMatrix. // // Each expression maps 1:1 to cuSOLVER library calls: // LltSolveExpr → cusolverDnXpotrf + cusolverDnXpotrs // LuSolveExpr → cusolverDnXgetrf + cusolverDnXgetrs // // Usage: // d_X = d_A.llt().solve(d_B); // Cholesky solve // d_X.device(ctx) = d_A.lu().solve(d_B); // LU solve on explicit stream #ifndef EIGEN_GPU_DEVICE_SOLVER_EXPR_H #define EIGEN_GPU_DEVICE_SOLVER_EXPR_H // IWYU pragma: private #include "./InternalHeaderCheck.h" namespace Eigen { // Forward declarations. template class DeviceMatrix; class GpuContext; // ---- LLT solve expression --------------------------------------------------- // d_A.llt().solve(d_B) → LltSolveExpr → cusolverDnXpotrf + cusolverDnXpotrs template class LltSolveExpr { public: using Scalar = Scalar_; enum { UpLo = UpLo_ }; LltSolveExpr(const DeviceMatrix& A, const DeviceMatrix& B) : A_(A), B_(B) {} const DeviceMatrix& matrix() const { return A_; } const DeviceMatrix& rhs() const { return B_; } private: const DeviceMatrix& A_; const DeviceMatrix& B_; }; // ---- LU solve expression ---------------------------------------------------- // d_A.lu().solve(d_B) → LuSolveExpr → cusolverDnXgetrf + cusolverDnXgetrs template class LuSolveExpr { public: using Scalar = Scalar_; LuSolveExpr(const DeviceMatrix& A, const DeviceMatrix& B) : A_(A), B_(B) {} const DeviceMatrix& matrix() const { return A_; } const DeviceMatrix& rhs() const { return B_; } private: const DeviceMatrix& A_; const DeviceMatrix& B_; }; // ---- DeviceLLTView: d_A.llt() → view with .solve() and .device() ----------- template class DeviceLLTView { public: using Scalar = Scalar_; explicit DeviceLLTView(const DeviceMatrix& m) : mat_(m) {} /** Build a solve expression: d_A.llt().solve(d_B). * The expression is evaluated when assigned to a DeviceMatrix. */ LltSolveExpr solve(const DeviceMatrix& rhs) const { return {mat_, rhs}; } // For cached factorizations, use the explicit GpuLLT API directly: // GpuLLT llt; // llt.compute(d_A); // auto d_X1 = llt.solve(d_B1); // auto d_X2 = llt.solve(d_B2); private: const DeviceMatrix& mat_; }; // ---- DeviceLUView: d_A.lu() → view with .solve() and .device() ------------- template class DeviceLUView { public: using Scalar = Scalar_; explicit DeviceLUView(const DeviceMatrix& m) : mat_(m) {} /** Build a solve expression: d_A.lu().solve(d_B). */ LuSolveExpr solve(const DeviceMatrix& rhs) const { return {mat_, rhs}; } // For cached factorizations, use the explicit GpuLU API directly: // GpuLU lu; // lu.compute(d_A); // auto d_X1 = lu.solve(d_B1); // auto d_X2 = lu.solve(d_B2); private: const DeviceMatrix& mat_; }; } // namespace Eigen #endif // EIGEN_GPU_DEVICE_SOLVER_EXPR_H