Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
GpuLLT.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2026 Eigen Authors
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 Cholesky (LLT) decomposition using cuSOLVER. Requires CUDA 11.4+ for the
12// cusolverDnX generic API.
13
14#ifndef EIGEN_GPU_LLT_H
15#define EIGEN_GPU_LLT_H
16
17// IWYU pragma: private
18#include "./InternalHeaderCheck.h"
19
20#include "./GpuSolverContext.h"
21
22namespace Eigen {
23namespace gpu {
40template <typename Scalar_, int UpLo_ = Lower>
41class LLT {
42 public:
43 using Scalar = Scalar_;
44 using RealScalar = typename NumTraits<Scalar>::Real;
46
47 static constexpr int UpLo = UpLo_;
48
50 LLT() = default;
51
55 explicit LLT(Context& ctx) : solver_ctx_(ctx) {}
56
58 template <typename InputType>
59 explicit LLT(const DenseBase<InputType>& A) {
60 compute(A);
61 }
62
64 explicit LLT(const DeviceMatrix<Scalar>& d_A) { compute(d_A); }
65
67 explicit LLT(DeviceMatrix<Scalar>&& d_A) { compute(std::move(d_A)); }
68
70 template <typename InputType>
71 LLT(Context& ctx, const DenseBase<InputType>& A) : solver_ctx_(ctx) {
72 compute(A);
73 }
74
76 LLT(Context& ctx, const DeviceMatrix<Scalar>& d_A) : solver_ctx_(ctx) { compute(d_A); }
77
79 LLT(Context& ctx, DeviceMatrix<Scalar>&& d_A) : solver_ctx_(ctx) { compute(std::move(d_A)); }
80
81 ~LLT() = default;
82
83 // Non-copyable (owns device memory and library handles).
84 LLT(const LLT&) = delete;
85 LLT& operator=(const LLT&) = delete;
86
87 // Movable.
88 LLT(LLT&& o) noexcept
89 : solver_ctx_(std::move(o.solver_ctx_)), d_factor_(std::move(o.d_factor_)), n_(o.n_), lda_(o.lda_) {
90 o.n_ = 0;
91 o.lda_ = 0;
92 }
93
94 LLT& operator=(LLT&& o) noexcept {
95 if (this != &o) {
96 solver_ctx_ = std::move(o.solver_ctx_);
97 d_factor_ = std::move(o.d_factor_);
98 n_ = o.n_;
99 lda_ = o.lda_;
100 o.n_ = 0;
101 o.lda_ = 0;
102 }
103 return *this;
104 }
105
108 template <typename InputType>
110 eigen_assert(A.rows() == A.cols());
111 if (!begin_compute(A.rows())) return *this;
112
113 // Ref binds column-major direct-access input in place (no host copy);
114 // row-major layouts and expressions evaluate into its temporary.
115 const Ref<const PlainMatrix> mat(A.derived());
116 lda_ = static_cast<int64_t>(mat.rows());
117 allocate_factor_storage();
118 internal::upload_host_matrix(static_cast<Scalar*>(d_factor_.get()), mat.rows(), mat.data(), mat.outerStride(),
119 mat.rows(), mat.cols(), solver_ctx_.stream());
120 EIGEN_CUDA_RUNTIME_CHECK(cudaStreamSynchronize(solver_ctx_.stream()));
121
122 factorize();
123 return *this;
124 }
125
128 eigen_assert(d_A.rows() == d_A.cols());
129 if (!begin_compute(d_A.rows())) return *this;
130
131 lda_ = static_cast<int64_t>(d_A.rows());
132 d_A.waitReady(solver_ctx_.stream());
133 allocate_factor_storage();
134 EIGEN_CUDA_RUNTIME_CHECK(
135 cudaMemcpyAsync(d_factor_.get(), d_A.data(), factorBytes(), cudaMemcpyDeviceToDevice, solver_ctx_.stream()));
136
137 factorize();
138 return *this;
139 }
140
144 if (d_A.isView()) return compute(static_cast<const DeviceMatrix<Scalar>&>(d_A));
145 eigen_assert(d_A.rows() == d_A.cols());
146 if (!begin_compute(d_A.rows())) return *this;
147
148 lda_ = static_cast<int64_t>(d_A.rows());
149 d_A.waitReady(solver_ctx_.stream());
150 d_factor_ = internal::DeviceBuffer::adopt(static_cast<void*>(d_A.release()), factorBytes());
151
152 factorize();
153 return *this;
154 }
155
157 template <typename Rhs>
158 PlainMatrix solve(const MatrixBase<Rhs>& B) const {
159 // Debug builds verify the factorization (info() synchronizes the stream on
160 // the first call after compute()); release builds skip both the check and
161 // the sync — use info() explicitly when failure must be detected.
162 eigen_assert(solver_ctx_.info() == Success && "LLT::solve called on a failed or uninitialized factorization");
163 eigen_assert(B.rows() == n_);
164
165 const Ref<const PlainMatrix> rhs(B.derived());
166 const int64_t nrhs = static_cast<int64_t>(rhs.cols());
167 const int64_t ldb = static_cast<int64_t>(rhs.rows());
168 internal::DeviceBuffer d_x(rhsBytes(nrhs, ldb));
169 internal::upload_host_matrix(static_cast<Scalar*>(d_x.get()), ldb, rhs.data(), rhs.outerStride(), rhs.rows(),
170 rhs.cols(), solver_ctx_.stream());
171 DeviceMatrix<Scalar> d_X = solve_impl(nrhs, ldb, std::move(d_x));
172
173 PlainMatrix X(n_, B.cols());
174 int solve_info = 0;
175 EIGEN_CUDA_RUNTIME_CHECK(
176 cudaMemcpyAsync(X.data(), d_X.data(), rhsBytes(nrhs, ldb), cudaMemcpyDeviceToHost, solver_ctx_.stream()));
177 EIGEN_CUDA_RUNTIME_CHECK(cudaMemcpyAsync(&solve_info, solver_ctx_.scratch_info(), sizeof(int),
178 cudaMemcpyDeviceToHost, solver_ctx_.stream()));
179 EIGEN_CUDA_RUNTIME_CHECK(cudaStreamSynchronize(solver_ctx_.stream()));
180
181 eigen_assert(solve_info == 0 && "cusolverDnXpotrs reported an error");
182 return X;
183 }
184
191 eigen_assert(solver_ctx_.info() == Success && "LLT::solve called on a failed or uninitialized factorization");
192 eigen_assert(d_B.rows() == n_);
193 d_B.waitReady(solver_ctx_.stream());
194 const int64_t nrhs = static_cast<int64_t>(d_B.cols());
195 const int64_t ldb = static_cast<int64_t>(d_B.rows());
196 internal::DeviceBuffer d_x(rhsBytes(nrhs, ldb));
197 EIGEN_CUDA_RUNTIME_CHECK(
198 cudaMemcpyAsync(d_x.get(), d_B.data(), rhsBytes(nrhs, ldb), cudaMemcpyDeviceToDevice, solver_ctx_.stream()));
199 return solve_impl(nrhs, ldb, std::move(d_x));
200 }
201
206 if (d_B.isView()) return solve(static_cast<const DeviceMatrix<Scalar>&>(d_B));
207 eigen_assert(solver_ctx_.info() == Success && "LLT::solve called on a failed or uninitialized factorization");
208 eigen_assert(d_B.rows() == n_);
209 d_B.waitReady(solver_ctx_.stream());
210 const int64_t nrhs = static_cast<int64_t>(d_B.cols());
211 const int64_t ldb = static_cast<int64_t>(d_B.rows());
212 internal::DeviceBuffer d_x = internal::DeviceBuffer::adopt(static_cast<void*>(d_B.release()), rhsBytes(nrhs, ldb));
213 return solve_impl(nrhs, ldb, std::move(d_x));
214 }
215
216 ComputationInfo info() const { return solver_ctx_.info(); }
217 Index rows() const { return n_; }
218 Index cols() const { return n_; }
219 cudaStream_t stream() const { return solver_ctx_.stream(); }
220
221 private:
222 mutable internal::GpuSolverContext solver_ctx_;
223 internal::DeviceBuffer d_factor_; // adopted from an rvalue input, else grow-only
224 int64_t n_ = 0;
225 int64_t lda_ = 0;
226
227 bool begin_compute(Index rows) {
228 n_ = rows;
229 return solver_ctx_.begin_compute(n_ != 0);
230 }
231
232 size_t factorBytes() const { return rhsBytes(n_, lda_); }
233
234 static size_t rhsBytes(int64_t cols, int64_t ld) {
235 return static_cast<size_t>(ld) * static_cast<size_t>(cols) * sizeof(Scalar);
236 }
237
238 void allocate_factor_storage() { internal::ensure_sized(d_factor_, factorBytes()); }
239
240 // Solve in place on `d_x` (which already holds B), then re-wrap as a typed
241 // DeviceMatrix carrying shape and a ready event. The release/adopt hop hands
242 // ownership of the raw cudaMalloc pointer from the untyped DeviceBuffer to
243 // the typed DeviceMatrix without copying.
244 DeviceMatrix<Scalar> solve_impl(int64_t nrhs, int64_t ldb, internal::DeviceBuffer&& d_x) const {
245 constexpr cudaDataType_t dtype = internal::cusolver_data_type<Scalar>::value;
246 constexpr cublasFillMode_t uplo = internal::cusolver_fill_mode<UpLo_>::value;
247
248 EIGEN_CUSOLVER_CHECK(cusolverDnXpotrs(solver_ctx_.cusolverHandle(), solver_ctx_.params_.p, uplo, n_, nrhs, dtype,
249 d_factor_.get(), lda_, dtype, d_x.get(), ldb, solver_ctx_.scratch_info()));
250
251 DeviceMatrix<Scalar> result =
252 DeviceMatrix<Scalar>::adopt(static_cast<Scalar*>(d_x.release()), n_, static_cast<Index>(nrhs));
253 result.recordReady(solver_ctx_.stream());
254 return result;
255 }
256
257 void factorize() {
258 constexpr cudaDataType_t dtype = internal::cusolver_data_type<Scalar>::value;
259 constexpr cublasFillMode_t uplo = internal::cusolver_fill_mode<UpLo_>::value;
260
261 solver_ctx_.mark_pending();
262
263 size_t dev_ws_bytes = 0, host_ws_bytes = 0;
264 EIGEN_CUSOLVER_CHECK(cusolverDnXpotrf_bufferSize(solver_ctx_.cusolverHandle(), solver_ctx_.params_.p, uplo, n_,
265 dtype, d_factor_.get(), lda_, dtype, &dev_ws_bytes,
266 &host_ws_bytes));
267
268 solver_ctx_.ensure_scratch(dev_ws_bytes);
269 solver_ctx_.h_workspace_.resize(host_ws_bytes);
270
271 EIGEN_CUSOLVER_CHECK(cusolverDnXpotrf(solver_ctx_.cusolverHandle(), solver_ctx_.params_.p, uplo, n_, dtype,
272 d_factor_.get(), lda_, dtype, solver_ctx_.scratch_workspace(), dev_ws_bytes,
273 host_ws_bytes > 0 ? solver_ctx_.h_workspace_.data() : nullptr, host_ws_bytes,
274 solver_ctx_.scratch_info()));
275
276 solver_ctx_.enqueue_info_copy();
277 }
278};
279} // namespace gpu
280} // namespace Eigen
281
282#endif // EIGEN_GPU_LLT_H
constexpr Scalar * data()
Unified GPU execution context owning a CUDA stream and library handles.
Definition GpuContext.h:81
RAII wrapper for a dense column-major matrix in GPU device memory.
Definition DeviceMatrix.h:122
static DeviceMatrix adopt(Scalar *device_ptr, Index rows, Index cols)
Definition DeviceMatrix.h:560
void waitReady(cudaStream_t stream) const
Definition DeviceMatrix.h:372
GPU Cholesky (LL^T) decomposition via cuSOLVER.
Definition GpuLLT.h:41
LLT(DeviceMatrix< Scalar > &&d_A)
Definition GpuLLT.h:67
LLT & compute(const DenseBase< InputType > &A)
Definition GpuLLT.h:109
DeviceMatrix< Scalar > solve(DeviceMatrix< Scalar > &&d_B) const
Definition GpuLLT.h:205
LLT & compute(DeviceMatrix< Scalar > &&d_A)
Definition GpuLLT.h:143
LLT(const DenseBase< InputType > &A)
Definition GpuLLT.h:59
LLT(const DeviceMatrix< Scalar > &d_A)
Definition GpuLLT.h:64
PlainMatrix solve(const MatrixBase< Rhs > &B) const
Definition GpuLLT.h:158
LLT(Context &ctx, const DeviceMatrix< Scalar > &d_A)
Definition GpuLLT.h:76
DeviceMatrix< Scalar > solve(const DeviceMatrix< Scalar > &d_B) const
Definition GpuLLT.h:190
LLT(Context &ctx, const DenseBase< InputType > &A)
Definition GpuLLT.h:71
LLT(Context &ctx)
Definition GpuLLT.h:55
LLT(Context &ctx, DeviceMatrix< Scalar > &&d_A)
Definition GpuLLT.h:79
LLT & compute(const DeviceMatrix< Scalar > &d_A)
Definition GpuLLT.h:127
Internal RAII owner for an untyped GPU device allocation.
Definition GpuSupport.h:293
ComputationInfo
Namespace containing all symbols from the Eigen library.