Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
DeviceExpr.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// Lightweight expression types for DeviceMatrix operations.
12//
13// These are NOT Eigen expression templates. Each type maps 1:1 to a single
14// NVIDIA library call (cuBLAS or cuSOLVER). There is no coefficient-level
15// evaluation, no lazy fusion, no packet operations.
16//
17// Expression types:
18// AdjointView<S> — d_A.adjoint() → marks ConjTrans for GEMM
19// TransposeView<S> — d_A.transpose() → marks Trans for GEMM
20// Scaled<Expr> — alpha * expr → carries scalar factor
21// gpu::GemmExpr<Lhs, Rhs> — lhs * rhs → dispatches to cublasLtMatmul (cublasGemmEx fallback)
22
23#ifndef EIGEN_GPU_DEVICE_EXPR_H
24#define EIGEN_GPU_DEVICE_EXPR_H
25
26// IWYU pragma: private
27#include "./InternalHeaderCheck.h"
28
29#include "./CuBlasSupport.h"
30#include "./type_traits.h"
31
32namespace Eigen {
33namespace gpu {
34namespace internal {
35// SFINAE gate for scalar factors: any type convertible to the expression's
36// scalar (so `2 * d_A` and `2.0 * d_cplx` work), except DeviceScalar.
37template <typename T, typename S>
38using require_host_scalar_convertible_t =
39 typename std::enable_if<std::is_convertible<T, S>::value && !is_device_scalar<typename std::decay<T>::type>::value,
40 int>::type;
41
42} // namespace internal
43
45template <typename Scalar_>
46class AdjointView {
47 public:
48 using Scalar = Scalar_;
49 explicit AdjointView(const DeviceMatrix<Scalar>& m) : mat_(m) {}
50 const DeviceMatrix<Scalar>& matrix() const { return mat_; }
51
52 private:
53 const DeviceMatrix<Scalar>& mat_;
54};
55
57template <typename Scalar_>
58class TransposeView {
59 public:
60 using Scalar = Scalar_;
61 explicit TransposeView(const DeviceMatrix<Scalar>& m) : mat_(m) {}
62 const DeviceMatrix<Scalar>& matrix() const { return mat_; }
63
64 private:
65 const DeviceMatrix<Scalar>& mat_;
66};
67
76template <typename Inner_>
77class Scaled {
78 public:
79 using Inner = std::decay_t<Inner_>;
80 using Scalar = internal::scalar_type_t<Inner>;
81 Scaled(Scalar alpha, const Inner_& inner) : alpha_(alpha), inner_(inner) {}
82 Scalar scalar() const { return alpha_; }
83 const Inner_& inner() const { return inner_; }
84
85 private:
86 Scalar alpha_;
87 const Inner_& inner_;
88};
89
91template <typename Lhs, typename Rhs>
92class GemmExpr {
93 public:
94 using Scalar = internal::scalar_type_t<Lhs>;
95 static_assert(std::is_same<Scalar, internal::scalar_type_t<Rhs>>::value,
96 "DeviceMatrix GEMM: LHS and RHS must have the same scalar type");
97
98 GemmExpr(const Lhs& lhs, const Rhs& rhs) : lhs_(lhs), rhs_(rhs) {}
99 const Lhs& lhs() const { return lhs_; }
100 const Rhs& rhs() const { return rhs_; }
101
102 private:
103 // Stored by reference — like Eigen's CPU expression templates, these must
104 // not be captured with auto (the references will dangle). Assign to (or
105 // construct) a DeviceMatrix immediately.
106 const Lhs& lhs_;
107 const Rhs& rhs_;
108};
109
110// Defined after device_expr_traits so it can accept any supported view pair.
111
112// The scalar factor accepts any type convertible to the matrix scalar (int
113// and double literals included), in either operand order. Division by a
114// scalar and unary minus fold into the same Scaled wrapper.
115
116template <typename T, typename S, internal::require_host_scalar_convertible_t<T, S> = 0>
117Scaled<DeviceMatrix<S>> operator*(T alpha, const DeviceMatrix<S>& m) {
118 return {static_cast<S>(alpha), m};
119}
120
121template <typename T, typename S, internal::require_host_scalar_convertible_t<T, S> = 0>
122Scaled<DeviceMatrix<S>> operator*(const DeviceMatrix<S>& m, T alpha) {
123 return {static_cast<S>(alpha), m};
124}
125
126template <typename T, typename S, internal::require_host_scalar_convertible_t<T, S> = 0>
127Scaled<DeviceMatrix<S>> operator/(const DeviceMatrix<S>& m, T alpha) {
128 return {S(1) / static_cast<S>(alpha), m};
129}
130
131template <typename S>
132Scaled<DeviceMatrix<S>> operator-(const DeviceMatrix<S>& m) {
133 return {S(-1), m};
134}
135
136template <typename T, typename S, internal::require_host_scalar_convertible_t<T, S> = 0>
137Scaled<AdjointView<S>> operator*(T alpha, const AdjointView<S>& m) {
138 return {static_cast<S>(alpha), m};
139}
140
141template <typename T, typename S, internal::require_host_scalar_convertible_t<T, S> = 0>
142Scaled<AdjointView<S>> operator*(const AdjointView<S>& m, T alpha) {
143 return {static_cast<S>(alpha), m};
144}
145
146template <typename T, typename S, internal::require_host_scalar_convertible_t<T, S> = 0>
147Scaled<TransposeView<S>> operator*(T alpha, const TransposeView<S>& m) {
148 return {static_cast<S>(alpha), m};
149}
150
151template <typename T, typename S, internal::require_host_scalar_convertible_t<T, S> = 0>
152Scaled<TransposeView<S>> operator*(const TransposeView<S>& m, T alpha) {
153 return {static_cast<S>(alpha), m};
154}
155
156// Rescale / negate an already-scaled expression: T * (alpha * m), -(alpha * m).
157template <typename T, typename Inner,
158 internal::require_host_scalar_convertible_t<T, internal::scalar_type_t<Inner>> = 0>
159Scaled<Inner> operator*(T alpha, const Scaled<Inner>& s) {
160 using S = internal::scalar_type_t<Inner>;
161 return {static_cast<S>(alpha) * s.scalar(), s.inner()};
162}
163
164template <typename Inner>
165Scaled<Inner> operator-(const Scaled<Inner>& s) {
166 using S = internal::scalar_type_t<Inner>;
167 return {S(-1) * s.scalar(), s.inner()};
168}
169
170namespace internal {
171// Default: a DeviceMatrix is NoTrans. Documented on the FwdDecl.h forward declaration.
172template <typename T>
174 static constexpr bool is_device_expr = false;
175};
176
177template <typename Scalar>
178struct device_expr_traits<DeviceMatrix<Scalar>> {
179 using scalar_type = Scalar;
180 static constexpr GpuOp op = GpuOp::NoTrans;
181 static constexpr bool is_device_expr = true;
182 static const DeviceMatrix<Scalar>& matrix(const DeviceMatrix<Scalar>& x) { return x; }
183 static Scalar alpha(const DeviceMatrix<Scalar>&) { return Scalar(1); }
184};
185
186template <typename Scalar>
187struct device_expr_traits<AdjointView<Scalar>> {
188 using scalar_type = Scalar;
189 static constexpr GpuOp op = GpuOp::ConjTrans;
190 static constexpr bool is_device_expr = true;
191 static const DeviceMatrix<Scalar>& matrix(const AdjointView<Scalar>& x) { return x.matrix(); }
192 static Scalar alpha(const AdjointView<Scalar>&) { return Scalar(1); }
193};
194
195template <typename Scalar>
196struct device_expr_traits<TransposeView<Scalar>> {
197 using scalar_type = Scalar;
198 static constexpr GpuOp op = GpuOp::Trans;
199 static constexpr bool is_device_expr = true;
200 static const DeviceMatrix<Scalar>& matrix(const TransposeView<Scalar>& x) { return x.matrix(); }
201 static Scalar alpha(const TransposeView<Scalar>&) { return Scalar(1); }
202};
203
204template <typename Inner>
205struct device_expr_traits<Scaled<Inner>> {
206 using scalar_type = scalar_type_t<Inner>;
207 static constexpr GpuOp op = device_expr_traits<Inner>::op;
208 static constexpr bool is_device_expr = true;
209 static const DeviceMatrix<scalar_type>& matrix(const Scaled<Inner>& x) {
210 return device_expr_traits<Inner>::matrix(x.inner());
211 }
212 static scalar_type alpha(const Scaled<Inner>& x) { return x.scalar() * device_expr_traits<Inner>::alpha(x.inner()); }
213};
214} // namespace internal
215
216template <typename Lhs, typename Rhs,
217 std::enable_if_t<internal::device_expr_traits<Lhs>::is_device_expr &&
218 internal::device_expr_traits<Rhs>::is_device_expr,
219 int> = 0>
220GemmExpr<Lhs, Rhs> operator*(const Lhs& a, const Rhs& b) {
221 return {a, b};
222}
223
230template <typename Scalar_>
231class DeviceScaledDevice {
232 public:
233 using Scalar = Scalar_;
234 DeviceScaledDevice(const DeviceScalar<Scalar>& alpha, const DeviceMatrix<Scalar>& mat) : alpha_(alpha), mat_(mat) {}
235 const DeviceScalar<Scalar>& alpha() const { return alpha_; }
236 const DeviceMatrix<Scalar>& matrix() const { return mat_; }
237
238 private:
239 const DeviceScalar<Scalar>& alpha_;
240 const DeviceMatrix<Scalar>& mat_;
241};
242
243// DeviceScalar * DeviceMatrix → DeviceScaledDevice
244template <typename S>
245DeviceScaledDevice<S> operator*(const DeviceScalar<S>& alpha, const DeviceMatrix<S>& m) {
246 return {alpha, m};
247}
248
249// Captures `DeviceMatrix + Scaled<DeviceMatrix>` (and reverse).
250// Dispatched to geam: C = alpha * A + beta * B.
251//
252// Note: These operator+/- overloads are intentionally free functions on
253// DeviceMatrix, not Eigen expression templates. DeviceMatrix does not inherit
254// from MatrixBase, so there is no ambiguity with Eigen's own operator+/-.
255// If DeviceMatrix is ever made an Eigen expression type, these would need to
256// be revisited.
257
259template <typename Scalar_>
260class DeviceAddExpr {
261 public:
262 using Scalar = Scalar_;
263 DeviceAddExpr(Scalar alpha, const DeviceMatrix<Scalar>& A, Scalar beta, const DeviceMatrix<Scalar>& B)
264 : alpha_(alpha), A_(A), beta_(beta), B_(B) {}
265 Scalar alpha() const { return alpha_; }
266 Scalar beta() const { return beta_; }
267 const DeviceMatrix<Scalar>& A() const { return A_; }
268 const DeviceMatrix<Scalar>& B() const { return B_; }
269
270 private:
271 Scalar alpha_;
272 const DeviceMatrix<Scalar>& A_;
273 Scalar beta_;
274 const DeviceMatrix<Scalar>& B_;
275};
276
277// DeviceMatrix + DeviceMatrix → DeviceAddExpr (alpha=1, beta=1)
278template <typename S>
279DeviceAddExpr<S> operator+(const DeviceMatrix<S>& a, const DeviceMatrix<S>& b) {
280 return {S(1), a, S(1), b};
281}
282
283// DeviceMatrix + Scaled<DeviceMatrix> → DeviceAddExpr (alpha=1, beta=scaled)
284template <typename S>
285DeviceAddExpr<S> operator+(const DeviceMatrix<S>& a, const Scaled<DeviceMatrix<S>>& b) {
286 return {S(1), a, b.scalar(), b.inner()};
287}
288
289// Scaled<DeviceMatrix> + DeviceMatrix → DeviceAddExpr (alpha=scaled, beta=1)
290template <typename S>
291DeviceAddExpr<S> operator+(const Scaled<DeviceMatrix<S>>& a, const DeviceMatrix<S>& b) {
292 return {a.scalar(), a.inner(), S(1), b};
293}
294
295// DeviceMatrix - DeviceMatrix → DeviceAddExpr (alpha=1, beta=-1)
296template <typename S>
297DeviceAddExpr<S> operator-(const DeviceMatrix<S>& a, const DeviceMatrix<S>& b) {
298 return {S(1), a, S(-1), b};
299}
300
301// DeviceMatrix - Scaled<DeviceMatrix> → DeviceAddExpr (alpha=1, beta=-scaled)
302template <typename S>
303DeviceAddExpr<S> operator-(const DeviceMatrix<S>& a, const Scaled<DeviceMatrix<S>>& b) {
304 return {S(1), a, -b.scalar(), b.inner()};
305}
306
307// Scaled<DeviceMatrix> - DeviceMatrix → DeviceAddExpr (alpha=scaled, beta=-1)
308template <typename S>
309DeviceAddExpr<S> operator-(const Scaled<DeviceMatrix<S>>& a, const DeviceMatrix<S>& b) {
310 return {a.scalar(), a.inner(), S(-1), b};
311}
312
313// Scaled<DeviceMatrix> ± Scaled<DeviceMatrix> → DeviceAddExpr
314template <typename S>
315DeviceAddExpr<S> operator+(const Scaled<DeviceMatrix<S>>& a, const Scaled<DeviceMatrix<S>>& b) {
316 return {a.scalar(), a.inner(), b.scalar(), b.inner()};
317}
318
319template <typename S>
320DeviceAddExpr<S> operator-(const Scaled<DeviceMatrix<S>>& a, const Scaled<DeviceMatrix<S>>& b) {
321 return {a.scalar(), a.inner(), -b.scalar(), b.inner()};
322}
323} // namespace gpu
324} // namespace Eigen
325
326#endif // EIGEN_GPU_DEVICE_EXPR_H
View returned by DeviceMatrix::adjoint(); maps to the cuBLAS conjugate-transpose operand flag.
Definition DeviceExpr.h:46
Linear combination of two device matrices.
Definition DeviceExpr.h:260
RAII wrapper for a dense column-major matrix in GPU device memory.
Definition DeviceMatrix.h:122
RAII wrapper for a scalar in GPU device memory.
Definition DeviceScalar.h:33
Expression that scales a device matrix by a DeviceScalar.
Definition DeviceExpr.h:231
Expression returned by operator*(lhs_expr, rhs_expr), dispatched to cuBLAS GEMM.
Definition DeviceExpr.h:92
Expression returned by operator*(Scalar, DeviceMatrix/View), carrying the scalar factor.
Definition DeviceExpr.h:77
View returned by DeviceMatrix::transpose(); maps to the cuBLAS transpose operand flag.
Definition DeviceExpr.h:58
Namespace containing all symbols from the Eigen library.
Describes GPU device expression types.
Definition DeviceExpr.h:173
Definition type_traits.h:873