Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
GpuContext.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2026 Rasmus Munk Larsen <rmlarsen@gmail.com>
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// Unified GPU execution context: a CUDA stream plus the NVIDIA library handles
12// used by gpu::DeviceMatrix operations.
13
14#ifndef EIGEN_GPU_CONTEXT_H
15#define EIGEN_GPU_CONTEXT_H
16
17// IWYU pragma: private
18#include "./InternalHeaderCheck.h"
19
20#include "./DeviceScalarOps.h"
21#include "./CuBlasSupport.h"
22#include "./CuSolverSupport.h"
23#include "./CuSparseSupport.h"
24#include <vector>
25
26namespace Eigen {
27namespace gpu {
28
29namespace internal {
30
31// cuSOLVER writes the factorization/solve status words to device memory in
32// every build, so d_info must exist even under EIGEN_NO_DEBUG. Only the
33// pinned host mirror is debug-only: it feeds the oneshot_check_info assert,
34// which release builds compile out.
35constexpr size_t kOneShotInfoBytes = 2 * sizeof(int);
36#ifdef EIGEN_NO_DEBUG
37constexpr size_t kOneShotHostInfoBytes = 0;
38#else
39constexpr size_t kOneShotHostInfoBytes = kOneShotInfoBytes;
40#endif
41// Grow-only scratch shared by the one-shot solver expressions
42// (d_A.llt().solve(d_B), d_A.lu().solve(d_B)), so repeated one-shot solves on
43// a Context perform no per-call device or pinned-host allocations. Holds only
44// CUDA-runtime types (no cuSOLVER types) to keep the lazy-linking property of
45// Context. Used by the one-shot solve dispatches in DeviceDispatch.h.
46struct OneShotSolverScratch {
47 DeviceBuffer d_factor;
48 DeviceBuffer d_ipiv;
49 DeviceBuffer d_workspace;
50 DeviceBuffer d_info{kOneShotInfoBytes}; // 2 ints: {factorization, solve}
51 PinnedHostBuffer h_info{kOneShotHostInfoBytes}; // debug-build info check only
52 std::vector<char> h_workspace;
53};
54
55inline void ensure_sized(DeviceBuffer& buf, size_t needed) {
56 if (needed > buf.size()) {
57 // Replacing an in-use buffer is safe: device_free is stream-ordered (or
58 // fully synchronous on the cudaMalloc fallback path) and DeviceBufferPool
59 // holds a released block back until the device has retired the work
60 // enqueued before the release, so nothing reuses the old buffer early.
61 buf = DeviceBuffer(needed);
62 }
63}
64} // namespace internal
65
81class Context {
82 public:
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());
89 init_cublas();
90 }
91
94 explicit Context(cudaStream_t stream) : stream_(stream, internal::CudaStreamDeleter{/*owns=*/false}) {
95 npp_stream_ctx_ = internal::make_npp_stream_ctx(stream_.get());
96 init_cublas();
97 }
98
99 ~Context() = default;
100
101 Context(const Context&) = delete;
102 Context& operator=(const Context&) = delete;
103 Context(Context&&) = delete;
104 Context& operator=(Context&&) = delete;
105
122 Context* override = tl_override_ptr();
123 if (override) return *override;
124 thread_local Context ctx;
125 return ctx;
126 }
127
131 static void setThreadLocal(Context* ctx) { tl_override_ptr() = ctx; }
132
133 cudaStream_t stream() const { return stream_.get(); }
134 cublasHandle_t cublasHandle() const { return cublas_.get(); }
135
139 const NppStreamContext& nppStreamContext() const { return npp_stream_ctx_; }
140
142 cusolverDnHandle_t cusolverHandle() {
143 if (!cusolver_) {
144 cusolverDnHandle_t h = nullptr;
145 EIGEN_CUSOLVER_CHECK(cusolverDnCreate(&h));
146 cusolver_ = LazyCusolverHandle(h, &destroyCusolver);
147 EIGEN_CUSOLVER_CHECK(cusolverDnSetStream(h, stream_.get()));
148 }
149 return cusolver_.get();
150 }
151
153 cublasLtHandle_t cublasLtHandle() {
154 if (!cublas_lt_) {
155 cublasLtHandle_t h = nullptr;
156 EIGEN_CUBLAS_CHECK(cublasLtCreate(&h));
157 cublas_lt_ = internal::UniqueCublasLtHandle(h);
158 }
159 return cublas_lt_.get();
160 }
161
164 internal::DeviceBuffer& gemmWorkspace() { return gemm_workspace_; }
165
168 internal::CublasLtPlanCache& gemmPlanCache() { return gemm_plan_cache_; }
169
173 internal::OneShotSolverScratch& oneshotSolverScratch() { return oneshot_solver_scratch_; }
174
178 std::size_t cublasLtMaxWorkspaceBytes() const { return cublaslt_max_workspace_bytes_; }
179
184 void setCublasLtMaxWorkspaceBytes(std::size_t bytes) { cublaslt_max_workspace_bytes_ = bytes; }
185
187 cusparseHandle_t cusparseHandle() {
188 if (!cusparse_) {
189 cusparseHandle_t h = nullptr;
190 EIGEN_CUSPARSE_CHECK(cusparseCreate(&h));
191 cusparse_ = LazyCusparseHandle(h, &destroyCusparse);
192 EIGEN_CUSPARSE_CHECK(cusparseSetStream(h, stream_.get()));
193 }
194 return cusparse_.get();
195 }
196
197 private:
198 static cusolverStatus_t destroyCusolver(cusolverDnHandle_t h) { return cusolverDnDestroy(h); }
199 static cusparseStatus_t destroyCusparse(cusparseHandle_t h) { return cusparseDestroy(h); }
200
201 // Function-pointer deleters keep cusolverDnDestroy / cusparseDestroy referenced only by TUs that create handles.
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)>;
206
207 // Destroyed in reverse declaration order: the plan cache before the cuBLASLt handle, the stream last.
208 internal::UniqueStream stream_;
209 NppStreamContext npp_stream_ctx_ = {};
210 internal::UniqueCublasHandle cublas_;
211 internal::DeviceBuffer cublas_workspace_; // cublasSetWorkspace; freed before the handle
212 LazyCusolverHandle cusolver_{nullptr, nullptr};
213 LazyCusparseHandle cusparse_{nullptr, nullptr};
214 internal::UniqueCublasLtHandle cublas_lt_; // lazy
215 internal::DeviceBuffer gemm_workspace_; // lazy
216 internal::CublasLtPlanCache gemm_plan_cache_{internal::kCublasLtPlanCacheCapacity};
217 internal::OneShotSolverScratch oneshot_solver_scratch_; // grow-only
218 std::size_t cublaslt_max_workspace_bytes_ = internal::kCublasLtMaxWorkspaceBytes;
219
220 static Context*& tl_override_ptr() {
221 thread_local Context* ptr = nullptr;
222 return ptr;
223 }
224
225 void init_cublas() {
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()));
232 }
233};
234
235} // namespace gpu
236} // namespace Eigen
237
238#endif // EIGEN_GPU_CONTEXT_H
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.