14#ifndef EIGEN_GPU_CONTEXT_H
15#define EIGEN_GPU_CONTEXT_H
18#include "./InternalHeaderCheck.h"
20#include "./DeviceScalarOps.h"
21#include "./CuBlasSupport.h"
22#include "./CuSolverSupport.h"
23#include "./CuSparseSupport.h"
35constexpr size_t kOneShotInfoBytes = 2 *
sizeof(int);
37constexpr size_t kOneShotHostInfoBytes = 0;
39constexpr size_t kOneShotHostInfoBytes = kOneShotInfoBytes;
46struct OneShotSolverScratch {
47 DeviceBuffer d_factor;
49 DeviceBuffer d_workspace;
50 DeviceBuffer d_info{kOneShotInfoBytes};
51 PinnedHostBuffer h_info{kOneShotHostInfoBytes};
52 std::vector<char> h_workspace;
55inline void ensure_sized(
DeviceBuffer& buf,
size_t needed) {
56 if (needed > buf.size()) {
85 cudaStream_t s =
nullptr;
86 EIGEN_CUDA_RUNTIME_CHECK(cudaStreamCreate(&s));
87 stream_ = internal::UniqueStream(s);
88 npp_stream_ctx_ = internal::make_npp_stream_ctx(stream_.get());
94 explicit Context(cudaStream_t stream) : stream_(stream, internal::CudaStreamDeleter{false}) {
95 npp_stream_ctx_ = internal::make_npp_stream_ctx(stream_.get());
122 Context*
override = tl_override_ptr();
123 if (
override)
return *
override;
133 cudaStream_t stream()
const {
return stream_.get(); }
134 cublasHandle_t cublasHandle()
const {
return cublas_.get(); }
144 cusolverDnHandle_t h =
nullptr;
145 EIGEN_CUSOLVER_CHECK(cusolverDnCreate(&h));
146 cusolver_ = LazyCusolverHandle(h, &destroyCusolver);
147 EIGEN_CUSOLVER_CHECK(cusolverDnSetStream(h, stream_.get()));
149 return cusolver_.get();
155 cublasLtHandle_t h =
nullptr;
156 EIGEN_CUBLAS_CHECK(cublasLtCreate(&h));
157 cublas_lt_ = internal::UniqueCublasLtHandle(h);
159 return cublas_lt_.get();
189 cusparseHandle_t h =
nullptr;
190 EIGEN_CUSPARSE_CHECK(cusparseCreate(&h));
191 cusparse_ = LazyCusparseHandle(h, &destroyCusparse);
192 EIGEN_CUSPARSE_CHECK(cusparseSetStream(h, stream_.get()));
194 return cusparse_.get();
198 static cusolverStatus_t destroyCusolver(cusolverDnHandle_t h) {
return cusolverDnDestroy(h); }
199 static cusparseStatus_t destroyCusparse(cusparseHandle_t h) {
return cusparseDestroy(h); }
202 using LazyCusolverHandle =
203 std::unique_ptr<std::remove_pointer_t<cusolverDnHandle_t>, cusolverStatus_t (*)(cusolverDnHandle_t)>;
204 using LazyCusparseHandle =
205 std::unique_ptr<std::remove_pointer_t<cusparseHandle_t>, cusparseStatus_t (*)(cusparseHandle_t)>;
208 internal::UniqueStream stream_;
209 NppStreamContext npp_stream_ctx_ = {};
210 internal::UniqueCublasHandle cublas_;
211 internal::DeviceBuffer cublas_workspace_;
212 LazyCusolverHandle cusolver_{
nullptr,
nullptr};
213 LazyCusparseHandle cusparse_{
nullptr,
nullptr};
214 internal::UniqueCublasLtHandle cublas_lt_;
215 internal::DeviceBuffer gemm_workspace_;
216 internal::CublasLtPlanCache gemm_plan_cache_{internal::kCublasLtPlanCacheCapacity};
217 internal::OneShotSolverScratch oneshot_solver_scratch_;
218 std::size_t cublaslt_max_workspace_bytes_ = internal::kCublasLtMaxWorkspaceBytes;
220 static Context*& tl_override_ptr() {
221 thread_local Context* ptr =
nullptr;
226 cublasHandle_t h =
nullptr;
227 EIGEN_CUBLAS_CHECK(cublasCreate(&h));
228 cublas_ = internal::UniqueCublasHandle(h);
229 EIGEN_CUBLAS_CHECK(cublasSetStream(h, stream_.get()));
230 cublas_workspace_ = internal::DeviceBuffer(internal::kCublasWorkspaceBytes);
231 EIGEN_CUBLAS_CHECK(cublasSetWorkspace(h, cublas_workspace_.get(), cublas_workspace_.size()));
Unified GPU execution context owning a CUDA stream and library handles.
Definition GpuContext.h:81
Context()
Definition GpuContext.h:84
cusolverDnHandle_t cusolverHandle()
Definition GpuContext.h:142
std::size_t cublasLtMaxWorkspaceBytes() const
Definition GpuContext.h:178
internal::CublasLtPlanCache & gemmPlanCache()
Definition GpuContext.h:168
cusparseHandle_t cusparseHandle()
Definition GpuContext.h:187
internal::OneShotSolverScratch & oneshotSolverScratch()
Definition GpuContext.h:173
cublasLtHandle_t cublasLtHandle()
Definition GpuContext.h:153
Context(cudaStream_t stream)
Definition GpuContext.h:94
internal::DeviceBuffer & gemmWorkspace()
Definition GpuContext.h:164
const NppStreamContext & nppStreamContext() const
Definition GpuContext.h:139
void setCublasLtMaxWorkspaceBytes(std::size_t bytes)
Definition GpuContext.h:184
static void setThreadLocal(Context *ctx)
Definition GpuContext.h:131
static Context & threadLocal()
Definition GpuContext.h:121
Internal RAII owner for an untyped GPU device allocation.
Definition GpuSupport.h:293
Namespace containing all symbols from the Eigen library.