15#ifndef EIGEN_GPU_DEVICE_DISPATCH_H
16#define EIGEN_GPU_DEVICE_DISPATCH_H
19#include "./InternalHeaderCheck.h"
23#include "./DeviceExpr.h"
24#include "./DeviceBlasExpr.h"
25#include "./DeviceSolverExpr.h"
26#include "./GpuContext.h"
27#include "./CuSolverSupport.h"
32template <
typename Scalar>
33bool aliases_device_memory(
const DeviceMatrix<Scalar>& a,
const DeviceMatrix<Scalar>& b) {
34 return a.data() !=
nullptr && a.data() == b.data();
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>;
44 const DeviceMatrix<Scalar>& A = traits_lhs::matrix(expr.lhs());
45 const DeviceMatrix<Scalar>& B = traits_rhs::matrix(expr.rhs());
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)");
51 constexpr cublasOperation_t transA = to_cublas_op(traits_lhs::op);
52 constexpr cublasOperation_t transB = to_cublas_op(traits_rhs::op);
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();
59 eigen_assert(k == rhs_k &&
"DeviceMatrix GEMM dimension mismatch");
61 const int64_t lda = A.rows();
62 const int64_t ldb = B.rows();
65 dst.waitReady(ctx.stream());
68 const bool resized = dst.empty() || dst.rows() != m || dst.cols() != n;
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());
83 const int64_t ldc = dst.rows();
85 Scalar alpha_local = alpha_scale * traits_lhs::alpha(expr.lhs()) * traits_rhs::alpha(expr.rhs());
87 A.waitReady(ctx.stream());
88 B.waitReady(ctx.stream());
90 if (resized && beta_val != Scalar(0) && dst.sizeInBytes() > 0) {
91 EIGEN_CUDA_RUNTIME_CHECK(cudaMemsetAsync(dst.data(), 0, dst.sizeInBytes(), ctx.stream()));
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());
98 dst.recordReady(ctx.stream());
106inline void oneshot_check_info(Context& ctx, OneShotSolverScratch& scratch,
const char* what) {
108 EIGEN_UNUSED_VARIABLE(ctx);
109 EIGEN_UNUSED_VARIABLE(scratch);
110 EIGEN_UNUSED_VARIABLE(what);
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);
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();
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");
130 if (A.rows() == 0 || B.cols() == 0) {
131 if (!dst.empty()) dst.waitReady(ctx.stream());
132 dst.resize(A.rows(), B.cols());
136 A.waitReady(ctx.stream());
137 B.waitReady(ctx.stream());
138 if (!dst.empty()) dst.waitReady(ctx.stream());
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();
148 const size_t mat_bytes = A.sizeInBytes();
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()));
154 const int64_t lda =
static_cast<int64_t
>(A.rows());
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);
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));
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()));
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());
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();
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");
188 if (A.rows() == 0 || B.cols() == 0) {
189 if (!dst.empty()) dst.waitReady(ctx.stream());
190 dst.resize(A.rows(), B.cols());
194 A.waitReady(ctx.stream());
195 B.waitReady(ctx.stream());
196 if (!dst.empty()) dst.waitReady(ctx.stream());
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();
205 const size_t mat_bytes = A.sizeInBytes();
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()));
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());
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,
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()));
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());
238template <
typename Scalar,
int UpLo>
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");
246 const int64_t n = A.rows();
247 const int64_t nrhs = B.cols();
249 if (n == 0 || nrhs == 0) {
250 if (!dst.empty()) dst.
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());
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()));
265 constexpr cublasFillMode_t uplo = (UpLo ==
Lower) ? CUBLAS_FILL_MODE_LOWER : CUBLAS_FILL_MODE_UPPER;
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()));
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();
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");
282 const int64_t m = A.rows();
283 const int64_t n = B.cols();
285 if (m == 0 || n == 0) {
286 if (!dst.empty()) dst.waitReady(ctx.stream());
287 dst.resize(m, B.cols());
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());
299 constexpr cublasFillMode_t uplo = (UpLo ==
Lower) ? CUBLAS_FILL_MODE_LOWER : CUBLAS_FILL_MODE_UPPER;
300 const Scalar one(1), zero(0);
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()));
305 dst.recordReady(ctx.stream());
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();
314 const int64_t n = A.rows();
315 const int64_t k = A.cols();
318 if (!dst.empty()) dst.waitReady(ctx.stream());
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());
327 if (dst.empty() || dst.rows() != n || dst.cols() != n) {
329 if (beta_val != RealScalar(0)) {
330 EIGEN_CUDA_RUNTIME_CHECK(cudaMemsetAsync(dst.data(), 0, dst.sizeInBytes(), ctx.stream()));
334 constexpr cublasFillMode_t uplo = (UpLo ==
Lower) ? CUBLAS_FILL_MODE_LOWER : CUBLAS_FILL_MODE_UPPER;
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()));
339 dst.recordReady(ctx.stream());
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();
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());
368template <
typename Scalar_>
371 using Scalar = Scalar_;
373 Assignment(DeviceMatrix<Scalar>& dst, Context& ctx) : dst_(dst), ctx_(ctx) {}
375 template <
typename Lhs,
typename Rhs>
376 DeviceMatrix<Scalar>& operator=(
const GemmExpr<Lhs, Rhs>& expr) {
377 internal::dispatch(ctx_, dst_, expr, Scalar(0));
381 template <
typename Lhs,
typename Rhs>
382 DeviceMatrix<Scalar>& operator+=(
const GemmExpr<Lhs, Rhs>& expr) {
383 internal::dispatch(ctx_, dst_, expr, Scalar(1));
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));
394 DeviceMatrix<Scalar>& operator=(
const LltSolveExpr<Scalar, UpLo>& expr) {
395 internal::dispatch(ctx_, dst_, expr);
399 DeviceMatrix<Scalar>& operator=(
const LuSolveExpr<Scalar>& expr) {
400 internal::dispatch(ctx_, dst_, expr);
405 DeviceMatrix<Scalar>& operator=(
const TrsmExpr<Scalar, UpLo>& expr) {
406 internal::dispatch(ctx_, dst_, expr);
411 DeviceMatrix<Scalar>& operator=(
const SymmExpr<Scalar, UpLo>& expr) {
412 internal::dispatch(ctx_, dst_, expr);
416 DeviceMatrix<Scalar>& operator=(
const DeviceAddExpr<Scalar>& expr) {
417 internal::dispatch(ctx_, dst_, expr);
421 DeviceMatrix<Scalar>& operator=(
const Scaled<DeviceMatrix<Scalar>>& expr) {
423 internal::dispatch(ctx_, dst_, DeviceAddExpr<Scalar>(expr.scalar(), expr.inner(), Scalar(0), expr.inner()));
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()).");
438 DeviceMatrix<Scalar>& dst_;
445template <
typename Scalar_>
446template <
typename Lhs,
typename Rhs>
452template <
typename Scalar_>
453template <
typename Lhs,
typename Rhs>
459template <
typename Scalar_>
460template <
typename Lhs,
typename Rhs>
466template <
typename Scalar_>
473template <
typename Scalar_>
479template <
typename Scalar_>
486template <
typename Scalar_>
493template <
typename Scalar_>
503template <
typename Scalar_>
504template <
typename Lhs,
typename Rhs>
509template <
typename Scalar_>
514template <
typename Scalar_>
519template <
typename Scalar_>
525template <
typename Scalar_>
530template <
typename Scalar_>
536template <
typename Scalar_>
542template <
typename Scalar_,
int UpLo_>
545 RealScalar beta = matrix().empty() ? RealScalar(0) : RealScalar(1);
554void with_device_pointer_mode(cublasHandle_t h, F&& f) {
555 struct RestoreOnThrow {
556 cublasHandle_t handle;
557 cublasPointerMode_t mode;
560 if (armed) (void)cublasSetPointerMode(handle, mode);
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};
568 restore.armed =
false;
569 EIGEN_CUBLAS_CHECK(cublasSetPointerMode(h, prev));
578inline int64_t blas1_size(Index rows, Index cols) {
return static_cast<int64_t
>(rows) *
static_cast<int64_t
>(cols); }
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");
589 internal::with_device_pointer_mode(ctx.cublasHandle(), [&] {
591 internal::cublasXdot(ctx.cublasHandle(), n, data_.get(), 1, other.data_.get(), 1, result.devicePtr()));
594 EIGEN_CUDA_RUNTIME_CHECK(cudaMemsetAsync(result.devicePtr(), 0,
sizeof(Scalar), ctx.stream()));
598template <
typename Scalar_>
603 dot(ctx, other, result);
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");
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()));
622 EIGEN_CUDA_RUNTIME_CHECK(cudaMemsetAsync(result.devicePtr(), 0,
sizeof(RealScalar), ctx.stream()));
626template <
typename Scalar_>
633template <
typename Scalar_>
637 squaredNorm(ctx, result);
641template <
typename Scalar_>
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");
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()));
658 EIGEN_CUDA_RUNTIME_CHECK(cudaMemsetAsync(result.devicePtr(), 0,
sizeof(RealScalar), ctx.stream()));
662template <
typename Scalar_>
664 if (sizeInBytes() > 0) {
666 EIGEN_CUDA_RUNTIME_CHECK(cudaMemsetAsync(data_.get(), 0, sizeInBytes(), stream));
671template <
typename Scalar_>
676template <
typename Scalar_>
678 const int64_t n = internal::blas1_size(rows_, cols_);
679 eigen_assert(n == internal::blas1_size(x.rows_, x.cols_));
683 EIGEN_CUBLAS_CHECK(internal::cublasXaxpy(ctx.cublasHandle(), n, &alpha, x.data_.get(), 1, data_.get(), 1));
688template <
typename Scalar_>
690 const int64_t n = internal::blas1_size(rows_, cols_);
693 EIGEN_CUBLAS_CHECK(internal::cublasXscal(ctx.cublasHandle(), n, &alpha, data_.get(), 1));
698template <
typename Scalar_>
703 resize(other.rows_, other.cols_);
704 const int64_t n = internal::blas1_size(rows_, cols_);
707 EIGEN_CUBLAS_CHECK(internal::cublasXcopy(ctx.cublasHandle(), n, other.data_.get(), 1, data_.get(), 1));
713template <
typename Scalar_>
720template <
typename Scalar_>
727template <
typename Scalar_>
735template <
typename Scalar_>
743template <
typename Scalar_>
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());
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());
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));
766template <
typename Scalar_>
768 const int64_t n = internal::blas1_size(rows_, cols_);
771 internal::divide_in_place(ctx, data_.get(), n, alpha);
777template <
typename Scalar_>
784template <
typename Scalar_>
789template <
typename Scalar_>
795template <
typename Scalar_>
802template <
typename Scalar_>
808template <
typename Scalar_>
810 const int64_t n = internal::blas1_size(rows_, cols_);
814 internal::with_device_pointer_mode(ctx.cublasHandle(), [&] {
815 EIGEN_CUBLAS_CHECK(internal::cublasXscal(ctx.cublasHandle(), n, alpha.devicePtr(), data_.get(), 1));
823template <
typename Scalar_>
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_));
831 x.waitReady(ctx.stream());
832 internal::with_device_pointer_mode(ctx.cublasHandle(), [&] {
834 internal::cublasXaxpy(ctx.cublasHandle(), n, expr.alpha().devicePtr(), x.data_.get(), 1, data_.get(), 1));
842template <
typename Scalar_>
844 auto neg_alpha = -expr.alpha();
846 return operator+=(neg_expr);
850template <
typename Scalar_>
857template <
typename Scalar_>
859 const int64_t n = internal::blas1_size(rows_, cols_);
860 eigen_assert(n == internal::blas1_size(other.rows_, other.cols_));
865 internal::device_cwiseProduct(data_.get(), other.data_.get(), result.data_.get(),
866 Eigen::internal::convert_index<int>(n), ctx.stream());
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_));
882 internal::device_cwiseProduct(a.data_.get(), b.data_.get(), data_.get(), Eigen::internal::convert_index<int>(n),
889template <
typename Scalar_>
894template <
typename Scalar_>
895DeviceScalar<typename NumTraits<Scalar_>::Real> DeviceMatrix<Scalar_>::squaredNorm()
const {
899template <
typename Scalar_>
900DeviceScalar<typename NumTraits<Scalar_>::Real> DeviceMatrix<Scalar_>::norm()
const {
904template <
typename Scalar_>
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