43 using Scalar = Scalar_;
47 static constexpr int UpLo = UpLo_;
58 template <
typename InputType>
70 template <
typename InputType>
85 LLT& operator=(
const LLT&) =
delete;
89 : solver_ctx_(std::move(o.solver_ctx_)), d_factor_(std::move(o.d_factor_)), n_(o.n_), lda_(o.lda_) {
94 LLT& operator=(
LLT&& o)
noexcept {
96 solver_ctx_ = std::move(o.solver_ctx_);
97 d_factor_ = std::move(o.d_factor_);
108 template <
typename InputType>
110 eigen_assert(A.rows() == A.cols());
111 if (!begin_compute(A.rows()))
return *
this;
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()));
128 eigen_assert(d_A.rows() == d_A.cols());
129 if (!begin_compute(d_A.rows()))
return *
this;
131 lda_ =
static_cast<int64_t
>(d_A.rows());
133 allocate_factor_storage();
134 EIGEN_CUDA_RUNTIME_CHECK(
135 cudaMemcpyAsync(d_factor_.get(), d_A.data(), factorBytes(), cudaMemcpyDeviceToDevice, solver_ctx_.stream()));
145 eigen_assert(d_A.rows() == d_A.cols());
146 if (!begin_compute(d_A.rows()))
return *
this;
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());
157 template <
typename Rhs>
162 eigen_assert(solver_ctx_.info() ==
Success &&
"LLT::solve called on a failed or uninitialized factorization");
163 eigen_assert(B.rows() == n_);
166 const int64_t nrhs =
static_cast<int64_t
>(rhs.cols());
167 const int64_t ldb =
static_cast<int64_t
>(rhs.rows());
169 internal::upload_host_matrix(
static_cast<Scalar*
>(d_x.get()), ldb, rhs.data(), rhs.outerStride(), rhs.rows(),
170 rhs.cols(), solver_ctx_.stream());
173 PlainMatrix X(n_, B.cols());
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()));
181 eigen_assert(solve_info == 0 &&
"cusolverDnXpotrs reported an error");
191 eigen_assert(solver_ctx_.info() ==
Success &&
"LLT::solve called on a failed or uninitialized factorization");
192 eigen_assert(d_B.rows() == n_);
194 const int64_t nrhs =
static_cast<int64_t
>(d_B.cols());
195 const int64_t ldb =
static_cast<int64_t
>(d_B.rows());
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));
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));
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_factor_;
227 bool begin_compute(Index rows) {
229 return solver_ctx_.begin_compute(n_ != 0);
232 size_t factorBytes()
const {
return rhsBytes(n_, lda_); }
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);
238 void allocate_factor_storage() { internal::ensure_sized(d_factor_, factorBytes()); }
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;
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()));
251 DeviceMatrix<Scalar> result =
253 result.recordReady(solver_ctx_.stream());
258 constexpr cudaDataType_t dtype = internal::cusolver_data_type<Scalar>::value;
259 constexpr cublasFillMode_t uplo = internal::cusolver_fill_mode<UpLo_>::value;
261 solver_ctx_.mark_pending();
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,
268 solver_ctx_.ensure_scratch(dev_ws_bytes);
269 solver_ctx_.h_workspace_.resize(host_ws_bytes);
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()));
276 solver_ctx_.enqueue_info_copy();
Unified GPU execution context owning a CUDA stream and library handles.
Definition GpuContext.h:81