23#include "./InternalHeaderCheck.h"
25#include "./GpuSolverContext.h"
29template <
typename Scalar_>
32 using Scalar = Scalar_;
33 using RealScalar =
typename NumTraits<Scalar>::Real;
34 using PlainMatrix = Eigen::Matrix<Scalar, Dynamic, Dynamic, ColMajor>;
41 explicit QR(Context& ctx) : solver_ctx_(ctx) {}
43 template <
typename InputType>
44 explicit QR(
const DenseBase<InputType>& A) {
48 explicit QR(
const DeviceMatrix<Scalar>& d_A) { compute(d_A); }
51 explicit QR(DeviceMatrix<Scalar>&& d_A) { compute(std::move(d_A)); }
54 template <
typename InputType>
55 QR(Context& ctx,
const DenseBase<InputType>& A) : solver_ctx_(ctx) {
60 QR(Context& ctx,
const DeviceMatrix<Scalar>& d_A) : solver_ctx_(ctx) { compute(d_A); }
63 QR(Context& ctx, DeviceMatrix<Scalar>&& d_A) : solver_ctx_(ctx) { compute(std::move(d_A)); }
67 QR(
const QR&) =
delete;
68 QR& operator=(
const QR&) =
delete;
71 : solver_ctx_(std::move(o.solver_ctx_)),
72 d_qr_(std::move(o.d_qr_)),
73 d_tau_(std::move(o.d_tau_)),
77 transposed_(o.transposed_) {
81 o.transposed_ =
false;
84 QR& operator=(QR&& o)
noexcept {
86 solver_ctx_ = std::move(o.solver_ctx_);
87 d_qr_ = std::move(o.d_qr_);
88 d_tau_ = std::move(o.d_tau_);
92 transposed_ = o.transposed_;
96 o.transposed_ =
false;
101 template <
typename InputType>
102 QR& compute(
const DenseBase<InputType>& A) {
110 QR& compute(
const DeviceMatrix<Scalar>& d_A) {
111 if (!begin_compute(d_A))
return *
this;
114 transpose_into_factor(d_A);
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()));
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;
134 transpose_into_factor(d_A);
136 d_qr_ = internal::DeviceBuffer::adopt(
static_cast<void*
>(d_A.release()), factorBytes());
146 template <
typename Rhs>
147 PlainMatrix solve(
const MatrixBase<Rhs>& B)
const {
151 eigen_assert(solver_ctx_.info() ==
Success &&
"QR::solve called on a failed or uninitialized factorization");
152 eigen_assert(B.rows() == m_);
154 const Ref<const PlainMatrix> rhs(B.derived());
155 const Index nrhs = rhs.cols();
158 return solve_overdetermined_host(rhs);
160 return solve_underdetermined_host(rhs, nrhs);
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());
171 return solve_overdetermined_device(d_B);
173 return solve_underdetermined_device(d_B);
178 Index rows()
const {
return m_; }
179 Index cols()
const {
return n_; }
180 cudaStream_t stream()
const {
return solver_ctx_.stream(); }
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));
191 PlainMatrix R = qr_full.topRows(k()).template triangularView<Upper>();
196 mutable internal::GpuSolverContext solver_ctx_;
199 internal::DeviceBuffer d_qr_;
200 internal::DeviceBuffer d_tau_;
204 bool transposed_ =
false;
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_); }
211 size_t factorBytes()
const {
return static_cast<size_t>(lda_) *
static_cast<size_t>(factor_cols()) *
sizeof(Scalar); }
215 bool begin_compute(
const DeviceMatrix<Scalar>& d_A) {
218 if (!solver_ctx_.begin_compute(m_ != 0 && n_ != 0)) {
219 d_qr_ = internal::DeviceBuffer();
220 d_tau_ = internal::DeviceBuffer();
223 transposed_ = (m_ < n_);
224 lda_ =
static_cast<int64_t
>(transposed_ ? n_ : m_);
225 d_A.waitReady(solver_ctx_.stream());
229 void allocate_factor_storage(
size_t mat_bytes) { internal::ensure_sized(d_qr_, mat_bytes); }
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_));
241 constexpr cudaDataType_t dtype = internal::cusolver_data_type<Scalar>::value;
243 solver_ctx_.mark_pending();
245 internal::ensure_sized(d_tau_,
static_cast<size_t>(k()) *
sizeof(Scalar));
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));
253 solver_ctx_.ensure_scratch(dev_ws);
254 solver_ctx_.h_workspace_.resize(host_ws);
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()));
261 solver_ctx_.enqueue_info_copy();
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);
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));
279 solver_ctx_.ensure_scratch(
static_cast<size_t>(lwork) *
sizeof(Scalar));
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()));
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);
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);
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());
300 apply_QH(d_B.get(), m_, nrhs);
301 trsm_R(d_B.get(), m_, nrhs, CUBLAS_OP_N);
303 PlainMatrix X(n_, nrhs);
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()));
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()));
314 EIGEN_CUDA_RUNTIME_CHECK(cudaStreamSynchronize(solver_ctx_.stream()));
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);
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()));
326 apply_QH(d_work.get(), m_, nrhs);
327 trsm_R(d_work.get(), m_, nrhs, CUBLAS_OP_N);
330 DeviceMatrix<Scalar> result =
332 result.recordReady(solver_ctx_.stream());
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());
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);
352 internal::DeviceBuffer d_X(x_bytes);
354 EIGEN_CUDA_RUNTIME_CHECK(cudaMemsetAsync(d_X.get(), 0, x_bytes, solver_ctx_.stream()));
357 internal::upload_host_matrix(
static_cast<Scalar*
>(d_X.get()), n_, rhs.data(), rhs.outerStride(), m_, nrhs,
358 solver_ctx_.stream());
360 trsm_R(d_X.get(), n_, nrhs, trsm_op_conj_trans());
361 apply_Q(CUBLAS_OP_N, d_X.get(), n_, nrhs);
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()));
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);
374 internal::DeviceBuffer d_X(x_bytes);
375 EIGEN_CUDA_RUNTIME_CHECK(cudaMemsetAsync(d_X.get(), 0, x_bytes, solver_ctx_.stream()));
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()));
384 trsm_R(d_X.get(), n_, nrhs, trsm_op_conj_trans());
385 apply_Q(CUBLAS_OP_N, d_X.get(), n_, nrhs);
387 DeviceMatrix<Scalar> result =
389 result.recordReady(solver_ctx_.stream());
393 static cublasOperation_t trsm_op_conj_trans() {
return NumTraits<Scalar>::IsComplex ? CUBLAS_OP_C : CUBLAS_OP_T; }
397 void trsm_R(
void* d_B, int64_t ldb, int64_t nrhs, cublasOperation_t op)
const {
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));
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
Namespace containing all symbols from the Eigen library.