Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
GpuQR.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 QR decomposition using cuSOLVER, wrapping cusolverDnXgeqrf, cusolverDnXormqr
12// (apply Q), and cublasXtrsm (triangular solve on R). Q is never formed
13// explicitly.
14//
15// Both shapes are handled transparently: for m >= n the factorization is
16// A = Q R and solve() is least-squares; for m < n it is A^H = Q R internally and
17// solve() is minimum-norm.
18
19#ifndef EIGEN_GPU_QR_H
20#define EIGEN_GPU_QR_H
21
22// IWYU pragma: private
23#include "./InternalHeaderCheck.h"
24
25#include "./GpuSolverContext.h"
26
27namespace Eigen {
28namespace gpu {
29template <typename Scalar_>
30class QR {
31 public:
32 using Scalar = Scalar_;
33 using RealScalar = typename NumTraits<Scalar>::Real;
34 using PlainMatrix = Eigen::Matrix<Scalar, Dynamic, Dynamic, ColMajor>;
35
36 QR() = default;
37
41 explicit QR(Context& ctx) : solver_ctx_(ctx) {}
42
43 template <typename InputType>
44 explicit QR(const DenseBase<InputType>& A) {
45 compute(A);
46 }
47
48 explicit QR(const DeviceMatrix<Scalar>& d_A) { compute(d_A); }
49
51 explicit QR(DeviceMatrix<Scalar>&& d_A) { compute(std::move(d_A)); }
52
54 template <typename InputType>
55 QR(Context& ctx, const DenseBase<InputType>& A) : solver_ctx_(ctx) {
56 compute(A);
57 }
58
60 QR(Context& ctx, const DeviceMatrix<Scalar>& d_A) : solver_ctx_(ctx) { compute(d_A); }
61
63 QR(Context& ctx, DeviceMatrix<Scalar>&& d_A) : solver_ctx_(ctx) { compute(std::move(d_A)); }
64
65 ~QR() = default;
66
67 QR(const QR&) = delete;
68 QR& operator=(const QR&) = delete;
69
70 QR(QR&& o) noexcept
71 : solver_ctx_(std::move(o.solver_ctx_)),
72 d_qr_(std::move(o.d_qr_)),
73 d_tau_(std::move(o.d_tau_)),
74 m_(o.m_),
75 n_(o.n_),
76 lda_(o.lda_),
77 transposed_(o.transposed_) {
78 o.m_ = 0;
79 o.n_ = 0;
80 o.lda_ = 0;
81 o.transposed_ = false;
82 }
83
84 QR& operator=(QR&& o) noexcept {
85 if (this != &o) {
86 solver_ctx_ = std::move(o.solver_ctx_);
87 d_qr_ = std::move(o.d_qr_);
88 d_tau_ = std::move(o.d_tau_);
89 m_ = o.m_;
90 n_ = o.n_;
91 lda_ = o.lda_;
92 transposed_ = o.transposed_;
93 o.m_ = 0;
94 o.n_ = 0;
95 o.lda_ = 0;
96 o.transposed_ = false;
97 }
98 return *this;
99 }
100
101 template <typename InputType>
102 QR& compute(const DenseBase<InputType>& A) {
103 // Upload to device, then delegate to the adopting overload — the freshly
104 // uploaded matrix is factored in place (geqrf overwrites its input), so no
105 // second device copy is made. The wide-matrix transpose runs on the GPU
106 // (via cublasXgeam) inside the device-input path; no host transpose.
107 return compute(DeviceMatrix<Scalar>::fromHost(A.derived(), solver_ctx_.stream()));
108 }
109
110 QR& compute(const DeviceMatrix<Scalar>& d_A) {
111 if (!begin_compute(d_A)) return *this;
112
113 if (transposed_) {
114 transpose_into_factor(d_A);
115 } else {
116 const size_t mat_bytes = factorBytes();
117 allocate_factor_storage(mat_bytes);
118 EIGEN_CUDA_RUNTIME_CHECK(
119 cudaMemcpyAsync(d_qr_.get(), d_A.data(), mat_bytes, cudaMemcpyDeviceToDevice, solver_ctx_.stream()));
120 }
121
122 factorize();
123 return *this;
124 }
125
129 QR& compute(DeviceMatrix<Scalar>&& d_A) {
130 if (d_A.isView()) return compute(static_cast<const DeviceMatrix<Scalar>&>(d_A));
131 if (!begin_compute(d_A)) return *this;
132
133 if (transposed_) {
134 transpose_into_factor(d_A);
135 } else {
136 d_qr_ = internal::DeviceBuffer::adopt(static_cast<void*>(d_A.release()), factorBytes());
137 }
138
139 factorize();
140 return *this;
141 }
142
146 template <typename Rhs>
147 PlainMatrix solve(const MatrixBase<Rhs>& B) const {
148 // Debug builds verify the factorization (info() synchronizes the stream on
149 // the first call after compute()); release builds skip both the check and
150 // the sync — use info() explicitly when failure must be detected.
151 eigen_assert(solver_ctx_.info() == Success && "QR::solve called on a failed or uninitialized factorization");
152 eigen_assert(B.rows() == m_);
153
154 const Ref<const PlainMatrix> rhs(B.derived());
155 const Index nrhs = rhs.cols();
156
157 if (!transposed_) {
158 return solve_overdetermined_host(rhs);
159 }
160 return solve_underdetermined_host(rhs, nrhs);
161 }
162
165 DeviceMatrix<Scalar> solve(const DeviceMatrix<Scalar>& d_B) const {
166 eigen_assert(solver_ctx_.info() == Success && "QR::solve called on a failed or uninitialized factorization");
167 eigen_assert(d_B.rows() == m_);
168 d_B.waitReady(solver_ctx_.stream());
169
170 if (!transposed_) {
171 return solve_overdetermined_device(d_B);
172 }
173 return solve_underdetermined_device(d_B);
174 }
175
176 ComputationInfo info() const { return solver_ctx_.info(); }
177
178 Index rows() const { return m_; }
179 Index cols() const { return n_; }
180 cudaStream_t stream() const { return solver_ctx_.stream(); }
181
183 PlainMatrix matrixR() const {
184 eigen_assert(solver_ctx_.info() == Success);
185 eigen_assert(!transposed_ && "matrixR() not available when m < n (we factored A^H internally)");
186 PlainMatrix qr_full(m_, n_);
187 if (m_ > 0 && n_ > 0) {
188 solver_ctx_.download(qr_full.data(), d_qr_.get(),
189 static_cast<size_t>(lda_) * static_cast<size_t>(n_) * sizeof(Scalar));
190 }
191 PlainMatrix R = qr_full.topRows(k()).template triangularView<Upper>();
192 return R;
193 }
194
195 private:
196 mutable internal::GpuSolverContext solver_ctx_;
197 // QR factors (reflectors below diag, R above). Host and rvalue input are
198 // adopted when m >= n; the copying paths are grow-only.
199 internal::DeviceBuffer d_qr_;
200 internal::DeviceBuffer d_tau_; // grow-only; Householder scalars (length k)
201 int64_t m_ = 0; // original A.rows()
202 int64_t n_ = 0; // original A.cols()
203 int64_t lda_ = 0; // factor leading dim = max(m_, n_)
204 bool transposed_ = false; // true iff m_ < n_, i.e. A^H was factored
205
206 // The factored matrix is always tall: rows >= cols.
207 int64_t factor_rows() const { return transposed_ ? n_ : m_; }
208 int64_t factor_cols() const { return transposed_ ? m_ : n_; }
209 int64_t k() const { return (std::min)(m_, n_); }
210
211 size_t factorBytes() const { return static_cast<size_t>(lda_) * static_cast<size_t>(factor_cols()) * sizeof(Scalar); }
212
213 // Common compute() prologue: record shape, reset info, wait on the input.
214 // Returns false (and clears stale factors) for empty input.
215 bool begin_compute(const DeviceMatrix<Scalar>& d_A) {
216 m_ = d_A.rows();
217 n_ = d_A.cols();
218 if (!solver_ctx_.begin_compute(m_ != 0 && n_ != 0)) {
219 d_qr_ = internal::DeviceBuffer();
220 d_tau_ = internal::DeviceBuffer();
221 return false;
222 }
223 transposed_ = (m_ < n_);
224 lda_ = static_cast<int64_t>(transposed_ ? n_ : m_);
225 d_A.waitReady(solver_ctx_.stream());
226 return true;
227 }
228
229 void allocate_factor_storage(size_t mat_bytes) { internal::ensure_sized(d_qr_, mat_bytes); }
230
231 // Wide input (m < n): factor A^H, produced on device via cuBLAS geam.
232 void transpose_into_factor(const DeviceMatrix<Scalar>& d_A) {
233 allocate_factor_storage(factorBytes());
234 Scalar alpha_one(1), beta_zero(0);
235 EIGEN_CUBLAS_CHECK(internal::cublasXgeam(solver_ctx_.cublasHandle(), CUBLAS_OP_C, CUBLAS_OP_N, n_, m_, &alpha_one,
236 d_A.data(), d_A.rows(), &beta_zero, static_cast<const Scalar*>(nullptr),
237 n_, static_cast<Scalar*>(d_qr_.get()), n_));
238 }
239
240 void factorize() {
241 constexpr cudaDataType_t dtype = internal::cusolver_data_type<Scalar>::value;
242
243 solver_ctx_.mark_pending();
244
245 internal::ensure_sized(d_tau_, static_cast<size_t>(k()) * sizeof(Scalar));
246
247 const int64_t fm = factor_rows();
248 const int64_t fn = factor_cols();
249 size_t dev_ws = 0, host_ws = 0;
250 EIGEN_CUSOLVER_CHECK(cusolverDnXgeqrf_bufferSize(solver_ctx_.cusolverHandle(), solver_ctx_.params_.p, fm, fn, dtype,
251 d_qr_.get(), lda_, dtype, d_tau_.get(), dtype, &dev_ws, &host_ws));
252
253 solver_ctx_.ensure_scratch(dev_ws);
254 solver_ctx_.h_workspace_.resize(host_ws);
255
256 EIGEN_CUSOLVER_CHECK(
257 cusolverDnXgeqrf(solver_ctx_.cusolverHandle(), solver_ctx_.params_.p, fm, fn, dtype, d_qr_.get(), lda_, dtype,
258 d_tau_.get(), dtype, solver_ctx_.scratch_workspace(), dev_ws,
259 host_ws > 0 ? solver_ctx_.h_workspace_.data() : nullptr, host_ws, solver_ctx_.scratch_info()));
260
261 solver_ctx_.enqueue_info_copy();
262 }
263
264 // Applies Q (CUBLAS_OP_N) or Q^H (CUBLAS_OP_T for real, CUBLAS_OP_C for complex)
265 // in place. Workspace comes from solver_ctx_'s grow-only scratch, so there is no
266 // per-call malloc/free.
267 void apply_Q(cublasOperation_t op, void* d_B, int64_t ldb, int64_t nrhs) const {
268 const int im = internal::to_blas_int(factor_rows());
269 const int in = internal::to_blas_int(nrhs);
270 const int ik = internal::to_blas_int(k());
271 const int ilda = internal::to_blas_int(lda_);
272 const int ildb = internal::to_blas_int(ldb);
273
274 int lwork = 0;
275 EIGEN_CUSOLVER_CHECK(internal::cusolverDnXormqr_bufferSize(
276 solver_ctx_.cusolverHandle(), CUBLAS_SIDE_LEFT, op, im, in, ik, static_cast<const Scalar*>(d_qr_.get()), ilda,
277 static_cast<const Scalar*>(d_tau_.get()), static_cast<const Scalar*>(d_B), ildb, &lwork));
278
279 solver_ctx_.ensure_scratch(static_cast<size_t>(lwork) * sizeof(Scalar));
280
281 EIGEN_CUSOLVER_CHECK(internal::cusolverDnXormqr(
282 solver_ctx_.cusolverHandle(), CUBLAS_SIDE_LEFT, op, im, in, ik, static_cast<const Scalar*>(d_qr_.get()), ilda,
283 static_cast<const Scalar*>(d_tau_.get()), static_cast<Scalar*>(d_B), ildb,
284 static_cast<Scalar*>(solver_ctx_.scratch_workspace()), lwork, solver_ctx_.scratch_info()));
285 }
286
287 void apply_QH(void* d_B, int64_t ldb, int64_t nrhs) const {
288 constexpr cublasOperation_t trans = NumTraits<Scalar>::IsComplex ? CUBLAS_OP_C : CUBLAS_OP_T;
289 apply_Q(trans, d_B, ldb, nrhs);
290 }
291
292 PlainMatrix solve_overdetermined_host(const Ref<const PlainMatrix>& rhs) const {
293 const Index nrhs = rhs.cols();
294 const size_t b_bytes = static_cast<size_t>(m_) * static_cast<size_t>(nrhs) * sizeof(Scalar);
295
296 internal::DeviceBuffer d_B(b_bytes);
297 internal::upload_host_matrix(static_cast<Scalar*>(d_B.get()), m_, rhs.data(), rhs.outerStride(), m_, nrhs,
298 solver_ctx_.stream());
299
300 apply_QH(d_B.get(), m_, nrhs);
301 trsm_R(d_B.get(), m_, nrhs, /*op=*/CUBLAS_OP_N);
302
303 PlainMatrix X(n_, nrhs);
304 if (m_ == n_) {
305 EIGEN_CUDA_RUNTIME_CHECK(cudaMemcpyAsync(X.data(), d_B.get(),
306 static_cast<size_t>(n_) * static_cast<size_t>(nrhs) * sizeof(Scalar),
307 cudaMemcpyDeviceToHost, solver_ctx_.stream()));
308 } else {
309 EIGEN_CUDA_RUNTIME_CHECK(cudaMemcpy2DAsync(X.data(), static_cast<size_t>(n_) * sizeof(Scalar), d_B.get(),
310 static_cast<size_t>(m_) * sizeof(Scalar),
311 static_cast<size_t>(n_) * sizeof(Scalar), static_cast<size_t>(nrhs),
312 cudaMemcpyDeviceToHost, solver_ctx_.stream()));
313 }
314 EIGEN_CUDA_RUNTIME_CHECK(cudaStreamSynchronize(solver_ctx_.stream()));
315 return X;
316 }
317
318 DeviceMatrix<Scalar> solve_overdetermined_device(const DeviceMatrix<Scalar>& d_B) const {
319 const Index nrhs = d_B.cols();
320 const size_t b_bytes = static_cast<size_t>(m_) * static_cast<size_t>(nrhs) * sizeof(Scalar);
321
322 internal::DeviceBuffer d_work(b_bytes);
323 EIGEN_CUDA_RUNTIME_CHECK(
324 cudaMemcpyAsync(d_work.get(), d_B.data(), b_bytes, cudaMemcpyDeviceToDevice, solver_ctx_.stream()));
325
326 apply_QH(d_work.get(), m_, nrhs);
327 trsm_R(d_work.get(), m_, nrhs, /*op=*/CUBLAS_OP_N);
328
329 if (m_ == n_) {
330 DeviceMatrix<Scalar> result =
331 DeviceMatrix<Scalar>::adopt(static_cast<Scalar*>(d_work.release()), n_, static_cast<Index>(nrhs));
332 result.recordReady(solver_ctx_.stream());
333 return result;
334 }
335 DeviceMatrix<Scalar> result(n_, static_cast<Index>(nrhs));
336 EIGEN_CUDA_RUNTIME_CHECK(cudaMemcpy2DAsync(result.data(), static_cast<size_t>(n_) * sizeof(Scalar), d_work.get(),
337 static_cast<size_t>(m_) * sizeof(Scalar),
338 static_cast<size_t>(n_) * sizeof(Scalar), static_cast<size_t>(nrhs),
339 cudaMemcpyDeviceToDevice, solver_ctx_.stream()));
340 result.recordReady(solver_ctx_.stream());
341 return result;
342 }
343
344 //
345 // We factored A^H = Q R, so A = R^H Q^H. Solving A X = B for X with min ||X||:
346 // z = R^{-H} B (m × nrhs, occupies top m rows of an n × nrhs buffer)
347 // X = Q [z; 0] (n × nrhs)
348
349 PlainMatrix solve_underdetermined_host(const Ref<const PlainMatrix>& rhs, Index nrhs) const {
350 const size_t x_bytes = static_cast<size_t>(n_) * static_cast<size_t>(nrhs) * sizeof(Scalar);
351
352 internal::DeviceBuffer d_X(x_bytes);
353 // Zero the full n × nrhs buffer; B will overwrite the top m × nrhs block.
354 EIGEN_CUDA_RUNTIME_CHECK(cudaMemsetAsync(d_X.get(), 0, x_bytes, solver_ctx_.stream()));
355
356 // B (m × nrhs) into the top of d_X (leading dim n).
357 internal::upload_host_matrix(static_cast<Scalar*>(d_X.get()), n_, rhs.data(), rhs.outerStride(), m_, nrhs,
358 solver_ctx_.stream());
359
360 trsm_R(d_X.get(), n_, nrhs, trsm_op_conj_trans());
361 apply_Q(CUBLAS_OP_N, d_X.get(), n_, nrhs);
362
363 PlainMatrix X(n_, nrhs);
364 EIGEN_CUDA_RUNTIME_CHECK(
365 cudaMemcpyAsync(X.data(), d_X.get(), x_bytes, cudaMemcpyDeviceToHost, solver_ctx_.stream()));
366 EIGEN_CUDA_RUNTIME_CHECK(cudaStreamSynchronize(solver_ctx_.stream()));
367 return X;
368 }
369
370 DeviceMatrix<Scalar> solve_underdetermined_device(const DeviceMatrix<Scalar>& d_B) const {
371 const Index nrhs = d_B.cols();
372 const size_t x_bytes = static_cast<size_t>(n_) * static_cast<size_t>(nrhs) * sizeof(Scalar);
373
374 internal::DeviceBuffer d_X(x_bytes);
375 EIGEN_CUDA_RUNTIME_CHECK(cudaMemsetAsync(d_X.get(), 0, x_bytes, solver_ctx_.stream()));
376
377 if (m_ > 0 && nrhs > 0) {
378 EIGEN_CUDA_RUNTIME_CHECK(cudaMemcpy2DAsync(d_X.get(), static_cast<size_t>(n_) * sizeof(Scalar), d_B.data(),
379 static_cast<size_t>(m_) * sizeof(Scalar),
380 static_cast<size_t>(m_) * sizeof(Scalar), static_cast<size_t>(nrhs),
381 cudaMemcpyDeviceToDevice, solver_ctx_.stream()));
382 }
383
384 trsm_R(d_X.get(), n_, nrhs, trsm_op_conj_trans());
385 apply_Q(CUBLAS_OP_N, d_X.get(), n_, nrhs);
386
387 DeviceMatrix<Scalar> result =
388 DeviceMatrix<Scalar>::adopt(static_cast<Scalar*>(d_X.release()), n_, static_cast<Index>(nrhs));
389 result.recordReady(solver_ctx_.stream());
390 return result;
391 }
392
393 static cublasOperation_t trsm_op_conj_trans() { return NumTraits<Scalar>::IsComplex ? CUBLAS_OP_C : CUBLAS_OP_T; }
394
395 // X := op(R)^{-1} B, in place on B. The m >= n branch passes CUBLAS_OP_N to
396 // solve R X = (Q^H B)[:k,:]; the m < n branch passes OP_T/OP_C to solve R^H z = B.
397 void trsm_R(void* d_B, int64_t ldb, int64_t nrhs, cublasOperation_t op) const {
398 Scalar alpha(1);
399 EIGEN_CUBLAS_CHECK(internal::cublasXtrsm(
400 solver_ctx_.cublasHandle(), CUBLAS_SIDE_LEFT, CUBLAS_FILL_MODE_UPPER, op, CUBLAS_DIAG_NON_UNIT, k(), nrhs,
401 &alpha, static_cast<const Scalar*>(d_qr_.get()), lda_, static_cast<Scalar*>(d_B), ldb));
402 }
403};
404} // namespace gpu
405} // namespace Eigen
406
407#endif // EIGEN_GPU_QR_H
static DeviceMatrix fromHost(const DenseBase< Derived > &host, cudaStream_t stream=nullptr)
Definition DeviceMatrix.h:226
static DeviceMatrix adopt(Scalar *device_ptr, Index rows, Index cols)
Definition DeviceMatrix.h:560
ComputationInfo
Namespace containing all symbols from the Eigen library.