Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
DeviceDispatch.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// Dispatch functions mapping DeviceMatrix expressions to NVIDIA library calls,
12// plus the DeviceMatrix members that need a complete gpu::Context. The
13// expression argument selects the dispatch() overload.
14
15#ifndef EIGEN_GPU_DEVICE_DISPATCH_H
16#define EIGEN_GPU_DEVICE_DISPATCH_H
17
18// IWYU pragma: private
19#include "./InternalHeaderCheck.h"
20
21#include <cstdint>
22
23#include "./DeviceExpr.h"
24#include "./DeviceBlasExpr.h"
25#include "./DeviceSolverExpr.h"
26#include "./GpuContext.h"
27#include "./CuSolverSupport.h"
28
29namespace Eigen {
30namespace gpu {
31namespace internal {
32template <typename Scalar>
33bool aliases_device_memory(const DeviceMatrix<Scalar>& a, const DeviceMatrix<Scalar>& b) {
34 return a.data() != nullptr && a.data() == b.data();
35}
36
37template <typename Lhs, typename Rhs>
38void dispatch(Context& ctx, DeviceMatrix<scalar_type_t<Lhs>>& dst, const GemmExpr<Lhs, Rhs>& expr,
39 scalar_type_t<Lhs> beta_val, scalar_type_t<Lhs> alpha_scale = scalar_type_t<Lhs>(1)) {
40 using Scalar = scalar_type_t<Lhs>;
41 using traits_lhs = device_expr_traits<Lhs>;
42 using traits_rhs = device_expr_traits<Rhs>;
43
44 const DeviceMatrix<Scalar>& A = traits_lhs::matrix(expr.lhs());
45 const DeviceMatrix<Scalar>& B = traits_rhs::matrix(expr.rhs());
46
47 // cuBLAS leaves C aliasing A or B undefined.
48 eigen_assert(!aliases_device_memory(dst, A) && "GEMM: output aliases left operand (use a temporary)");
49 eigen_assert(!aliases_device_memory(dst, B) && "GEMM: output aliases right operand (use a temporary)");
50
51 constexpr cublasOperation_t transA = to_cublas_op(traits_lhs::op);
52 constexpr cublasOperation_t transB = to_cublas_op(traits_rhs::op);
53
54 const int64_t m = (traits_lhs::op == GpuOp::NoTrans) ? A.rows() : A.cols();
55 const int64_t k = (traits_lhs::op == GpuOp::NoTrans) ? A.cols() : A.rows();
56 const int64_t n = (traits_rhs::op == GpuOp::NoTrans) ? B.cols() : B.rows();
57 const int64_t rhs_k = (traits_rhs::op == GpuOp::NoTrans) ? B.rows() : B.cols();
58
59 eigen_assert(k == rhs_k && "DeviceMatrix GEMM dimension mismatch");
60
61 const int64_t lda = A.rows();
62 const int64_t ldb = B.rows();
63
64 if (!dst.empty()) {
65 dst.waitReady(ctx.stream());
66 }
67
68 const bool resized = dst.empty() || dst.rows() != m || dst.cols() != n;
69 if (resized) {
70 dst.resize(m, n);
71 }
72
73 // cuBLAS rejects ld = 0 (ld >= max(1, rows)) and the cublasLt heuristic
74 // faults on n == 0, so empty products stop here. For k == 0, C = beta * C
75 // with beta in {0, 1}; a resized C starts from zero.
76 if (m == 0 || n == 0 || k == 0) {
77 if ((resized || beta_val == Scalar(0)) && dst.sizeInBytes() > 0) {
78 EIGEN_CUDA_RUNTIME_CHECK(cudaMemsetAsync(dst.data(), 0, dst.sizeInBytes(), ctx.stream()));
79 dst.recordReady(ctx.stream());
80 }
81 return;
82 }
83 const int64_t ldc = dst.rows();
84
85 Scalar alpha_local = alpha_scale * traits_lhs::alpha(expr.lhs()) * traits_rhs::alpha(expr.rhs());
86
87 A.waitReady(ctx.stream());
88 B.waitReady(ctx.stream());
89
90 if (resized && beta_val != Scalar(0) && dst.sizeInBytes() > 0) {
91 EIGEN_CUDA_RUNTIME_CHECK(cudaMemsetAsync(dst.data(), 0, dst.sizeInBytes(), ctx.stream()));
92 }
93
94 cublaslt_gemm(ctx.cublasLtHandle(), ctx.cublasHandle(), transA, transB, m, n, k, &alpha_local, A.data(), lda,
95 B.data(), ldb, &beta_val, dst.data(), ldc, ctx.gemmWorkspace(), ctx.gemmPlanCache(),
96 ctx.cublasLtMaxWorkspaceBytes(), ctx.stream());
97
98 dst.recordReady(ctx.stream());
99}
100
101// Debug-build status check shared by the one-shot solver dispatches: syncs
102// the stream and asserts on the two info words (factorize, solve). Release
103// builds skip both the check and the sync — one-shot expressions are then
104// fully async with no failure detection; use gpu::LLT / gpu::LU + info()
105// when failures must be detected.
106inline void oneshot_check_info(Context& ctx, OneShotSolverScratch& scratch, const char* what) {
107#ifdef EIGEN_NO_DEBUG
108 EIGEN_UNUSED_VARIABLE(ctx);
109 EIGEN_UNUSED_VARIABLE(scratch);
110 EIGEN_UNUSED_VARIABLE(what);
111#else
112 int* info_words = static_cast<int*>(scratch.h_info.get());
113 EIGEN_CUDA_RUNTIME_CHECK(
114 cudaMemcpyAsync(info_words, scratch.d_info.get(), kOneShotInfoBytes, cudaMemcpyDeviceToHost, ctx.stream()));
115 EIGEN_CUDA_RUNTIME_CHECK(cudaStreamSynchronize(ctx.stream()));
116 eigen_assert(info_words[0] == 0 && "cuSOLVER one-shot factorization failed" && what);
117 eigen_assert(info_words[1] == 0 && "cuSOLVER one-shot solve failed" && what);
118 EIGEN_UNUSED_VARIABLE(what);
119#endif
120}
121
122template <typename Scalar, int UpLo>
123void dispatch(Context& ctx, DeviceMatrix<Scalar>& dst, const LltSolveExpr<Scalar, UpLo>& expr) {
124 const DeviceMatrix<Scalar>& A = expr.matrix();
125 const DeviceMatrix<Scalar>& B = expr.rhs();
126
127 eigen_assert(A.rows() == A.cols() && "LLT requires a square matrix");
128 eigen_assert(B.rows() == A.rows() && "LLT solve: RHS rows must match matrix size");
129
130 if (A.rows() == 0 || B.cols() == 0) {
131 if (!dst.empty()) dst.waitReady(ctx.stream());
132 dst.resize(A.rows(), B.cols());
133 return;
134 }
135
136 A.waitReady(ctx.stream());
137 B.waitReady(ctx.stream());
138 if (!dst.empty()) dst.waitReady(ctx.stream());
139
140 // thread_local: must outlive the async kernels (no end-of-call sync), and
141 // only TUs that instantiate the one-shot path pull in cuSOLVER symbols.
142 static thread_local CusolverParams params;
143 constexpr cublasFillMode_t uplo = cusolver_fill_mode<UpLo>::value;
144 const int64_t n = static_cast<int64_t>(A.rows());
145 constexpr cudaDataType_t dtype = cuda_data_type<Scalar>::value;
146 OneShotSolverScratch& scratch = ctx.oneshotSolverScratch();
147 {
148 const size_t mat_bytes = A.sizeInBytes();
149 // Context-owned grow-only scratch: no per-call allocation, no end-of-call sync.
150 ensure_sized(scratch.d_factor, mat_bytes);
151 EIGEN_CUDA_RUNTIME_CHECK(
152 cudaMemcpyAsync(scratch.d_factor.get(), A.data(), mat_bytes, cudaMemcpyDeviceToDevice, ctx.stream()));
153 }
154 const int64_t lda = static_cast<int64_t>(A.rows());
155 size_t dev_ws = 0;
156 size_t host_ws = 0;
157 EIGEN_CUSOLVER_CHECK(cusolverDnXpotrf_bufferSize(ctx.cusolverHandle(), params.p, uplo, n, dtype,
158 scratch.d_factor.get(), lda, dtype, &dev_ws, &host_ws));
159 ensure_sized(scratch.d_workspace, dev_ws);
160 if (scratch.h_workspace.size() < host_ws) scratch.h_workspace.resize(host_ws);
161 // Two info slots (potrf, potrs) so both kernels queue back-to-back. If potrf
162 // fails, potrs runs on garbage but the debug check catches both at once.
163 int* d_info_potrf = static_cast<int*>(scratch.d_info.get());
164 int* d_info_potrs = d_info_potrf + 1;
165 EIGEN_CUSOLVER_CHECK(cusolverDnXpotrf(ctx.cusolverHandle(), params.p, uplo, n, dtype, scratch.d_factor.get(), lda,
166 dtype, scratch.d_workspace.get(), dev_ws,
167 host_ws > 0 ? scratch.h_workspace.data() : nullptr, host_ws, d_info_potrf));
168
169 dst.resize(n, B.cols());
170 const size_t rhs_bytes = B.sizeInBytes();
171 EIGEN_CUDA_RUNTIME_CHECK(cudaMemcpyAsync(dst.data(), B.data(), rhs_bytes, cudaMemcpyDeviceToDevice, ctx.stream()));
172
173 const int64_t nrhs = static_cast<int64_t>(B.cols());
174 EIGEN_CUSOLVER_CHECK(cusolverDnXpotrs(ctx.cusolverHandle(), params.p, uplo, n, nrhs, dtype, scratch.d_factor.get(),
175 lda, dtype, dst.data(), static_cast<int64_t>(dst.rows()), d_info_potrs));
176 oneshot_check_info(ctx, scratch, "llt");
177 dst.recordReady(ctx.stream());
178}
179
180template <typename Scalar>
181void dispatch(Context& ctx, DeviceMatrix<Scalar>& dst, const LuSolveExpr<Scalar>& expr) {
182 const DeviceMatrix<Scalar>& A = expr.matrix();
183 const DeviceMatrix<Scalar>& B = expr.rhs();
184
185 eigen_assert(A.rows() == A.cols() && "LU requires a square matrix");
186 eigen_assert(B.rows() == A.rows() && "LU solve: RHS rows must match matrix size");
187
188 if (A.rows() == 0 || B.cols() == 0) {
189 if (!dst.empty()) dst.waitReady(ctx.stream());
190 dst.resize(A.rows(), B.cols());
191 return;
192 }
193
194 A.waitReady(ctx.stream());
195 B.waitReady(ctx.stream());
196 if (!dst.empty()) dst.waitReady(ctx.stream());
197
198 // thread_local: must outlive the async kernels (no end-of-call sync), and
199 // only TUs that instantiate the one-shot path pull in cuSOLVER symbols.
200 static thread_local CusolverParams params;
201 const int64_t n = static_cast<int64_t>(A.rows());
202 constexpr cudaDataType_t dtype = cuda_data_type<Scalar>::value;
203 OneShotSolverScratch& scratch = ctx.oneshotSolverScratch();
204 {
205 const size_t mat_bytes = A.sizeInBytes();
206 // Context-owned grow-only scratch: no per-call allocation, no end-of-call sync.
207 ensure_sized(scratch.d_factor, mat_bytes);
208 EIGEN_CUDA_RUNTIME_CHECK(
209 cudaMemcpyAsync(scratch.d_factor.get(), A.data(), mat_bytes, cudaMemcpyDeviceToDevice, ctx.stream()));
210 }
211 ensure_sized(scratch.d_ipiv, static_cast<size_t>(n) * sizeof(int64_t));
212 const int64_t lda = static_cast<int64_t>(A.rows());
213 size_t dev_ws = 0;
214 size_t host_ws = 0;
215 EIGEN_CUSOLVER_CHECK(cusolverDnXgetrf_bufferSize(ctx.cusolverHandle(), params.p, n, n, dtype, scratch.d_factor.get(),
216 lda, dtype, &dev_ws, &host_ws));
217 ensure_sized(scratch.d_workspace, dev_ws);
218 if (scratch.h_workspace.size() < host_ws) scratch.h_workspace.resize(host_ws);
219 int* d_info_getrf = static_cast<int*>(scratch.d_info.get());
220 int* d_info_getrs = d_info_getrf + 1;
221 EIGEN_CUSOLVER_CHECK(cusolverDnXgetrf(ctx.cusolverHandle(), params.p, n, n, dtype, scratch.d_factor.get(), lda,
222 static_cast<int64_t*>(scratch.d_ipiv.get()), dtype, scratch.d_workspace.get(),
223 dev_ws, host_ws > 0 ? scratch.h_workspace.data() : nullptr, host_ws,
224 d_info_getrf));
225
226 dst.resize(n, B.cols());
227 const size_t rhs_bytes = B.sizeInBytes();
228 EIGEN_CUDA_RUNTIME_CHECK(cudaMemcpyAsync(dst.data(), B.data(), rhs_bytes, cudaMemcpyDeviceToDevice, ctx.stream()));
229
230 const int64_t nrhs = static_cast<int64_t>(B.cols());
231 EIGEN_CUSOLVER_CHECK(cusolverDnXgetrs(ctx.cusolverHandle(), params.p, CUBLAS_OP_N, n, nrhs, dtype,
232 scratch.d_factor.get(), lda, static_cast<const int64_t*>(scratch.d_ipiv.get()),
233 dtype, dst.data(), static_cast<int64_t>(dst.rows()), d_info_getrs));
234 oneshot_check_info(ctx, scratch, "lu");
235 dst.recordReady(ctx.stream());
236}
237
238template <typename Scalar, int UpLo>
239void dispatch(Context& ctx, DeviceMatrix<Scalar>& dst, const TrsmExpr<Scalar, UpLo>& expr) {
240 const DeviceMatrix<Scalar>& A = expr.matrix();
241 const DeviceMatrix<Scalar>& B = expr.rhs();
242
243 eigen_assert(A.rows() == A.cols() && "TRSM requires a square triangular matrix");
244 eigen_assert(B.rows() == A.rows() && "TRSM: RHS rows must match matrix size");
245
246 const int64_t n = A.rows();
247 const int64_t nrhs = B.cols();
248
249 if (n == 0 || nrhs == 0) {
250 if (!dst.empty()) dst.waitReady(ctx.stream());
251 dst.resize(n, B.cols());
252 return;
253 }
254
255 A.waitReady(ctx.stream());
256 B.waitReady(ctx.stream());
257 eigen_assert(!aliases_device_memory(dst, A) && "DeviceMatrix TRSM destination aliases triangular operand");
258 eigen_assert(!aliases_device_memory(dst, B) && "DeviceMatrix TRSM destination aliases RHS operand");
259 if (!dst.empty()) dst.waitReady(ctx.stream());
260
261 dst.resize(n, B.cols());
262 const size_t rhs_bytes = static_cast<size_t>(dst.rows()) * static_cast<size_t>(nrhs) * sizeof(Scalar);
263 EIGEN_CUDA_RUNTIME_CHECK(cudaMemcpyAsync(dst.data(), B.data(), rhs_bytes, cudaMemcpyDeviceToDevice, ctx.stream()));
264
265 constexpr cublasFillMode_t uplo = (UpLo == Lower) ? CUBLAS_FILL_MODE_LOWER : CUBLAS_FILL_MODE_UPPER;
266 Scalar alpha(1);
267
268 EIGEN_CUBLAS_CHECK(cublasXtrsm(ctx.cublasHandle(), CUBLAS_SIDE_LEFT, uplo, CUBLAS_OP_N, CUBLAS_DIAG_NON_UNIT, n, nrhs,
269 &alpha, A.data(), A.rows(), dst.data(), dst.rows()));
270
271 dst.recordReady(ctx.stream());
272}
273
274template <typename Scalar, int UpLo>
275void dispatch(Context& ctx, DeviceMatrix<Scalar>& dst, const SymmExpr<Scalar, UpLo>& expr) {
276 const DeviceMatrix<Scalar>& A = expr.matrix();
277 const DeviceMatrix<Scalar>& B = expr.rhs();
278
279 eigen_assert(A.rows() == A.cols() && "SYMM requires a square matrix");
280 eigen_assert(B.rows() == A.rows() && "SYMM: RHS rows must match matrix size");
281
282 const int64_t m = A.rows();
283 const int64_t n = B.cols();
284
285 if (m == 0 || n == 0) {
286 if (!dst.empty()) dst.waitReady(ctx.stream());
287 dst.resize(m, B.cols());
288 return;
289 }
290
291 A.waitReady(ctx.stream());
292 B.waitReady(ctx.stream());
293 eigen_assert(!aliases_device_memory(dst, A) && "DeviceMatrix SYMM destination aliases self-adjoint operand");
294 eigen_assert(!aliases_device_memory(dst, B) && "DeviceMatrix SYMM destination aliases RHS operand");
295 if (!dst.empty()) dst.waitReady(ctx.stream());
296
297 dst.resize(m, n);
298
299 constexpr cublasFillMode_t uplo = (UpLo == Lower) ? CUBLAS_FILL_MODE_LOWER : CUBLAS_FILL_MODE_UPPER;
300 const Scalar one(1), zero(0);
301
302 EIGEN_CUBLAS_CHECK(cublasXsymm(ctx.cublasHandle(), CUBLAS_SIDE_LEFT, uplo, m, n, &one, A.data(), A.rows(), B.data(),
303 B.rows(), &zero, dst.data(), dst.rows()));
304
305 dst.recordReady(ctx.stream());
306}
307
308template <typename Scalar, int UpLo>
309void dispatch(Context& ctx, DeviceMatrix<Scalar>& dst, const SyrkExpr<Scalar, UpLo>& expr,
310 typename NumTraits<Scalar>::Real alpha_val, typename NumTraits<Scalar>::Real beta_val) {
311 using RealScalar = typename NumTraits<Scalar>::Real;
312 const DeviceMatrix<Scalar>& A = expr.matrix();
313
314 const int64_t n = A.rows();
315 const int64_t k = A.cols();
316
317 if (n == 0) {
318 if (!dst.empty()) dst.waitReady(ctx.stream());
319 dst.resize(0, 0);
320 return;
321 }
322
323 A.waitReady(ctx.stream());
324 eigen_assert(!aliases_device_memory(dst, A) && "DeviceMatrix SYRK destination aliases input operand");
325 if (!dst.empty()) dst.waitReady(ctx.stream());
326
327 if (dst.empty() || dst.rows() != n || dst.cols() != n) {
328 dst.resize(n, n);
329 if (beta_val != RealScalar(0)) {
330 EIGEN_CUDA_RUNTIME_CHECK(cudaMemsetAsync(dst.data(), 0, dst.sizeInBytes(), ctx.stream()));
331 }
332 }
333
334 constexpr cublasFillMode_t uplo = (UpLo == Lower) ? CUBLAS_FILL_MODE_LOWER : CUBLAS_FILL_MODE_UPPER;
335
336 EIGEN_CUBLAS_CHECK(cublasXsyrk(ctx.cublasHandle(), uplo, CUBLAS_OP_N, n, k, &alpha_val, A.data(), A.rows(), &beta_val,
337 dst.data(), dst.rows()));
338
339 dst.recordReady(ctx.stream());
340}
341
342// DeviceAddExpr → cublasXgeam: dst = alpha * A + beta * B. Safe when dst
343// aliases A and/or B (geam supports in-place operation with equal leading
344// dimensions, which always holds here since DeviceMatrix is fully dense).
345
346template <typename Scalar>
347void dispatch(Context& ctx, DeviceMatrix<Scalar>& dst, const DeviceAddExpr<Scalar>& expr) {
348 const DeviceMatrix<Scalar>& A = expr.A();
349 const DeviceMatrix<Scalar>& B = expr.B();
350 eigen_assert(A.rows() == B.rows() && A.cols() == B.cols());
351 const int64_t m = A.rows();
352 const int64_t n = A.cols();
353 // Wait on dst before resize — resize may free the old buffer while another
354 // stream is still reading it.
355 if (!dst.empty()) dst.waitReady(ctx.stream());
356 dst.resize(A.rows(), A.cols());
357 if (m > 0 && n > 0) {
358 A.waitReady(ctx.stream());
359 B.waitReady(ctx.stream());
360 const Scalar alpha_val = expr.alpha(), beta_val = expr.beta();
361 EIGEN_CUBLAS_CHECK(cublasXgeam(ctx.cublasHandle(), CUBLAS_OP_N, CUBLAS_OP_N, m, n, &alpha_val, A.data(), m,
362 &beta_val, B.data(), m, dst.data(), m));
363 dst.recordReady(ctx.stream());
364 }
365}
366} // namespace internal
367
368template <typename Scalar_>
369class Assignment {
370 public:
371 using Scalar = Scalar_;
372
373 Assignment(DeviceMatrix<Scalar>& dst, Context& ctx) : dst_(dst), ctx_(ctx) {}
374
375 template <typename Lhs, typename Rhs>
376 DeviceMatrix<Scalar>& operator=(const GemmExpr<Lhs, Rhs>& expr) {
377 internal::dispatch(ctx_, dst_, expr, Scalar(0));
378 return dst_;
379 }
380
381 template <typename Lhs, typename Rhs>
382 DeviceMatrix<Scalar>& operator+=(const GemmExpr<Lhs, Rhs>& expr) {
383 internal::dispatch(ctx_, dst_, expr, Scalar(1));
384 return dst_;
385 }
386
387 template <typename Lhs, typename Rhs>
388 DeviceMatrix<Scalar>& operator-=(const GemmExpr<Lhs, Rhs>& expr) {
389 internal::dispatch(ctx_, dst_, expr, Scalar(1), Scalar(-1));
390 return dst_;
391 }
392
393 template <int UpLo>
394 DeviceMatrix<Scalar>& operator=(const LltSolveExpr<Scalar, UpLo>& expr) {
395 internal::dispatch(ctx_, dst_, expr);
396 return dst_;
397 }
398
399 DeviceMatrix<Scalar>& operator=(const LuSolveExpr<Scalar>& expr) {
400 internal::dispatch(ctx_, dst_, expr);
401 return dst_;
402 }
403
404 template <int UpLo>
405 DeviceMatrix<Scalar>& operator=(const TrsmExpr<Scalar, UpLo>& expr) {
406 internal::dispatch(ctx_, dst_, expr);
407 return dst_;
408 }
409
410 template <int UpLo>
411 DeviceMatrix<Scalar>& operator=(const SymmExpr<Scalar, UpLo>& expr) {
412 internal::dispatch(ctx_, dst_, expr);
413 return dst_;
414 }
415
416 DeviceMatrix<Scalar>& operator=(const DeviceAddExpr<Scalar>& expr) {
417 internal::dispatch(ctx_, dst_, expr);
418 return dst_;
419 }
420
421 DeviceMatrix<Scalar>& operator=(const Scaled<DeviceMatrix<Scalar>>& expr) {
422 // geam with beta == 0: cuBLAS documents B as unread, so pass A twice.
423 internal::dispatch(ctx_, dst_, DeviceAddExpr<Scalar>(expr.scalar(), expr.inner(), Scalar(0), expr.inner()));
424 return dst_;
425 }
426
427 template <typename Expr>
428 DeviceMatrix<Scalar>& operator=(const Expr&) {
429 static_assert(sizeof(Expr) == 0,
430 "DeviceMatrix expression not supported: no cuBLAS/cuSOLVER mapping. "
431 "Supported: GEMM (A*B), geam (A + alpha*B, alpha*A), "
432 "TRSM (.triangularView().solve()), SYMM (.selfadjointView()*B), "
433 "LLT (.llt().solve()), LU (.lu().solve()).");
434 return dst_;
435 }
436
437 private:
438 DeviceMatrix<Scalar>& dst_;
439 Context& ctx_;
440};
441
442// The definitions below call Context::threadLocal(), so they cannot live in
443// DeviceMatrix.h, where Context is still incomplete.
444
445template <typename Scalar_>
446template <typename Lhs, typename Rhs>
447DeviceMatrix<Scalar_>& DeviceMatrix<Scalar_>::operator=(const GemmExpr<Lhs, Rhs>& expr) {
449 return *this;
450}
451
452template <typename Scalar_>
453template <typename Lhs, typename Rhs>
454DeviceMatrix<Scalar_>& DeviceMatrix<Scalar_>::operator+=(const GemmExpr<Lhs, Rhs>& expr) {
455 device(Context::threadLocal()) += expr;
456 return *this;
457}
459template <typename Scalar_>
460template <typename Lhs, typename Rhs>
462 device(Context::threadLocal()) -= expr;
463 return *this;
465
466template <typename Scalar_>
467template <int UpLo>
468DeviceMatrix<Scalar_>& DeviceMatrix<Scalar_>::operator=(const LltSolveExpr<Scalar_, UpLo>& expr) {
469 device(Context::threadLocal()) = expr;
470 return *this;
471}
472
473template <typename Scalar_>
474DeviceMatrix<Scalar_>& DeviceMatrix<Scalar_>::operator=(const LuSolveExpr<Scalar_>& expr) {
476 return *this;
478
479template <typename Scalar_>
480template <int UpLo>
481DeviceMatrix<Scalar_>& DeviceMatrix<Scalar_>::operator=(const TrsmExpr<Scalar_, UpLo>& expr) {
483 return *this;
484}
486template <typename Scalar_>
487template <int UpLo>
488DeviceMatrix<Scalar_>& DeviceMatrix<Scalar_>::operator=(const SymmExpr<Scalar_, UpLo>& expr) {
490 return *this;
491}
492
493template <typename Scalar_>
494DeviceMatrix<Scalar_>& DeviceMatrix<Scalar_>::operator=(const Scaled<DeviceMatrix>& expr) {
496 return *this;
497}
498
499// Enable copy-initialization straight from an expression, e.g.
500// DeviceMatrix<double> d_C = d_A * d_B;
501// Each default-constructs and delegates to the matching operator=.
503template <typename Scalar_>
504template <typename Lhs, typename Rhs>
506 *this = expr;
507}
509template <typename Scalar_>
511 *this = expr;
512}
513
514template <typename Scalar_>
516 *this = expr;
518
519template <typename Scalar_>
520template <int UpLo>
522 *this = expr;
523}
525template <typename Scalar_>
527 *this = expr;
528}
529
530template <typename Scalar_>
531template <int UpLo>
533 *this = expr;
534}
535
536template <typename Scalar_>
537template <int UpLo>
539 *this = expr;
540}
541
542template <typename Scalar_, int UpLo_>
545 RealScalar beta = matrix().empty() ? RealScalar(0) : RealScalar(1);
546 internal::dispatch(Context::threadLocal(), matrix(), expr, alpha, beta);
547}
548
549namespace internal {
550// Runs `f` with the handle temporarily in CUBLAS_POINTER_MODE_DEVICE, restoring
551// the caller's mode afterwards, also when a failure in `f` throws: the handle
552// would otherwise read host scalars of later calls as device pointers.
553template <typename F>
554void with_device_pointer_mode(cublasHandle_t h, F&& f) {
555 struct RestoreOnThrow {
556 cublasHandle_t handle;
557 cublasPointerMode_t mode;
558 bool armed;
559 ~RestoreOnThrow() {
560 if (armed) (void)cublasSetPointerMode(handle, mode); // unchecked: may run during unwinding
561 }
562 };
563 cublasPointerMode_t prev;
564 EIGEN_CUBLAS_CHECK(cublasGetPointerMode(h, &prev));
565 EIGEN_CUBLAS_CHECK(cublasSetPointerMode(h, CUBLAS_POINTER_MODE_DEVICE));
566 RestoreOnThrow restore{h, prev, true};
567 f();
568 restore.armed = false;
569 EIGEN_CUBLAS_CHECK(cublasSetPointerMode(h, prev));
570}
571} // namespace internal
572
573// The reductions below (dot, norm, squaredNorm) run under
574// CUBLAS_POINTER_MODE_DEVICE: the scalar result is written to device memory and
575// stays there until DeviceScalar's conversion to Scalar syncs and reads it.
576
577namespace internal {
578inline int64_t blas1_size(Index rows, Index cols) { return static_cast<int64_t>(rows) * static_cast<int64_t>(cols); }
579} // namespace internal
580
581template <typename Scalar_>
583 const int64_t n = internal::blas1_size(rows_, cols_);
584 eigen_assert(n == internal::blas1_size(other.rows_, other.cols_));
585 eigen_assert(result.stream() == ctx.stream() && "DeviceMatrix::dot: result must live on ctx's stream");
586 if (n > 0) {
587 waitReady(ctx.stream());
588 other.waitReady(ctx.stream());
589 internal::with_device_pointer_mode(ctx.cublasHandle(), [&] {
590 EIGEN_CUBLAS_CHECK(
591 internal::cublasXdot(ctx.cublasHandle(), n, data_.get(), 1, other.data_.get(), 1, result.devicePtr()));
592 });
593 } else {
594 EIGEN_CUDA_RUNTIME_CHECK(cudaMemsetAsync(result.devicePtr(), 0, sizeof(Scalar), ctx.stream()));
595 }
596}
597
598template <typename Scalar_>
600 const DeviceMatrix& other) const {
601 // Allocated uninitialized: the reduction overwrites the slot.
602 DeviceScalar<Scalar> result(ctx.stream());
603 dot(ctx, other, result);
604 return result;
605}
606
607template <typename Scalar_>
609 const int64_t n = internal::blas1_size(rows_, cols_);
610 eigen_assert(result.stream() == ctx.stream() && "DeviceMatrix::squaredNorm: result must live on ctx's stream");
611 if (n > 0) {
612 // ||x||^2 = x^T x over the 2n real and imaginary parts of a complex x
613 // (std::complex<T> has the layout of T[2]): real, so no host sync. Unscaled,
614 // and ~4.5x faster than nrm2^2; stableNorm() is the scaled form.
615 const int64_t reals = NumTraits<Scalar>::IsComplex ? 2 * n : n;
616 const RealScalar* x = reinterpret_cast<const RealScalar*>(data_.get());
617 waitReady(ctx.stream());
618 internal::with_device_pointer_mode(ctx.cublasHandle(), [&] {
619 EIGEN_CUBLAS_CHECK(internal::cublasXdot(ctx.cublasHandle(), reals, x, 1, x, 1, result.devicePtr()));
620 });
621 } else {
622 EIGEN_CUDA_RUNTIME_CHECK(cudaMemsetAsync(result.devicePtr(), 0, sizeof(RealScalar), ctx.stream()));
623 }
624}
625
626template <typename Scalar_>
628 DeviceScalar<RealScalar> result(ctx.stream());
629 squaredNorm(ctx, result);
630 return result;
631}
632
633template <typename Scalar_>
635 // sqrt of the dot product: a dot and a one-element NPP sqrt cost less than
636 // cuBLAS nrm2's scaled accumulation, see stableNorm().
637 squaredNorm(ctx, result);
638 internal::device_scalar_sqrt(result.devicePtr(), ctx.nppStreamContext());
639}
640
641template <typename Scalar_>
643 DeviceScalar<RealScalar> result(ctx.stream());
644 norm(ctx, result);
645 return result;
646}
647
648template <typename Scalar_>
650 const int64_t n = internal::blas1_size(rows_, cols_);
651 eigen_assert(result.stream() == ctx.stream() && "DeviceMatrix::stableNorm: result must live on ctx's stream");
652 if (n > 0) {
653 waitReady(ctx.stream());
654 internal::with_device_pointer_mode(ctx.cublasHandle(), [&] {
655 EIGEN_CUBLAS_CHECK(internal::cublasXnrm2(ctx.cublasHandle(), n, data_.get(), 1, result.devicePtr()));
656 });
657 } else {
658 EIGEN_CUDA_RUNTIME_CHECK(cudaMemsetAsync(result.devicePtr(), 0, sizeof(RealScalar), ctx.stream()));
659 }
660}
661
662template <typename Scalar_>
663void DeviceMatrix<Scalar_>::setZero(cudaStream_t stream) {
664 if (sizeInBytes() > 0) {
665 waitReady(stream);
666 EIGEN_CUDA_RUNTIME_CHECK(cudaMemsetAsync(data_.get(), 0, sizeInBytes(), stream));
667 recordReady(stream);
668 }
669}
670
671template <typename Scalar_>
673 setZero(ctx.stream());
674}
675
676template <typename Scalar_>
677void DeviceMatrix<Scalar_>::addScaled(Context& ctx, Scalar alpha, const DeviceMatrix& x) {
678 const int64_t n = internal::blas1_size(rows_, cols_);
679 eigen_assert(n == internal::blas1_size(x.rows_, x.cols_));
680 if (n > 0) {
681 waitReady(ctx.stream());
682 x.waitReady(ctx.stream());
683 EIGEN_CUBLAS_CHECK(internal::cublasXaxpy(ctx.cublasHandle(), n, &alpha, x.data_.get(), 1, data_.get(), 1));
684 recordReady(ctx.stream());
685 }
686}
687
688template <typename Scalar_>
689void DeviceMatrix<Scalar_>::scale(Context& ctx, Scalar alpha) {
690 const int64_t n = internal::blas1_size(rows_, cols_);
691 if (n > 0) {
692 waitReady(ctx.stream());
693 EIGEN_CUBLAS_CHECK(internal::cublasXscal(ctx.cublasHandle(), n, &alpha, data_.get(), 1));
694 recordReady(ctx.stream());
695 }
696}
697
698template <typename Scalar_>
700 // Wait on *this before resize — resize may free the old buffer while another
701 // stream is still reading it.
702 if (!empty()) waitReady(ctx.stream());
703 resize(other.rows_, other.cols_);
704 const int64_t n = internal::blas1_size(rows_, cols_);
705 if (n > 0) {
706 other.waitReady(ctx.stream());
707 EIGEN_CUBLAS_CHECK(internal::cublasXcopy(ctx.cublasHandle(), n, other.data_.get(), 1, data_.get(), 1));
708 recordReady(ctx.stream());
709 }
710}
711
712// this += alpha * x (axpy)
713template <typename Scalar_>
714DeviceMatrix<Scalar_>& DeviceMatrix<Scalar_>::operator+=(const Scaled<DeviceMatrix>& expr) {
716 return *this;
717}
718
719// this -= alpha * x (axpy with negated alpha)
720template <typename Scalar_>
725
726// this += x (axpy with alpha=1)
727template <typename Scalar_>
728DeviceMatrix<Scalar_>& DeviceMatrix<Scalar_>::operator+=(const DeviceMatrix& other) {
729 Scalar one(1);
730 addScaled(Context::threadLocal(), one, other);
731 return *this;
732}
733
734// this -= x (axpy with alpha=-1)
735template <typename Scalar_>
737 Scalar neg_one(-1);
738 addScaled(Context::threadLocal(), neg_one, other);
739 return *this;
740}
741
742// this *= alpha (scal, host pointer)
743template <typename Scalar_>
745 scale(Context::threadLocal(), alpha);
746 return *this;
747}
748
749namespace internal {
750// x[i] /= alpha. Real Scalar: NPP divides in place. Complex Scalar: cuBLAS scal by the
751// host reciprocal, which std::complex computes with scaling, so |alpha| > sqrt(max)
752// does not overflow; one extra rounding per element.
753inline void divide_in_place(Context& ctx, float* x, int64_t n, float alpha) {
754 device_divC(alpha, x, Eigen::internal::convert_index<int>(n), ctx.stream());
755}
756inline void divide_in_place(Context& ctx, double* x, int64_t n, double alpha) {
757 device_divC(alpha, x, Eigen::internal::convert_index<int>(n), ctx.stream());
758}
759template <typename Real>
760void divide_in_place(Context& ctx, std::complex<Real>* x, int64_t n, std::complex<Real> alpha) {
761 const std::complex<Real> inv = std::complex<Real>(1) / alpha;
762 EIGEN_CUBLAS_CHECK(cublasXscal(ctx.cublasHandle(), n, &inv, x, 1));
763}
764} // namespace internal
765
766template <typename Scalar_>
767void DeviceMatrix<Scalar_>::divide(Context& ctx, Scalar alpha) {
768 const int64_t n = internal::blas1_size(rows_, cols_);
769 if (n > 0) {
770 waitReady(ctx.stream());
771 internal::divide_in_place(ctx, data_.get(), n, alpha);
772 recordReady(ctx.stream());
773 }
774}
775
776// this /= alpha
777template <typename Scalar_>
780 return *this;
781}
782
783// Deep copies: device-to-device cuBLAS copy on the thread-local Context.
784template <typename Scalar_>
788
789template <typename Scalar_>
790DeviceMatrix<Scalar_>& DeviceMatrix<Scalar_>::operator=(const DeviceMatrix& other) {
791 if (this != &other) copyFrom(Context::threadLocal(), other);
792 return *this;
793}
794
795template <typename Scalar_>
797 DeviceScalar<RealScalar> result(ctx.stream());
798 stableNorm(ctx, result);
799 return result;
800}
801
802template <typename Scalar_>
804 return stableNorm(Context::threadLocal());
805}
806
807// this *= alpha (scal, device pointer — avoids host sync)
808template <typename Scalar_>
810 const int64_t n = internal::blas1_size(rows_, cols_);
811 if (n > 0) {
812 auto& ctx = Context::threadLocal();
813 waitReady(ctx.stream());
814 internal::with_device_pointer_mode(ctx.cublasHandle(), [&] {
815 EIGEN_CUBLAS_CHECK(internal::cublasXscal(ctx.cublasHandle(), n, alpha.devicePtr(), data_.get(), 1));
816 });
817 recordReady(ctx.stream());
818 }
819 return *this;
820}
821
822// this += DeviceScalar * x (axpy with CUBLAS_POINTER_MODE_DEVICE)
823template <typename Scalar_>
824DeviceMatrix<Scalar_>& DeviceMatrix<Scalar_>::operator+=(const DeviceScaledDevice<Scalar_>& expr) {
825 const int64_t n = internal::blas1_size(rows_, cols_);
826 const auto& x = expr.matrix();
827 eigen_assert(n == internal::blas1_size(x.rows_, x.cols_));
828 if (n > 0) {
829 auto& ctx = Context::threadLocal();
830 waitReady(ctx.stream());
831 x.waitReady(ctx.stream());
832 internal::with_device_pointer_mode(ctx.cublasHandle(), [&] {
833 EIGEN_CUBLAS_CHECK(
834 internal::cublasXaxpy(ctx.cublasHandle(), n, expr.alpha().devicePtr(), x.data_.get(), 1, data_.get(), 1));
835 });
836 recordReady(ctx.stream());
837 }
838 return *this;
839}
840
841// this -= DeviceScalar * x (axpy with negated device scalar)
842template <typename Scalar_>
844 auto neg_alpha = -expr.alpha();
845 DeviceScaledDevice<Scalar_> neg_expr(neg_alpha, expr.matrix());
846 return operator+=(neg_expr);
847}
848
849// this = alpha * A + beta * B (cuBLAS geam)
850template <typename Scalar_>
851DeviceMatrix<Scalar_>& DeviceMatrix<Scalar_>::operator=(const DeviceAddExpr<Scalar_>& expr) {
852 internal::dispatch(Context::threadLocal(), *this, expr);
853 return *this;
854}
855
856// cwiseProduct (allocating).
857template <typename Scalar_>
859 const int64_t n = internal::blas1_size(rows_, cols_);
860 eigen_assert(n == internal::blas1_size(other.rows_, other.cols_));
861 DeviceMatrix result(rows_, cols_);
862 if (n > 0) {
863 waitReady(ctx.stream());
864 other.waitReady(ctx.stream());
865 internal::device_cwiseProduct(data_.get(), other.data_.get(), result.data_.get(),
866 Eigen::internal::convert_index<int>(n), ctx.stream());
867 result.recordReady(ctx.stream());
868 }
869 return result;
870}
871
872// In-place cwiseProduct: this = a .* b (reuses this buffer, no allocation).
873template <typename Scalar_>
875 const int64_t n = internal::blas1_size(a.rows_, a.cols_);
876 eigen_assert(n == internal::blas1_size(b.rows_, b.cols_));
877 if (!empty()) waitReady(ctx.stream());
878 resize(a.rows_, a.cols_);
879 if (n > 0) {
880 a.waitReady(ctx.stream());
881 b.waitReady(ctx.stream());
882 internal::device_cwiseProduct(a.data_.get(), b.data_.get(), data_.get(), Eigen::internal::convert_index<int>(n),
883 ctx.stream());
884 recordReady(ctx.stream());
885 }
886}
887
888// Convenience overloads using thread-local default Context.
889template <typename Scalar_>
891 return dot(Context::threadLocal(), other);
892}
893
894template <typename Scalar_>
895DeviceScalar<typename NumTraits<Scalar_>::Real> DeviceMatrix<Scalar_>::squaredNorm() const {
896 return squaredNorm(Context::threadLocal());
897}
898
899template <typename Scalar_>
900DeviceScalar<typename NumTraits<Scalar_>::Real> DeviceMatrix<Scalar_>::norm() const {
901 return norm(Context::threadLocal());
902}
903
904template <typename Scalar_>
906 setZero(Context::threadLocal());
907}
908} // namespace gpu
909} // namespace Eigen
910
911#endif // EIGEN_GPU_DEVICE_DISPATCH_H
Unified GPU execution context owning a CUDA stream and library handles.
Definition GpuContext.h:81
const NppStreamContext & nppStreamContext() const
Definition GpuContext.h:139
static Context & threadLocal()
Definition GpuContext.h:121
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
DeviceMatrix & operator*=(Scalar alpha)
Definition DeviceDispatch.h:744
void addScaled(Context &ctx, Scalar alpha, const DeviceMatrix &x)
Definition DeviceDispatch.h:677
Assignment< Scalar > device(Context &ctx)
Definition DeviceMatrix.h:386
DeviceMatrix cwiseProduct(Context &ctx, const DeviceMatrix &other) const
Definition DeviceDispatch.h:858
void scale(Context &ctx, Scalar alpha)
Definition DeviceDispatch.h:689
void copyFrom(Context &ctx, const DeviceMatrix &other)
Definition DeviceDispatch.h:699
DeviceScalar< typename NumTraits< Scalar >::Real > squaredNorm(Context &ctx) const
Definition DeviceDispatch.h:627
DeviceMatrix & operator/=(Scalar alpha)
Definition DeviceDispatch.h:778
DeviceScalar< typename NumTraits< Scalar >::Real > stableNorm(Context &ctx) const
Definition DeviceDispatch.h:796
DeviceMatrix & operator-=(const GemmExpr< Lhs, Rhs > &expr)
void divide(Context &ctx, Scalar alpha)
Definition DeviceDispatch.h:767
void recordReady(cudaStream_t stream)
Definition DeviceMatrix.h:363
DeviceScalar< typename NumTraits< Scalar >::Real > norm(Context &ctx) const
Definition DeviceDispatch.h:642
void setZero(Context &ctx)
Definition DeviceDispatch.h:672
DeviceScalar< Scalar > dot(Context &ctx, const DeviceMatrix &other) const
Definition DeviceDispatch.h:599
void resize(Index rows, Index cols)
Definition DeviceMatrix.h:323
void waitReady(cudaStream_t stream) const
Definition DeviceMatrix.h:372
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
Definition DeviceSolverExpr.h:30
Definition DeviceSolverExpr.h:46
Expression returned by operator*(Scalar, DeviceMatrix/View), carrying the scalar factor.
Definition DeviceExpr.h:77
void rankUpdate(const DeviceMatrix< Scalar > &A, RealScalar alpha=RealScalar(1))
Definition DeviceDispatch.h:543
Definition DeviceBlasExpr.h:96
Definition DeviceBlasExpr.h:122
Definition DeviceBlasExpr.h:45
Namespace containing all symbols from the Eigen library.
Describes GPU device expression types.
Definition DeviceExpr.h:173