14#ifndef EIGEN_GPU_SOLVER_CONTEXT_H
15#define EIGEN_GPU_SOLVER_CONTEXT_H
18#include "./InternalHeaderCheck.h"
20#include "./CuSolverSupport.h"
21#include "./CuBlasSupport.h"
22#include "./GpuContext.h"
29struct GpuSolverContext {
30 Context* bound_ctx_ =
nullptr;
32 UniqueCusolverHandle cusolver_;
33 UniqueCublasHandle cublas_;
34 UniqueCublasLtHandle cublas_lt_;
35 CusolverParams params_;
36 DeviceBuffer d_scratch_;
37 std::vector<char> h_workspace_;
38 DeviceBuffer gemm_workspace_;
39 CublasLtPlanCache gemm_plan_cache_{kCublasLtPlanCacheCapacity};
42 std::size_t cublaslt_max_workspace_bytes_ = kCublasLtMaxWorkspaceBytes;
44 PinnedHostBuffer pinned_info_{
sizeof(int)};
45 bool info_synced_ =
true;
47 int& info_word() {
return *
static_cast<int*
>(pinned_info_.get()); }
48 int info_word()
const {
return *
static_cast<const int*
>(pinned_info_.get()); }
50 cudaStream_t stream()
const {
return stream_.get(); }
51 cusolverDnHandle_t cusolverHandle()
const {
return cusolver_.get(); }
52 cublasHandle_t cublasHandle()
const {
return cublas_.get(); }
55 cudaStream_t s =
nullptr;
56 EIGEN_CUDA_RUNTIME_CHECK(cudaStreamCreate(&s));
57 stream_ = UniqueStream(s);
58 cusolverDnHandle_t solver =
nullptr;
59 EIGEN_CUSOLVER_CHECK(cusolverDnCreate(&solver));
60 cusolver_ = UniqueCusolverHandle(solver);
61 EIGEN_CUSOLVER_CHECK(cusolverDnSetStream(solver, s));
62 cublasHandle_t blas =
nullptr;
63 EIGEN_CUBLAS_CHECK(cublasCreate(&blas));
64 cublas_ = UniqueCublasHandle(blas);
65 EIGEN_CUBLAS_CHECK(cublasSetStream(blas, s));
74 explicit GpuSolverContext(Context& ctx)
76 stream_(ctx.stream(), CudaStreamDeleter{false}),
77 cusolver_(ctx.cusolverHandle(), CusolverHandleDeleter{false}),
78 cublas_(ctx.cublasHandle(), CublasHandleDeleter{false}) {
82 ~GpuSolverContext() =
default;
83 GpuSolverContext(GpuSolverContext&& o)
noexcept =
default;
85 GpuSolverContext& operator=(GpuSolverContext&& o)
noexcept {
88 if (!info_synced_ && pinned_info_) (void)cudaStreamSynchronize(stream());
90 gemm_plan_cache_.clear();
91 bound_ctx_ = o.bound_ctx_;
92 stream_ = std::move(o.stream_);
93 cusolver_ = std::move(o.cusolver_);
94 cublas_ = std::move(o.cublas_);
95 cublas_lt_ = std::move(o.cublas_lt_);
96 params_ = std::move(o.params_);
97 d_scratch_ = std::move(o.d_scratch_);
98 h_workspace_ = std::move(o.h_workspace_);
99 gemm_workspace_ = std::move(o.gemm_workspace_);
100 gemm_plan_cache_ = std::move(o.gemm_plan_cache_);
101 cublaslt_max_workspace_bytes_ = o.cublaslt_max_workspace_bytes_;
103 pinned_info_ = std::move(o.pinned_info_);
104 info_synced_ = o.info_synced_;
105 o.bound_ctx_ =
nullptr;
112 cublasLtHandle_t cublasLtHandle() {
113 if (bound_ctx_)
return bound_ctx_->cublasLtHandle();
115 cublasLtHandle_t h =
nullptr;
116 EIGEN_CUBLAS_CHECK(cublasLtCreate(&h));
117 cublas_lt_ = UniqueCublasLtHandle(h);
119 return cublas_lt_.get();
124 CublasLtPlanCache& gemmPlanCache() {
return bound_ctx_ ? bound_ctx_->gemmPlanCache() : gemm_plan_cache_; }
125 DeviceBuffer& gemmWorkspace() {
return bound_ctx_ ? bound_ctx_->gemmWorkspace() : gemm_workspace_; }
126 std::size_t cublasLtMaxWorkspaceBytes()
const {
127 return bound_ctx_ ? bound_ctx_->cublasLtMaxWorkspaceBytes() : cublaslt_max_workspace_bytes_;
130 GpuSolverContext(
const GpuSolverContext&) =
delete;
131 GpuSolverContext& operator=(
const GpuSolverContext&) =
delete;
135 static constexpr size_t kInfoBytes =
sizeof(int);
136 static constexpr size_t kScratchAlign = 16;
138 static size_t scratchBytesFor(
size_t workspace_bytes) {
139 workspace_bytes = (workspace_bytes + kScratchAlign - 1) & ~(kScratchAlign - 1);
140 return workspace_bytes + kInfoBytes;
146 void ensure_scratch(
size_t workspace_bytes) {
147 size_t needed = scratchBytesFor(workspace_bytes);
148 if (needed > d_scratch_.size()) {
149 if (d_scratch_) EIGEN_CUDA_RUNTIME_CHECK(cudaStreamSynchronize(stream()));
150 d_scratch_ = DeviceBuffer(needed);
154 void* scratch_workspace()
const {
return d_scratch_.get(); }
156 int* scratch_info()
const {
157 eigen_assert(d_scratch_ && d_scratch_.size() >= kInfoBytes);
158 return reinterpret_cast<int*
>(
static_cast<char*
>(d_scratch_.get()) + d_scratch_.size() - kInfoBytes);
162 void mark_pending() {
163 info_synced_ =
false;
169 bool begin_compute(
bool nonempty) {
181 void enqueue_info_copy() {
182 EIGEN_CUDA_RUNTIME_CHECK(
183 cudaMemcpyAsync(&info_word(), scratch_info(),
sizeof(
int), cudaMemcpyDeviceToHost, stream()));
189 EIGEN_CUDA_RUNTIME_CHECK(cudaStreamSynchronize(stream()));
202 void download(
void* dst,
const void* src,
size_t bytes)
const {
203 if (bytes == 0)
return;
204 EIGEN_CUDA_RUNTIME_CHECK(cudaMemcpyAsync(dst, src, bytes, cudaMemcpyDeviceToHost, stream()));
205 EIGEN_CUDA_RUNTIME_CHECK(cudaStreamSynchronize(stream()));
Namespace containing all symbols from the Eigen library.