Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
GpuLU.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 partial-pivoting LU decomposition using cuSOLVER, wrapping
12// cusolverDnXgetrf and cusolverDnXgetrs.
13
14#ifndef EIGEN_GPU_LU_H
15#define EIGEN_GPU_LU_H
16
17// IWYU pragma: private
18#include "./InternalHeaderCheck.h"
19
20#include "./GpuSolverContext.h"
21
22namespace Eigen {
23namespace gpu {
36template <typename Scalar_>
37class LU {
38 public:
39 using Scalar = Scalar_;
40 using RealScalar = typename NumTraits<Scalar>::Real;
42
43 LU() = default;
44
48 explicit LU(Context& ctx) : solver_ctx_(ctx) {}
49
50 template <typename InputType>
51 explicit LU(const DenseBase<InputType>& A) {
52 compute(A);
53 }
54
56 explicit LU(const DeviceMatrix<Scalar>& d_A) { compute(d_A); }
57
59 explicit LU(DeviceMatrix<Scalar>&& d_A) { compute(std::move(d_A)); }
60
62 template <typename InputType>
63 LU(Context& ctx, const DenseBase<InputType>& A) : solver_ctx_(ctx) {
64 compute(A);
65 }
66
68 LU(Context& ctx, const DeviceMatrix<Scalar>& d_A) : solver_ctx_(ctx) { compute(d_A); }
69
71 LU(Context& ctx, DeviceMatrix<Scalar>&& d_A) : solver_ctx_(ctx) { compute(std::move(d_A)); }
72
73 ~LU() = default;
74
75 LU(const LU&) = delete;
76 LU& operator=(const LU&) = delete;
77
78 LU(LU&& o) noexcept
79 : solver_ctx_(std::move(o.solver_ctx_)),
80 d_lu_(std::move(o.d_lu_)),
81 d_ipiv_(std::move(o.d_ipiv_)),
82 n_(o.n_),
83 lda_(o.lda_) {
84 o.n_ = 0;
85 o.lda_ = 0;
86 }
87
88 LU& operator=(LU&& o) noexcept {
89 if (this != &o) {
90 solver_ctx_ = std::move(o.solver_ctx_);
91 d_lu_ = std::move(o.d_lu_);
92 d_ipiv_ = std::move(o.d_ipiv_);
93 n_ = o.n_;
94 lda_ = o.lda_;
95 o.n_ = 0;
96 o.lda_ = 0;
97 }
98 return *this;
99 }
100
103 template <typename InputType>
105 eigen_assert(A.rows() == A.cols() && "LU requires a square matrix");
106 if (!begin_compute(A.rows())) return *this;
107
108 // Ref binds column-major direct-access input in place (no host copy);
109 // row-major layouts and expressions evaluate into its temporary.
110 const Ref<const PlainMatrix> mat(A.derived());
111 lda_ = static_cast<int64_t>(mat.rows());
112 allocate_lu_storage();
113 internal::upload_host_matrix(static_cast<Scalar*>(d_lu_.get()), mat.rows(), mat.data(), mat.outerStride(),
114 mat.rows(), mat.cols(), solver_ctx_.stream());
115 EIGEN_CUDA_RUNTIME_CHECK(cudaStreamSynchronize(solver_ctx_.stream()));
116
117 factorize();
118 return *this;
119 }
120
123 eigen_assert(d_A.rows() == d_A.cols() && "LU requires a square matrix");
124 if (!begin_compute(d_A.rows())) return *this;
125
126 lda_ = static_cast<int64_t>(d_A.rows());
127 d_A.waitReady(solver_ctx_.stream());
128 allocate_lu_storage();
129 EIGEN_CUDA_RUNTIME_CHECK(
130 cudaMemcpyAsync(d_lu_.get(), d_A.data(), matrixBytes(), cudaMemcpyDeviceToDevice, solver_ctx_.stream()));
131
132 factorize();
133 return *this;
134 }
135
139 if (d_A.isView()) return compute(static_cast<const DeviceMatrix<Scalar>&>(d_A));
140 eigen_assert(d_A.rows() == d_A.cols() && "LU requires a square matrix");
141 if (!begin_compute(d_A.rows())) return *this;
142
143 lda_ = static_cast<int64_t>(d_A.rows());
144 d_A.waitReady(solver_ctx_.stream());
145 d_lu_ = internal::DeviceBuffer::adopt(static_cast<void*>(d_A.release()), matrixBytes());
146
147 factorize();
148 return *this;
149 }
150
156 template <typename Rhs>
157 PlainMatrix solve(const MatrixBase<Rhs>& B, GpuOp op = GpuOp::NoTrans) const {
158 // Debug builds verify the factorization (info() synchronizes the stream on
159 // the first call after compute()); release builds skip both the check and
160 // the sync — use info() explicitly when failure must be detected.
161 eigen_assert(solver_ctx_.info() == Success && "LU::solve called on a failed or uninitialized factorization");
162 eigen_assert(B.rows() == n_);
163
164 const Ref<const PlainMatrix> rhs(B.derived());
165 const int64_t nrhs = static_cast<int64_t>(rhs.cols());
166 const int64_t ldb = static_cast<int64_t>(rhs.rows());
167 internal::DeviceBuffer d_x(matrixBytes(nrhs, ldb));
168 internal::upload_host_matrix(static_cast<Scalar*>(d_x.get()), ldb, rhs.data(), rhs.outerStride(), rhs.rows(),
169 rhs.cols(), solver_ctx_.stream());
170 DeviceMatrix<Scalar> d_X = solve_impl(nrhs, ldb, op, std::move(d_x));
171
172 PlainMatrix X(n_, B.cols());
173 int solve_info = 0;
174 EIGEN_CUDA_RUNTIME_CHECK(
175 cudaMemcpyAsync(X.data(), d_X.data(), matrixBytes(nrhs, ldb), cudaMemcpyDeviceToHost, solver_ctx_.stream()));
176 EIGEN_CUDA_RUNTIME_CHECK(cudaMemcpyAsync(&solve_info, solver_ctx_.scratch_info(), sizeof(int),
177 cudaMemcpyDeviceToHost, solver_ctx_.stream()));
178 EIGEN_CUDA_RUNTIME_CHECK(cudaStreamSynchronize(solver_ctx_.stream()));
179
180 eigen_assert(solve_info == 0 && "cusolverDnXgetrs reported an error");
181 return X;
182 }
183
189 DeviceMatrix<Scalar> solve(const DeviceMatrix<Scalar>& d_B, GpuOp op = GpuOp::NoTrans) const {
190 eigen_assert(solver_ctx_.info() == Success && "LU::solve called on a failed or uninitialized factorization");
191 eigen_assert(d_B.rows() == n_);
192 d_B.waitReady(solver_ctx_.stream());
193 const int64_t nrhs = static_cast<int64_t>(d_B.cols());
194 const int64_t ldb = static_cast<int64_t>(d_B.rows());
195 internal::DeviceBuffer d_x(matrixBytes(nrhs, ldb));
196 EIGEN_CUDA_RUNTIME_CHECK(
197 cudaMemcpyAsync(d_x.get(), d_B.data(), matrixBytes(nrhs, ldb), cudaMemcpyDeviceToDevice, solver_ctx_.stream()));
198 return solve_impl(nrhs, ldb, op, std::move(d_x));
199 }
200
204 DeviceMatrix<Scalar> solve(DeviceMatrix<Scalar>&& d_B, GpuOp op = GpuOp::NoTrans) const {
205 if (d_B.isView()) return solve(static_cast<const DeviceMatrix<Scalar>&>(d_B), op);
206 eigen_assert(solver_ctx_.info() == Success && "LU::solve called on a failed or uninitialized factorization");
207 eigen_assert(d_B.rows() == n_);
208 d_B.waitReady(solver_ctx_.stream());
209 const int64_t nrhs = static_cast<int64_t>(d_B.cols());
210 const int64_t ldb = static_cast<int64_t>(d_B.rows());
212 internal::DeviceBuffer::adopt(static_cast<void*>(d_B.release()), matrixBytes(nrhs, ldb));
213 return solve_impl(nrhs, ldb, op, 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_lu_; // adopted from an rvalue input, else grow-only
224 internal::DeviceBuffer d_ipiv_; // grow-only
225 int64_t n_ = 0;
226 int64_t lda_ = 0;
227
228 bool begin_compute(Index rows) {
229 n_ = rows;
230 return solver_ctx_.begin_compute(n_ != 0);
231 }
232
233 size_t matrixBytes() const { return matrixBytes(n_, lda_); }
234
235 static size_t matrixBytes(int64_t cols, int64_t ld) {
236 return static_cast<size_t>(ld) * static_cast<size_t>(cols) * sizeof(Scalar);
237 }
238
239 void allocate_lu_storage() { internal::ensure_sized(d_lu_, matrixBytes()); }
240
241 // Solve in place on `d_x` (which already holds B), then re-wrap as a typed
242 // DeviceMatrix carrying shape and a ready event. The release/adopt hop hands
243 // ownership of the raw cudaMalloc pointer from the untyped DeviceBuffer to
244 // the typed DeviceMatrix without copying.
245 DeviceMatrix<Scalar> solve_impl(int64_t nrhs, int64_t ldb, GpuOp op, internal::DeviceBuffer&& d_x) const {
246 constexpr cudaDataType_t dtype = internal::cusolver_data_type<Scalar>::value;
247 const cublasOperation_t trans = internal::to_cublas_op(op);
248
249 EIGEN_CUSOLVER_CHECK(cusolverDnXgetrs(solver_ctx_.cusolverHandle(), solver_ctx_.params_.p, trans, n_, nrhs, dtype,
250 d_lu_.get(), lda_, static_cast<const int64_t*>(d_ipiv_.get()), dtype,
251 d_x.get(), ldb, solver_ctx_.scratch_info()));
252
253 DeviceMatrix<Scalar> result =
254 DeviceMatrix<Scalar>::adopt(static_cast<Scalar*>(d_x.release()), n_, static_cast<Index>(nrhs));
255 result.recordReady(solver_ctx_.stream());
256 return result;
257 }
258
259 void factorize() {
260 constexpr cudaDataType_t dtype = internal::cusolver_data_type<Scalar>::value;
261 const size_t ipiv_bytes = static_cast<size_t>(n_) * sizeof(int64_t);
262
263 solver_ctx_.mark_pending();
264
265 internal::ensure_sized(d_ipiv_, ipiv_bytes);
266
267 size_t dev_ws_bytes = 0, host_ws_bytes = 0;
268 EIGEN_CUSOLVER_CHECK(cusolverDnXgetrf_bufferSize(solver_ctx_.cusolverHandle(), solver_ctx_.params_.p, n_, n_, dtype,
269 d_lu_.get(), lda_, dtype, &dev_ws_bytes, &host_ws_bytes));
270
271 solver_ctx_.ensure_scratch(dev_ws_bytes);
272 solver_ctx_.h_workspace_.resize(host_ws_bytes);
273
274 EIGEN_CUSOLVER_CHECK(cusolverDnXgetrf(
275 solver_ctx_.cusolverHandle(), solver_ctx_.params_.p, n_, n_, dtype, d_lu_.get(), lda_,
276 static_cast<int64_t*>(d_ipiv_.get()), dtype, solver_ctx_.scratch_workspace(), dev_ws_bytes,
277 host_ws_bytes > 0 ? solver_ctx_.h_workspace_.data() : nullptr, host_ws_bytes, solver_ctx_.scratch_info()));
278
279 solver_ctx_.enqueue_info_copy();
280 }
281};
282} // namespace gpu
283} // namespace Eigen
284
285#endif // EIGEN_GPU_LU_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 LU decomposition with partial pivoting via cuSOLVER.
Definition GpuLU.h:37
LU & compute(const DenseBase< InputType > &A)
Definition GpuLU.h:104
LU & compute(DeviceMatrix< Scalar > &&d_A)
Definition GpuLU.h:138
LU & compute(const DeviceMatrix< Scalar > &d_A)
Definition GpuLU.h:122
PlainMatrix solve(const MatrixBase< Rhs > &B, GpuOp op=GpuOp::NoTrans) const
Definition GpuLU.h:157
LU(const DeviceMatrix< Scalar > &d_A)
Definition GpuLU.h:56
LU(Context &ctx)
Definition GpuLU.h:48
LU(Context &ctx, const DenseBase< InputType > &A)
Definition GpuLU.h:63
LU(DeviceMatrix< Scalar > &&d_A)
Definition GpuLU.h:59
LU(Context &ctx, DeviceMatrix< Scalar > &&d_A)
Definition GpuLU.h:71
LU(Context &ctx, const DeviceMatrix< Scalar > &d_A)
Definition GpuLU.h:68
DeviceMatrix< Scalar > solve(DeviceMatrix< Scalar > &&d_B, GpuOp op=GpuOp::NoTrans) const
Definition GpuLU.h:204
DeviceMatrix< Scalar > solve(const DeviceMatrix< Scalar > &d_B, GpuOp op=GpuOp::NoTrans) const
Definition GpuLU.h:189
Internal RAII owner for an untyped GPU device allocation.
Definition GpuSupport.h:293
ComputationInfo
Namespace containing all symbols from the Eigen library.