39 using Scalar = Scalar_;
50 template <
typename InputType>
62 template <
typename InputType>
75 LU(
const LU&) =
delete;
76 LU& operator=(
const LU&) =
delete;
79 : solver_ctx_(std::move(o.solver_ctx_)),
80 d_lu_(std::move(o.d_lu_)),
81 d_ipiv_(std::move(o.d_ipiv_)),
88 LU& operator=(LU&& o)
noexcept {
90 solver_ctx_ = std::move(o.solver_ctx_);
91 d_lu_ = std::move(o.d_lu_);
92 d_ipiv_ = std::move(o.d_ipiv_);
103 template <
typename InputType>
105 eigen_assert(A.rows() == A.cols() &&
"LU requires a square matrix");
106 if (!begin_compute(A.rows()))
return *
this;
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()));
123 eigen_assert(d_A.rows() == d_A.cols() &&
"LU requires a square matrix");
124 if (!begin_compute(d_A.rows()))
return *
this;
126 lda_ =
static_cast<int64_t
>(d_A.rows());
128 allocate_lu_storage();
129 EIGEN_CUDA_RUNTIME_CHECK(
130 cudaMemcpyAsync(d_lu_.get(), d_A.data(), matrixBytes(), cudaMemcpyDeviceToDevice, solver_ctx_.stream()));
140 eigen_assert(d_A.rows() == d_A.cols() &&
"LU requires a square matrix");
141 if (!begin_compute(d_A.rows()))
return *
this;
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());
156 template <
typename Rhs>
161 eigen_assert(solver_ctx_.info() ==
Success &&
"LU::solve called on a failed or uninitialized factorization");
162 eigen_assert(B.rows() == n_);
165 const int64_t nrhs =
static_cast<int64_t
>(rhs.cols());
166 const int64_t ldb =
static_cast<int64_t
>(rhs.rows());
168 internal::upload_host_matrix(
static_cast<Scalar*
>(d_x.get()), ldb, rhs.data(), rhs.outerStride(), rhs.rows(),
169 rhs.cols(), solver_ctx_.stream());
172 PlainMatrix X(n_, B.cols());
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()));
180 eigen_assert(solve_info == 0 &&
"cusolverDnXgetrs reported an error");
190 eigen_assert(solver_ctx_.info() ==
Success &&
"LU::solve called on a failed or uninitialized factorization");
191 eigen_assert(d_B.rows() == n_);
193 const int64_t nrhs =
static_cast<int64_t
>(d_B.cols());
194 const int64_t ldb =
static_cast<int64_t
>(d_B.rows());
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));
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));
217 Index rows()
const {
return n_; }
218 Index cols()
const {
return n_; }
219 cudaStream_t stream()
const {
return solver_ctx_.stream(); }
222 mutable internal::GpuSolverContext solver_ctx_;
223 internal::DeviceBuffer d_lu_;
224 internal::DeviceBuffer d_ipiv_;
228 bool begin_compute(Index rows) {
230 return solver_ctx_.begin_compute(n_ != 0);
233 size_t matrixBytes()
const {
return matrixBytes(n_, lda_); }
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);
239 void allocate_lu_storage() { internal::ensure_sized(d_lu_, matrixBytes()); }
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);
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()));
253 DeviceMatrix<Scalar> result =
255 result.recordReady(solver_ctx_.stream());
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);
263 solver_ctx_.mark_pending();
265 internal::ensure_sized(d_ipiv_, ipiv_bytes);
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));
271 solver_ctx_.ensure_scratch(dev_ws_bytes);
272 solver_ctx_.h_workspace_.resize(host_ws_bytes);
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()));
279 solver_ctx_.enqueue_info_copy();
Unified GPU execution context owning a CUDA stream and library handles.
Definition GpuContext.h:81