17#ifndef EIGEN_GPU_SVD_H
18#define EIGEN_GPU_SVD_H
21#include "./InternalHeaderCheck.h"
23#include "./GpuSolverContext.h"
27template <
typename Scalar_>
30 using Scalar = Scalar_;
31 using RealScalar =
typename NumTraits<Scalar>::Real;
32 using PlainMatrix = Eigen::Matrix<Scalar, Dynamic, Dynamic, ColMajor>;
33 using PlainVector = Eigen::Matrix<Scalar, Dynamic, 1>;
34 using RealVector = Eigen::Matrix<RealScalar, Dynamic, 1>;
41 explicit SVD(Context& ctx) : solver_ctx_(ctx) {}
43 template <
typename InputType>
49 compute(d_A, options);
54 compute(std::move(d_A), options);
58 template <
typename InputType>
67 compute(d_A, options);
71 SVD(Context& ctx, DeviceMatrix<Scalar>&& d_A,
unsigned int options =
ComputeThinU |
ComputeThinV) : solver_ctx_(ctx) {
72 compute(std::move(d_A), options);
77 SVD(
const SVD&) =
delete;
78 SVD& operator=(
const SVD&) =
delete;
81 : solver_ctx_(std::move(o.solver_ctx_)),
82 d_A_(std::move(o.d_A_)),
83 d_U_(std::move(o.d_U_)),
84 d_S_(std::move(o.d_S_)),
85 d_VT_(std::move(o.d_VT_)),
86 d_D_(std::move(o.d_D_)),
87 cached_diag_kk_(o.cached_diag_kk_),
88 cached_diag_lambda_(o.cached_diag_lambda_),
89 diag_valid_(o.diag_valid_),
94 transposed_(o.transposed_) {
95 o.diag_valid_ =
false;
100 o.transposed_ =
false;
103 SVD& operator=(SVD&& o)
noexcept {
105 solver_ctx_ = std::move(o.solver_ctx_);
106 d_A_ = std::move(o.d_A_);
107 d_U_ = std::move(o.d_U_);
108 d_S_ = std::move(o.d_S_);
109 d_VT_ = std::move(o.d_VT_);
110 d_D_ = std::move(o.d_D_);
111 cached_diag_kk_ = o.cached_diag_kk_;
112 cached_diag_lambda_ = o.cached_diag_lambda_;
113 diag_valid_ = o.diag_valid_;
114 options_ = o.options_;
118 transposed_ = o.transposed_;
119 o.diag_valid_ =
false;
124 o.transposed_ =
false;
129 template <
typename InputType>
139 if (!begin_compute(d_A, options))
return *
this;
142 transpose_into_input(d_A);
144 const size_t mat_bytes =
static_cast<size_t>(lda_) *
static_cast<size_t>(n_) *
sizeof(Scalar);
145 d_A_ = internal::DeviceBuffer(mat_bytes);
146 EIGEN_CUDA_RUNTIME_CHECK(
147 cudaMemcpyAsync(d_A_.get(), d_A.data(), mat_bytes, cudaMemcpyDeviceToDevice, solver_ctx_.stream()));
158 if (d_A.isView())
return compute(
static_cast<const DeviceMatrix<Scalar>&
>(d_A), options);
159 if (!begin_compute(d_A, options))
return *
this;
162 transpose_into_input(d_A);
164 const size_t a_bytes = d_A.sizeInBytes();
165 d_A_ = internal::DeviceBuffer::adopt(
static_cast<void*
>(d_A.release()), a_bytes);
174 Index rows()
const {
return transposed_ ? n_ : m_; }
175 Index cols()
const {
return transposed_ ? m_ : n_; }
178 RealVector singularValues()
const {
179 eigen_assert(solver_ctx_.info() ==
Success);
180 const Index k = (std::min)(m_, n_);
182 solver_ctx_.download(S.data(), d_S_.get(),
static_cast<size_t>(k) *
sizeof(RealScalar));
187 PlainMatrix matrixU()
const {
188 eigen_assert(solver_ctx_.info() ==
Success);
190 const Index m_orig = transposed_ ? n_ : m_;
191 const Index n_orig = transposed_ ? m_ : n_;
192 const Index k = (std::min)(m_orig, n_orig);
195 PlainMatrix U(m_, ucols);
196 solver_ctx_.download(U.data(), d_U_.get(),
static_cast<size_t>(m_) *
static_cast<size_t>(ucols) *
sizeof(Scalar));
199 const Index vtrows = (options_ &
ComputeFullU) ? m_orig : k;
200 PlainMatrix VT_stored(vtrows, n_);
201 solver_ctx_.download(VT_stored.data(), d_VT_.get(),
202 static_cast<size_t>(vtrows) *
static_cast<size_t>(n_) *
sizeof(Scalar));
203 return VT_stored.adjoint();
208 PlainMatrix matrixV()
const {
return matrixVT().adjoint(); }
211 PlainMatrix matrixVT()
const {
212 eigen_assert(solver_ctx_.info() ==
Success);
214 const Index m_orig = transposed_ ? n_ : m_;
215 const Index n_orig = transposed_ ? m_ : n_;
216 const Index k = (std::min)(m_orig, n_orig);
218 const Index vtrows = (options_ &
ComputeFullV) ? n_ : k;
219 PlainMatrix VT(vtrows, n_);
220 solver_ctx_.download(VT.data(), d_VT_.get(),
221 static_cast<size_t>(vtrows) *
static_cast<size_t>(n_) *
sizeof(Scalar));
224 const Index ucols = (options_ &
ComputeFullV) ? n_orig : k;
225 PlainMatrix U_stored(m_, ucols);
226 solver_ctx_.download(U_stored.data(), d_U_.get(),
227 static_cast<size_t>(m_) *
static_cast<size_t>(ucols) *
sizeof(Scalar));
228 return U_stored.adjoint();
243 DeviceMatrix<RealScalar> d_singularValues()
const {
244 eigen_assert(solver_ctx_.info() ==
Success);
245 const Index k = (std::min)(m_, n_);
247 v.recordReady(solver_ctx_.stream());
253 DeviceMatrix<Scalar> d_matrixU()
const {
254 eigen_assert(solver_ctx_.info() ==
Success);
256 const Index m_orig = transposed_ ? n_ : m_;
257 const Index n_orig = transposed_ ? m_ : n_;
258 const Index k = (std::min)(m_orig, n_orig);
262 v.recordReady(solver_ctx_.stream());
266 const Index vtrows_stored = (options_ &
ComputeFullU) ? n_ : k;
267 DeviceMatrix<Scalar> result(n_, vtrows_stored);
268 if (n_ > 0 && vtrows_stored > 0) {
269 Scalar alpha_one(1), beta_zero(0);
270 EIGEN_CUBLAS_CHECK(internal::cublasXgeam(solver_ctx_.cublasHandle(), CUBLAS_OP_C, CUBLAS_OP_N, n_, vtrows_stored,
271 &alpha_one,
static_cast<const Scalar*
>(d_VT_.get()), vtrows_stored,
272 &beta_zero,
static_cast<const Scalar*
>(
nullptr), n_, result.data(), n_));
273 result.recordReady(solver_ctx_.stream());
280 DeviceMatrix<Scalar> d_matrixVT()
const {
281 eigen_assert(solver_ctx_.info() ==
Success);
283 const Index m_orig = transposed_ ? n_ : m_;
284 const Index n_orig = transposed_ ? m_ : n_;
285 const Index k = (std::min)(m_orig, n_orig);
287 const Index vtrows = (options_ &
ComputeFullV) ? n_ : k;
289 v.recordReady(solver_ctx_.stream());
293 const Index ucols = (options_ &
ComputeFullV) ? n_orig : k;
294 DeviceMatrix<Scalar> result(ucols, m_);
295 if (ucols > 0 && m_ > 0) {
296 Scalar alpha_one(1), beta_zero(0);
297 EIGEN_CUBLAS_CHECK(internal::cublasXgeam(solver_ctx_.cublasHandle(), CUBLAS_OP_C, CUBLAS_OP_N, ucols, m_,
298 &alpha_one,
static_cast<const Scalar*
>(d_U_.get()), m_, &beta_zero,
299 static_cast<const Scalar*
>(
nullptr), ucols, result.data(), ucols));
300 result.recordReady(solver_ctx_.stream());
306 Index rank(RealScalar threshold = RealScalar(-1))
const {
307 RealVector S = singularValues();
308 if (S.size() == 0)
return 0;
310 threshold = (std::max)(m_, n_) * S(0) * NumTraits<RealScalar>::epsilon();
312 return (S.array() > threshold).count();
316 template <
typename Rhs>
317 PlainMatrix solve(
const MatrixBase<Rhs>& B)
const {
318 return solve_impl(B, (std::min)(m_, n_), RealScalar(0));
322 template <
typename Rhs>
323 PlainMatrix solve(
const MatrixBase<Rhs>& B, Index trunc)
const {
324 eigen_assert(trunc > 0 && trunc <= (std::min)(m_, n_));
325 return solve_impl(B, trunc, RealScalar(0));
329 template <
typename Rhs>
330 PlainMatrix solve(
const MatrixBase<Rhs>& B, RealScalar lambda)
const {
331 eigen_assert(lambda > 0);
332 return solve_impl(B, (std::min)(m_, n_), lambda);
340 DeviceMatrix<Scalar> solve(
const DeviceMatrix<Scalar>& d_B)
const {
341 return solve_device_impl(d_B, (std::min)(m_, n_), RealScalar(0));
345 DeviceMatrix<Scalar> solve(
const DeviceMatrix<Scalar>& d_B, Index trunc)
const {
346 eigen_assert(trunc > 0 && trunc <= (std::min)(m_, n_));
347 return solve_device_impl(d_B, trunc, RealScalar(0));
351 DeviceMatrix<Scalar> solve(
const DeviceMatrix<Scalar>& d_B, RealScalar lambda)
const {
352 eigen_assert(lambda > 0);
353 return solve_device_impl(d_B, (std::min)(m_, n_), lambda);
356 cudaStream_t stream()
const {
return solver_ctx_.stream(); }
359 mutable internal::GpuSolverContext solver_ctx_;
360 internal::DeviceBuffer d_A_;
361 internal::DeviceBuffer d_U_;
362 internal::DeviceBuffer d_S_;
363 internal::DeviceBuffer d_VT_;
365 mutable internal::DeviceBuffer d_D_;
366 mutable Index cached_diag_kk_ = -1;
367 mutable RealScalar cached_diag_lambda_ = RealScalar(-1);
368 mutable bool diag_valid_ =
false;
369 unsigned int options_ = 0;
373 bool transposed_ =
false;
377 bool begin_compute(
const DeviceMatrix<Scalar>& d_A,
unsigned int options) {
384 if (!solver_ctx_.begin_compute(m_ != 0 && n_ != 0)) {
385 d_A_ = internal::DeviceBuffer();
386 d_U_ = internal::DeviceBuffer();
387 d_S_ = internal::DeviceBuffer();
388 d_VT_ = internal::DeviceBuffer();
391 transposed_ = (m_ < n_);
396 lda_ =
static_cast<int64_t
>(d_A.rows());
398 d_A.waitReady(solver_ctx_.stream());
403 void transpose_into_input(
const DeviceMatrix<Scalar>& d_A) {
404 const size_t mat_bytes =
static_cast<size_t>(lda_) *
static_cast<size_t>(n_) *
sizeof(Scalar);
405 d_A_ = internal::DeviceBuffer(mat_bytes);
407 Scalar alpha_one(1), beta_zero(0);
408 EIGEN_CUBLAS_CHECK(internal::cublasXgeam(solver_ctx_.cublasHandle(), CUBLAS_OP_C, CUBLAS_OP_N, m_, n_, &alpha_one,
409 d_A.data(), d_A.rows(), &beta_zero,
static_cast<const Scalar*
>(
nullptr),
410 m_,
static_cast<Scalar*
>(d_A_.get()), m_));
414 static unsigned int swap_uv_options(
unsigned int opts) {
415 unsigned int result = 0;
423 static signed char jobu(
unsigned int opts) {
429 static signed char jobvt(
unsigned int opts) {
436 constexpr cudaDataType_t dtype = internal::cusolver_data_type<Scalar>::value;
437 constexpr cudaDataType_t rtype = internal::cuda_data_type<RealScalar>::value;
438 const Index k = (std::min)(m_, n_);
440 solver_ctx_.mark_pending();
442 internal::ensure_sized(d_S_,
static_cast<size_t>(k) *
sizeof(RealScalar));
444 const unsigned int int_opts = transposed_ ? swap_uv_options(options_) : options_;
448 const int64_t ldu = m_;
449 const int64_t ldvt = vtrows > 0 ? vtrows : 1;
452 internal::ensure_sized(d_U_,
static_cast<size_t>(m_) *
static_cast<size_t>(ucols) *
sizeof(Scalar));
455 internal::ensure_sized(d_VT_,
static_cast<size_t>(vtrows) *
static_cast<size_t>(n_) *
sizeof(Scalar));
458 eigen_assert(m_ >= n_ &&
"Internal error: m_ < n_ should have been handled by transpose in compute()");
459 size_t dev_ws = 0, host_ws = 0;
460 EIGEN_CUSOLVER_CHECK(cusolverDnXgesvd_bufferSize(
461 solver_ctx_.cusolverHandle(), solver_ctx_.params_.p, jobu(int_opts), jobvt(int_opts), m_, n_, dtype, d_A_.get(),
462 lda_, rtype, d_S_.get(), dtype, ucols > 0 ? d_U_.get() :
nullptr, ldu, dtype,
463 vtrows > 0 ? d_VT_.get() :
nullptr, ldvt, dtype, &dev_ws, &host_ws));
465 solver_ctx_.ensure_scratch(dev_ws);
466 solver_ctx_.h_workspace_.resize(host_ws);
468 EIGEN_CUSOLVER_CHECK(cusolverDnXgesvd(
469 solver_ctx_.cusolverHandle(), solver_ctx_.params_.p, jobu(int_opts), jobvt(int_opts), m_, n_, dtype, d_A_.get(),
470 lda_, rtype, d_S_.get(), dtype, ucols > 0 ? d_U_.get() :
nullptr, ldu, dtype,
471 vtrows > 0 ? d_VT_.get() :
nullptr, ldvt, dtype, solver_ctx_.scratch_workspace(), dev_ws,
472 host_ws > 0 ? solver_ctx_.h_workspace_.data() :
nullptr, host_ws, solver_ctx_.scratch_info()));
474 solver_ctx_.enqueue_info_copy();
480 d_A_ = internal::DeviceBuffer();
493 void build_diag(Index kk, RealScalar lambda)
const {
494 if (diag_valid_ && cached_diag_kk_ == kk && cached_diag_lambda_ == lambda)
return;
496 const Index k = (std::min)(m_, n_);
498 EIGEN_CUDA_RUNTIME_CHECK(cudaMemcpyAsync(S.data(), d_S_.get(),
static_cast<size_t>(k) *
sizeof(RealScalar),
499 cudaMemcpyDeviceToHost, solver_ctx_.stream()));
500 EIGEN_CUDA_RUNTIME_CHECK(cudaStreamSynchronize(solver_ctx_.stream()));
502 const RealScalar drop_threshold = S(0) * RealScalar(k) * NumTraits<RealScalar>::epsilon();
503 auto S_head = S.head(kk).array();
505 if (lambda == RealScalar(0)) {
506 D = (S_head > drop_threshold).select(S_head.inverse(), RealScalar(0)).matrix().template cast<Scalar>();
508 D = (S_head / (S_head.square() + lambda * lambda)).matrix().template cast<Scalar>();
511 const size_t d_bytes =
static_cast<size_t>(kk) *
sizeof(Scalar);
512 internal::ensure_sized(d_D_, d_bytes);
513 EIGEN_CUDA_RUNTIME_CHECK(
514 cudaMemcpyAsync(d_D_.get(), D.data(), d_bytes, cudaMemcpyHostToDevice, solver_ctx_.stream()));
515 cached_diag_kk_ = kk;
516 cached_diag_lambda_ = lambda;
523 void apply_pinv(
const Scalar* B_dev, Index kk, Index nrhs, Scalar* X_dev)
const {
524 const Index m_orig = transposed_ ? n_ : m_;
525 const Index n_orig = transposed_ ? m_ : n_;
526 const Index k = (std::min)(m_, n_);
528 auto* U_dev =
static_cast<const Scalar*
>(d_U_.get());
529 auto* VT_dev =
static_cast<const Scalar*
>(d_VT_.get());
531 Scalar scalars[2] = {Scalar(1), Scalar(0)};
534 internal::DeviceBuffer d_tmp(
static_cast<size_t>(kk) *
static_cast<size_t>(nrhs) *
sizeof(Scalar));
535 auto* tmp_dev =
static_cast<Scalar*
>(d_tmp.get());
537 internal::cublaslt_gemm(solver_ctx_.cublasLtHandle(), solver_ctx_.cublasHandle(), CUBLAS_OP_C, CUBLAS_OP_N, kk,
538 nrhs, m_, &scalars[0], U_dev, m_, B_dev, m_orig, &scalars[1], tmp_dev, kk,
539 solver_ctx_.gemmWorkspace(), solver_ctx_.gemmPlanCache(),
540 solver_ctx_.cublasLtMaxWorkspaceBytes(), solver_ctx_.stream());
542 const Index vtrows_stored = (swap_uv_options(options_) &
ComputeFullV) ? n_ : k;
543 internal::cublaslt_gemm(solver_ctx_.cublasLtHandle(), solver_ctx_.cublasHandle(), CUBLAS_OP_N, CUBLAS_OP_N, kk,
544 nrhs, m_orig, &scalars[0], VT_dev, vtrows_stored, B_dev, m_orig, &scalars[1], tmp_dev, kk,
545 solver_ctx_.gemmWorkspace(), solver_ctx_.gemmPlanCache(),
546 solver_ctx_.cublasLtMaxWorkspaceBytes(), solver_ctx_.stream());
550 EIGEN_CUBLAS_CHECK(internal::cublasXdgmm(solver_ctx_.cublasHandle(), CUBLAS_SIDE_LEFT, kk, nrhs, tmp_dev, kk,
551 static_cast<const Scalar*
>(d_D_.get()), 1, tmp_dev, kk));
555 const Index vtrows = (options_ &
ComputeFullV) ? n_ : k;
556 internal::cublaslt_gemm(solver_ctx_.cublasLtHandle(), solver_ctx_.cublasHandle(), CUBLAS_OP_C, CUBLAS_OP_N,
557 n_orig, nrhs, kk, &scalars[0], VT_dev, vtrows, tmp_dev, kk, &scalars[1], X_dev, n_orig,
558 solver_ctx_.gemmWorkspace(), solver_ctx_.gemmPlanCache(),
559 solver_ctx_.cublasLtMaxWorkspaceBytes(), solver_ctx_.stream());
561 internal::cublaslt_gemm(solver_ctx_.cublasLtHandle(), solver_ctx_.cublasHandle(), CUBLAS_OP_N, CUBLAS_OP_N,
562 n_orig, nrhs, kk, &scalars[0], U_dev, m_, tmp_dev, kk, &scalars[1], X_dev, n_orig,
563 solver_ctx_.gemmWorkspace(), solver_ctx_.gemmPlanCache(),
564 solver_ctx_.cublasLtMaxWorkspaceBytes(), solver_ctx_.stream());
568 template <
typename Rhs>
569 PlainMatrix solve_impl(
const MatrixBase<Rhs>& B, Index trunc, RealScalar lambda)
const {
570 eigen_assert(solver_ctx_.info() ==
Success &&
"SVD::solve called on a failed or uninitialized decomposition");
574 const Index m_orig = transposed_ ? n_ : m_;
575 const Index n_orig = transposed_ ? m_ : n_;
576 eigen_assert(B.rows() == m_orig);
578 const Index k = (std::min)(m_, n_);
579 const Index kk = (std::min)(trunc, k);
580 const Index nrhs = B.cols();
583 if (kk == 0 || nrhs == 0 || n_orig == 0) {
584 return PlainMatrix::Zero(n_orig, nrhs);
590 const Ref<const PlainMatrix> rhs(B.derived());
591 internal::DeviceBuffer d_B(
static_cast<size_t>(m_orig) *
static_cast<size_t>(nrhs) *
sizeof(Scalar));
592 internal::upload_host_matrix(
static_cast<Scalar*
>(d_B.get()), m_orig, rhs.data(), rhs.outerStride(), m_orig, nrhs,
593 solver_ctx_.stream());
594 build_diag(kk, lambda);
596 PlainMatrix X(n_orig, nrhs);
597 internal::DeviceBuffer d_X(
static_cast<size_t>(n_orig) *
static_cast<size_t>(nrhs) *
sizeof(Scalar));
598 apply_pinv(
static_cast<const Scalar*
>(d_B.get()), kk, nrhs,
static_cast<Scalar*
>(d_X.get()));
600 EIGEN_CUDA_RUNTIME_CHECK(cudaMemcpyAsync(X.data(), d_X.get(),
601 static_cast<size_t>(n_orig) *
static_cast<size_t>(nrhs) *
sizeof(Scalar),
602 cudaMemcpyDeviceToHost, solver_ctx_.stream()));
603 EIGEN_CUDA_RUNTIME_CHECK(cudaStreamSynchronize(solver_ctx_.stream()));
608 DeviceMatrix<Scalar> solve_device_impl(
const DeviceMatrix<Scalar>& d_B, Index trunc, RealScalar lambda)
const {
609 eigen_assert(solver_ctx_.info() ==
Success &&
"SVD::solve called on a failed or uninitialized decomposition");
613 const Index m_orig = transposed_ ? n_ : m_;
614 const Index n_orig = transposed_ ? m_ : n_;
615 eigen_assert(d_B.rows() == m_orig);
617 const Index k = (std::min)(m_, n_);
618 const Index kk = (std::min)(trunc, k);
619 const Index nrhs = d_B.cols();
621 if (kk == 0 || nrhs == 0 || n_orig == 0) {
622 DeviceMatrix<Scalar> X(n_orig, nrhs);
623 X.setZero(solver_ctx_.stream());
627 d_B.waitReady(solver_ctx_.stream());
628 build_diag(kk, lambda);
630 DeviceMatrix<Scalar> X(n_orig, nrhs);
631 apply_pinv(d_B.data(), kk, nrhs, X.data());
632 X.recordReady(solver_ctx_.stream());
static DeviceMatrix fromHost(const DenseBase< Derived > &host, cudaStream_t stream=nullptr)
Definition DeviceMatrix.h:226
static DeviceMatrix view(Scalar *device_ptr, Index rows, Index cols)
Definition DeviceMatrix.h:576
Namespace containing all symbols from the Eigen library.