Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
DeviceSolverExpr.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2026 Rasmus Munk Larsen <rmlarsen@gmail.com>
5//
6// This Source Code Form is subject to the terms of the Mozilla
7// Public License v. 2.0. If a copy of the MPL was not distributed
8// with this file, You can obtain one at http://mozilla.org/MPL/2.0/.
9// SPDX-License-Identifier: MPL-2.0
10
11// Solver expression types for gpu::DeviceMatrix. Each maps 1:1 onto a pair of
12// cuSOLVER calls and factors afresh on every assignment; use the gpu::LLT /
13// gpu::LU classes when a factorization should be cached across solves.
14
15#ifndef EIGEN_GPU_DEVICE_SOLVER_EXPR_H
16#define EIGEN_GPU_DEVICE_SOLVER_EXPR_H
17
18// IWYU pragma: private
19#include "./InternalHeaderCheck.h"
20
21#include <functional>
22
23#include "./FwdDecl.h"
24
25namespace Eigen {
26namespace gpu {
27
29template <typename Scalar_, int UpLo_ = Lower>
30class LltSolveExpr {
31 public:
32 using Scalar = Scalar_;
33 static constexpr int UpLo = UpLo_;
34
35 LltSolveExpr(const DeviceMatrix<Scalar>& A, const DeviceMatrix<Scalar>& B) : A_(A), B_(B) {}
36 const DeviceMatrix<Scalar>& matrix() const { return A_; }
37 const DeviceMatrix<Scalar>& rhs() const { return B_; }
38
39 private:
40 std::reference_wrapper<const DeviceMatrix<Scalar>> A_;
41 std::reference_wrapper<const DeviceMatrix<Scalar>> B_;
42};
43
45template <typename Scalar_>
46class LuSolveExpr {
47 public:
48 using Scalar = Scalar_;
49
50 LuSolveExpr(const DeviceMatrix<Scalar>& A, const DeviceMatrix<Scalar>& B) : A_(A), B_(B) {}
51 const DeviceMatrix<Scalar>& matrix() const { return A_; }
52 const DeviceMatrix<Scalar>& rhs() const { return B_; }
53
54 private:
55 std::reference_wrapper<const DeviceMatrix<Scalar>> A_;
56 std::reference_wrapper<const DeviceMatrix<Scalar>> B_;
57};
58
60template <typename Scalar_, int UpLo_ = Lower>
61class LLTView {
62 public:
63 using Scalar = Scalar_;
64
65 explicit LLTView(const DeviceMatrix<Scalar>& m) : mat_(m) {}
66
68 LltSolveExpr<Scalar, UpLo_> solve(const DeviceMatrix<Scalar>& rhs) const { return {mat_, rhs}; }
69
70 private:
71 std::reference_wrapper<const DeviceMatrix<Scalar>> mat_;
72};
73
75template <typename Scalar_>
76class LUView {
77 public:
78 using Scalar = Scalar_;
79
80 explicit LUView(const DeviceMatrix<Scalar>& m) : mat_(m) {}
81
83 LuSolveExpr<Scalar> solve(const DeviceMatrix<Scalar>& rhs) const { return {mat_, rhs}; }
84
85 private:
86 std::reference_wrapper<const DeviceMatrix<Scalar>> mat_;
87};
88
89} // namespace gpu
90} // namespace Eigen
91
92#endif // EIGEN_GPU_DEVICE_SOLVER_EXPR_H
RAII wrapper for a dense column-major matrix in GPU device memory.
Definition DeviceMatrix.h:122
LltSolveExpr< Scalar, UpLo_ > solve(const DeviceMatrix< Scalar > &rhs) const
Definition DeviceSolverExpr.h:68
LuSolveExpr< Scalar > solve(const DeviceMatrix< Scalar > &rhs) const
Definition DeviceSolverExpr.h:83
Definition DeviceSolverExpr.h:30
Definition DeviceSolverExpr.h:46
Namespace containing all symbols from the Eigen library.