Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
DeviceBlasExpr.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// BLAS Level 3 expression types for gpu::DeviceMatrix beyond GEMM: triangular
12// solve, self-adjoint multiply, and rank-k update.
13
14#ifndef EIGEN_GPU_DEVICE_BLAS_EXPR_H
15#define EIGEN_GPU_DEVICE_BLAS_EXPR_H
16
17// IWYU pragma: private
18#include "./InternalHeaderCheck.h"
19
20#include <functional>
21
22#include "./FwdDecl.h"
23
24namespace Eigen {
25namespace gpu {
26
28template <typename Scalar_, int UpLo_>
29class TriangularView {
30 public:
31 using Scalar = Scalar_;
32 static constexpr int UpLo = UpLo_;
33
34 explicit TriangularView(const DeviceMatrix<Scalar>& m) : mat_(m) {}
35 const DeviceMatrix<Scalar>& matrix() const { return mat_; }
36
37 TrsmExpr<Scalar, UpLo_> solve(const DeviceMatrix<Scalar>& rhs) const { return {mat_, rhs}; }
38
39 private:
40 std::reference_wrapper<const DeviceMatrix<Scalar>> mat_;
41};
42
44template <typename Scalar_, int UpLo_>
45class TrsmExpr {
46 public:
47 using Scalar = Scalar_;
48 static constexpr int UpLo = UpLo_;
49
50 TrsmExpr(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_>
61class SelfAdjointView {
62 public:
63 using Scalar = Scalar_;
64 using RealScalar = typename NumTraits<Scalar>::Real;
65 static constexpr int UpLo = UpLo_;
66
67 explicit SelfAdjointView(DeviceMatrix<Scalar>& m) : mat_(m) {}
68 const DeviceMatrix<Scalar>& matrix() const { return mat_; }
69 DeviceMatrix<Scalar>& matrix() { return mat_; }
70
73 void rankUpdate(const DeviceMatrix<Scalar>& A, RealScalar alpha = RealScalar(1));
74
75 private:
76 std::reference_wrapper<DeviceMatrix<Scalar>> mat_;
77};
78
80template <typename Scalar_, int UpLo_>
81class ConstSelfAdjointView {
82 public:
83 using Scalar = Scalar_;
84 static constexpr int UpLo = UpLo_;
85
86 explicit ConstSelfAdjointView(const DeviceMatrix<Scalar>& m) : mat_(m) {}
87 const DeviceMatrix<Scalar>& matrix() const { return mat_; }
88
89 private:
90 std::reference_wrapper<const DeviceMatrix<Scalar>> mat_;
91};
92
95template <typename Scalar_, int UpLo_>
96class SymmExpr {
97 public:
98 using Scalar = Scalar_;
99 static constexpr int UpLo = UpLo_;
100
101 SymmExpr(const DeviceMatrix<Scalar>& A, const DeviceMatrix<Scalar>& B) : A_(A), B_(B) {}
102 const DeviceMatrix<Scalar>& matrix() const { return A_; }
103 const DeviceMatrix<Scalar>& rhs() const { return B_; }
104
105 private:
106 std::reference_wrapper<const DeviceMatrix<Scalar>> A_;
107 std::reference_wrapper<const DeviceMatrix<Scalar>> B_;
108};
109
110template <typename S, int UpLo>
111SymmExpr<S, UpLo> operator*(const SelfAdjointView<S, UpLo>& a, const DeviceMatrix<S>& b) {
112 return {a.matrix(), b};
113}
114template <typename S, int UpLo>
115SymmExpr<S, UpLo> operator*(const ConstSelfAdjointView<S, UpLo>& a, const DeviceMatrix<S>& b) {
116 return {a.matrix(), b};
117}
118
121template <typename Scalar_, int UpLo_>
122class SyrkExpr {
123 public:
124 using Scalar = Scalar_;
125 static constexpr int UpLo = UpLo_;
126
127 SyrkExpr(const DeviceMatrix<Scalar>& A) : A_(A) {}
128 const DeviceMatrix<Scalar>& matrix() const { return A_; }
129
130 private:
131 std::reference_wrapper<const DeviceMatrix<Scalar>> A_;
132};
133
134} // namespace gpu
135} // namespace Eigen
136
137#endif // EIGEN_GPU_DEVICE_BLAS_EXPR_H
RAII wrapper for a dense column-major matrix in GPU device memory.
Definition DeviceMatrix.h:122
Definition DeviceBlasExpr.h:61
void rankUpdate(const DeviceMatrix< Scalar > &A, RealScalar alpha=RealScalar(1))
Definition DeviceDispatch.h:543
Definition DeviceBlasExpr.h:96
Definition DeviceBlasExpr.h:45
Namespace containing all symbols from the Eigen library.