Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
GpuSVD.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// GPU SVD using cuSOLVER's divide-and-conquer cusolverDnXgesvd. U, S, and VT
12// stay on device; solve() forms X = V * diag(D) * U^H * B with cuBLAS GEMM.
13//
14// cuSOLVER returns VT rather than V, so VT is what is stored; matrixV() adjoints
15// it to match JacobiSVD and BDCSVD.
16
17#ifndef EIGEN_GPU_SVD_H
18#define EIGEN_GPU_SVD_H
19
20// IWYU pragma: private
21#include "./InternalHeaderCheck.h"
22
23#include "./GpuSolverContext.h"
24
25namespace Eigen {
26namespace gpu {
27template <typename Scalar_>
28class SVD {
29 public:
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>;
35
36 SVD() = default;
37
41 explicit SVD(Context& ctx) : solver_ctx_(ctx) {}
42
43 template <typename InputType>
44 explicit SVD(const DenseBase<InputType>& A, unsigned int options = ComputeThinU | ComputeThinV) {
45 compute(A, options);
46 }
47
48 explicit SVD(const DeviceMatrix<Scalar>& d_A, unsigned int options = ComputeThinU | ComputeThinV) {
49 compute(d_A, options);
50 }
51
53 explicit SVD(DeviceMatrix<Scalar>&& d_A, unsigned int options = ComputeThinU | ComputeThinV) {
54 compute(std::move(d_A), options);
55 }
56
58 template <typename InputType>
59 SVD(Context& ctx, const DenseBase<InputType>& A, unsigned int options = ComputeThinU | ComputeThinV)
60 : solver_ctx_(ctx) {
61 compute(A, options);
62 }
63
65 SVD(Context& ctx, const DeviceMatrix<Scalar>& d_A, unsigned int options = ComputeThinU | ComputeThinV)
66 : solver_ctx_(ctx) {
67 compute(d_A, options);
68 }
69
71 SVD(Context& ctx, DeviceMatrix<Scalar>&& d_A, unsigned int options = ComputeThinU | ComputeThinV) : solver_ctx_(ctx) {
72 compute(std::move(d_A), options);
73 }
74
75 ~SVD() = default;
76
77 SVD(const SVD&) = delete;
78 SVD& operator=(const SVD&) = delete;
79
80 SVD(SVD&& o) noexcept
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_),
90 options_(o.options_),
91 m_(o.m_),
92 n_(o.n_),
93 lda_(o.lda_),
94 transposed_(o.transposed_) {
95 o.diag_valid_ = false;
96 o.options_ = 0;
97 o.m_ = 0;
98 o.n_ = 0;
99 o.lda_ = 0;
100 o.transposed_ = false;
101 }
102
103 SVD& operator=(SVD&& o) noexcept {
104 if (this != &o) {
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_;
115 m_ = o.m_;
116 n_ = o.n_;
117 lda_ = o.lda_;
118 transposed_ = o.transposed_;
119 o.diag_valid_ = false;
120 o.options_ = 0;
121 o.m_ = 0;
122 o.n_ = 0;
123 o.lda_ = 0;
124 o.transposed_ = false;
125 }
126 return *this;
127 }
128
129 template <typename InputType>
130 SVD& compute(const DenseBase<InputType>& A, unsigned int options = ComputeThinU | ComputeThinV) {
131 // Upload to device, then delegate to the adopting overload — the freshly
132 // uploaded matrix is consumed in place by gesvd, so no second device copy.
133 // The wide-matrix transpose runs on the GPU (via cublasXgeam) inside the
134 // device-input path; no host transpose.
135 return compute(DeviceMatrix<Scalar>::fromHost(A.derived(), solver_ctx_.stream()), options);
136 }
137
138 SVD& compute(const DeviceMatrix<Scalar>& d_A, unsigned int options = ComputeThinU | ComputeThinV) {
139 if (!begin_compute(d_A, options)) return *this;
140
141 if (transposed_) {
142 transpose_into_input(d_A);
143 } else {
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()));
148 }
149
150 factorize();
151 return *this;
152 }
153
157 SVD& compute(DeviceMatrix<Scalar>&& d_A, unsigned int options = ComputeThinU | ComputeThinV) {
158 if (d_A.isView()) return compute(static_cast<const DeviceMatrix<Scalar>&>(d_A), options);
159 if (!begin_compute(d_A, options)) return *this;
160
161 if (transposed_) {
162 transpose_into_input(d_A);
163 } else {
164 const size_t a_bytes = d_A.sizeInBytes();
165 d_A_ = internal::DeviceBuffer::adopt(static_cast<void*>(d_A.release()), a_bytes);
166 }
167
168 factorize();
169 return *this;
170 }
171
172 ComputationInfo info() const { return solver_ctx_.info(); }
173
174 Index rows() const { return transposed_ ? n_ : m_; }
175 Index cols() const { return transposed_ ? m_ : n_; }
176
178 RealVector singularValues() const {
179 eigen_assert(solver_ctx_.info() == Success);
180 const Index k = (std::min)(m_, n_);
181 RealVector S(k);
182 solver_ctx_.download(S.data(), d_S_.get(), static_cast<size_t>(k) * sizeof(RealScalar));
183 return S;
184 }
185
187 PlainMatrix matrixU() const {
188 eigen_assert(solver_ctx_.info() == Success);
189 eigen_assert((options_ & (ComputeThinU | ComputeFullU)) && "matrixU() requires ComputeThinU or ComputeFullU");
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);
193 if (!transposed_) {
194 const Index ucols = (options_ & ComputeFullU) ? m_ : k;
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));
197 return U;
198 } else {
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();
204 }
205 }
206
208 PlainMatrix matrixV() const { return matrixVT().adjoint(); }
209
211 PlainMatrix matrixVT() const {
212 eigen_assert(solver_ctx_.info() == Success);
213 eigen_assert((options_ & (ComputeThinV | ComputeFullV)) && "matrixVT() requires ComputeThinV or ComputeFullV");
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);
217 if (!transposed_) {
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));
222 return VT;
223 } else {
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();
229 }
230 }
231
232 //
233 // These return non-owning DeviceMatrix views over the SVD's internal device storage.
234 // The view borrows the pointer: destruction does not free; the SVD object must outlive
235 // any view derived from it. For the common case (m >= n) all three accessors are pure
236 // metadata: zero kernel launches, zero allocations.
237 //
238 // For wide matrices (m < n, internally factored as A^H), original U and V^T are the
239 // adjoints of the stored buffers, so d_matrixU() / d_matrixVT() build them via a
240 // cublasXgeam into an owning temporary. d_singularValues() remains zero-copy.
241
243 DeviceMatrix<RealScalar> d_singularValues() const {
244 eigen_assert(solver_ctx_.info() == Success);
245 const Index k = (std::min)(m_, n_);
246 auto v = DeviceMatrix<RealScalar>::view(static_cast<RealScalar*>(d_S_.get()), k, 1);
247 v.recordReady(solver_ctx_.stream());
248 return v;
249 }
250
253 DeviceMatrix<Scalar> d_matrixU() const {
254 eigen_assert(solver_ctx_.info() == Success);
255 eigen_assert((options_ & (ComputeThinU | ComputeFullU)) && "d_matrixU() requires ComputeThinU or ComputeFullU");
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);
259 if (!transposed_) {
260 const Index ucols = (options_ & ComputeFullU) ? m_ : k;
261 auto v = DeviceMatrix<Scalar>::view(static_cast<Scalar*>(d_U_.get()), m_, ucols);
262 v.recordReady(solver_ctx_.stream());
263 return v;
264 }
265 // transposed: U_orig = VT_stored^H -> conjugate-transpose via cublasXgeam.
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());
274 }
275 return result;
276 }
277
280 DeviceMatrix<Scalar> d_matrixVT() const {
281 eigen_assert(solver_ctx_.info() == Success);
282 eigen_assert((options_ & (ComputeThinV | ComputeFullV)) && "d_matrixVT() requires ComputeThinV or ComputeFullV");
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);
286 if (!transposed_) {
287 const Index vtrows = (options_ & ComputeFullV) ? n_ : k;
288 auto v = DeviceMatrix<Scalar>::view(static_cast<Scalar*>(d_VT_.get()), vtrows, n_);
289 v.recordReady(solver_ctx_.stream());
290 return v;
291 }
292 // transposed: VT_orig = U_stored^H.
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());
301 }
302 return result;
303 }
304
306 Index rank(RealScalar threshold = RealScalar(-1)) const {
307 RealVector S = singularValues();
308 if (S.size() == 0) return 0;
309 if (threshold < 0) {
310 threshold = (std::max)(m_, n_) * S(0) * NumTraits<RealScalar>::epsilon();
311 }
312 return (S.array() > threshold).count();
313 }
314
316 template <typename Rhs>
317 PlainMatrix solve(const MatrixBase<Rhs>& B) const {
318 return solve_impl(B, (std::min)(m_, n_), RealScalar(0));
319 }
320
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));
326 }
327
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);
333 }
334
340 DeviceMatrix<Scalar> solve(const DeviceMatrix<Scalar>& d_B) const {
341 return solve_device_impl(d_B, (std::min)(m_, n_), RealScalar(0));
342 }
343
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));
348 }
349
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);
354 }
355
356 cudaStream_t stream() const { return solver_ctx_.stream(); }
357
358 private:
359 mutable internal::GpuSolverContext solver_ctx_;
360 internal::DeviceBuffer d_A_; // gesvd input scratch; released after factorize()
361 internal::DeviceBuffer d_U_; // grow-only
362 internal::DeviceBuffer d_S_; // grow-only
363 internal::DeviceBuffer d_VT_; // grow-only
364 // Cached inverse-diagonal for solve (built lazily, reused across solves; grow-only).
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;
370 int64_t m_ = 0;
371 int64_t n_ = 0;
372 int64_t lda_ = 0;
373 bool transposed_ = false;
374
375 // Common compute() prologue: record shape/options, reset info and cached
376 // diagonal, wait on input. Returns false (clearing state) for empty input.
377 bool begin_compute(const DeviceMatrix<Scalar>& d_A, unsigned int options) {
378 options_ = options;
379 m_ = d_A.rows();
380 n_ = d_A.cols();
381 lda_ = 0;
382 transposed_ = false;
383 diag_valid_ = false;
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();
389 return false;
390 }
391 transposed_ = (m_ < n_);
392 if (transposed_) {
393 std::swap(m_, n_);
394 lda_ = m_;
395 } else {
396 lda_ = static_cast<int64_t>(d_A.rows());
397 }
398 d_A.waitReady(solver_ctx_.stream());
399 return true;
400 }
401
402 // Wide input (m < n): produce d_A_ = A^H on device via cuBLAS geam.
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);
406 // geam: C(m×n) = alpha * op(A) + beta * op(B). beta=0, B=nullptr.
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_));
411 }
412
413 // Swap U↔V flags for the transposed case.
414 static unsigned int swap_uv_options(unsigned int opts) {
415 unsigned int result = 0;
416 if (opts & ComputeThinU) result |= ComputeThinV;
417 if (opts & ComputeFullU) result |= ComputeFullV;
418 if (opts & ComputeThinV) result |= ComputeThinU;
419 if (opts & ComputeFullV) result |= ComputeFullU;
420 return result;
421 }
422
423 static signed char jobu(unsigned int opts) {
424 if (opts & ComputeFullU) return 'A';
425 if (opts & ComputeThinU) return 'S';
426 return 'N';
427 }
428
429 static signed char jobvt(unsigned int opts) {
430 if (opts & ComputeFullV) return 'A';
431 if (opts & ComputeThinV) return 'S';
432 return 'N';
433 }
434
435 void factorize() {
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_);
439
440 solver_ctx_.mark_pending();
441
442 internal::ensure_sized(d_S_, static_cast<size_t>(k) * sizeof(RealScalar));
443
444 const unsigned int int_opts = transposed_ ? swap_uv_options(options_) : options_;
445
446 const Index ucols = (int_opts & ComputeFullU) ? m_ : ((int_opts & ComputeThinU) ? k : 0);
447 const Index vtrows = (int_opts & ComputeFullV) ? n_ : ((int_opts & ComputeThinV) ? k : 0);
448 const int64_t ldu = m_;
449 const int64_t ldvt = vtrows > 0 ? vtrows : 1;
450
451 if (ucols > 0) {
452 internal::ensure_sized(d_U_, static_cast<size_t>(m_) * static_cast<size_t>(ucols) * sizeof(Scalar));
453 }
454 if (vtrows > 0) {
455 internal::ensure_sized(d_VT_, static_cast<size_t>(vtrows) * static_cast<size_t>(n_) * sizeof(Scalar));
456 }
457
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));
464
465 solver_ctx_.ensure_scratch(dev_ws);
466 solver_ctx_.h_workspace_.resize(host_ws);
467
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()));
473
474 solver_ctx_.enqueue_info_copy();
475
476 // The input copy is pure gesvd scratch — release it now. The free is
477 // stream-ordered (or synchronous on the fallback allocator), so it waits
478 // for gesvd to retire; the memory returns to the pool instead of staying
479 // resident for the solver's lifetime.
480 d_A_ = internal::DeviceBuffer();
481 }
482
483 // Ensure d_D_ holds the kk-entry inverse diagonal for (kk, lambda).
484 // Downloads S and synchronizes only when the cached diagonal doesn't match;
485 // repeated solves with the same truncation/regularization are then free of
486 // host syncs and H2D traffic for the diagonal.
487 //
488 // For lambda == 0 we mirror Eigen's SVDBase::_solve_impl: drop singular
489 // values below S(0) * k * eps (numerical-rank truncation), so this
490 // pseudoinverse solve agrees with CPU BDCSVD::solve on near-singular A.
491 // dgmm wants the diagonal in the matrix scalar type — for complex Scalar
492 // the diagonal is still real, so we build the real values then cast.
493 void build_diag(Index kk, RealScalar lambda) const {
494 if (diag_valid_ && cached_diag_kk_ == kk && cached_diag_lambda_ == lambda) return;
495
496 const Index k = (std::min)(m_, n_);
497 RealVector S(k);
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()));
501
502 const RealScalar drop_threshold = S(0) * RealScalar(k) * NumTraits<RealScalar>::epsilon();
503 auto S_head = S.head(kk).array();
504 PlainVector D(kk);
505 if (lambda == RealScalar(0)) {
506 D = (S_head > drop_threshold).select(S_head.inverse(), RealScalar(0)).matrix().template cast<Scalar>();
507 } else {
508 D = (S_head / (S_head.square() + lambda * lambda)).matrix().template cast<Scalar>();
509 }
510
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;
517 diag_valid_ = true;
518 }
519
520 // Shared pseudoinverse application: X = V_orig * diag(D) * U_orig^H * B,
521 // entirely on device. Assumes build_diag(kk, ...) has run and B_dev/X_dev
522 // are device pointers with leading dimensions m_orig / n_orig.
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_);
527
528 auto* U_dev = static_cast<const Scalar*>(d_U_.get());
529 auto* VT_dev = static_cast<const Scalar*>(d_VT_.get());
530
531 Scalar scalars[2] = {Scalar(1), Scalar(0)};
532
533 // Step 1: tmp = U_orig^H * B (kk × nrhs).
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());
536 if (!transposed_) {
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());
541 } else {
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());
547 }
548
549 // Step 2: tmp = diag(D) * tmp on device via cublasXdgmm.
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));
552
553 // Step 3: X = V_orig * tmp (n_orig × nrhs).
554 if (!transposed_) {
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());
560 } else {
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());
565 }
566 }
567
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");
571 eigen_assert((options_ & (ComputeThinU | ComputeFullU)) && "solve requires U");
572 eigen_assert((options_ & (ComputeThinV | ComputeFullV)) && "solve requires V");
573
574 const Index m_orig = transposed_ ? n_ : m_;
575 const Index n_orig = transposed_ ? m_ : n_;
576 eigen_assert(B.rows() == m_orig);
577
578 const Index k = (std::min)(m_, n_);
579 const Index kk = (std::min)(trunc, k);
580 const Index nrhs = B.cols();
581
582 // Empty problem: no rank, no RHS, or zero domain -> result is the zero matrix.
583 if (kk == 0 || nrhs == 0 || n_orig == 0) {
584 return PlainMatrix::Zero(n_orig, nrhs);
585 }
586
587 // Enqueue the B upload before build_diag: when the diagonal must be
588 // (re)built, its S-download sync then also covers the in-flight upload —
589 // one blocking wait instead of two.
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);
595
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()));
599
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()));
604
605 return X;
606 }
607
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");
610 eigen_assert((options_ & (ComputeThinU | ComputeFullU)) && "solve requires U");
611 eigen_assert((options_ & (ComputeThinV | ComputeFullV)) && "solve requires V");
612
613 const Index m_orig = transposed_ ? n_ : m_;
614 const Index n_orig = transposed_ ? m_ : n_;
615 eigen_assert(d_B.rows() == m_orig);
616
617 const Index k = (std::min)(m_, n_);
618 const Index kk = (std::min)(trunc, k);
619 const Index nrhs = d_B.cols();
620
621 if (kk == 0 || nrhs == 0 || n_orig == 0) {
622 DeviceMatrix<Scalar> X(n_orig, nrhs);
623 X.setZero(solver_ctx_.stream());
624 return X;
625 }
626
627 d_B.waitReady(solver_ctx_.stream());
628 build_diag(kk, lambda);
629
630 DeviceMatrix<Scalar> X(n_orig, nrhs);
631 apply_pinv(d_B.data(), kk, nrhs, X.data());
632 X.recordReady(solver_ctx_.stream());
633 return X;
634 }
635};
636} // namespace gpu
637} // namespace Eigen
638
639#endif // EIGEN_GPU_SVD_H
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
ComputationInfo
Namespace containing all symbols from the Eigen library.