Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
GpuSolverContext.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// Shared context for the dense GPU solvers. Each solver holds one by composition
12// and delegates handle lifetime and scratch management to it.
13
14#ifndef EIGEN_GPU_SOLVER_CONTEXT_H
15#define EIGEN_GPU_SOLVER_CONTEXT_H
16
17// IWYU pragma: private
18#include "./InternalHeaderCheck.h"
19
20#include "./CuSolverSupport.h"
21#include "./CuBlasSupport.h"
22#include "./GpuContext.h"
23#include <vector>
24
25namespace Eigen {
26namespace gpu {
27namespace internal {
28
29struct GpuSolverContext {
30 Context* bound_ctx_ = nullptr;
31 UniqueStream stream_;
32 UniqueCusolverHandle cusolver_;
33 UniqueCublasHandle cublas_;
34 UniqueCublasLtHandle cublas_lt_; // lazy: created on first GEMM-via-cublasLt call (standalone mode only)
35 CusolverParams params_;
36 DeviceBuffer d_scratch_;
37 std::vector<char> h_workspace_;
38 DeviceBuffer gemm_workspace_; // grown lazily by cublaslt_gemm
39 CublasLtPlanCache gemm_plan_cache_{kCublasLtPlanCacheCapacity};
40 // Workspace ceiling fed to the cublasLtMatmul heuristic at plan-creation time.
41 // See gpu::Context::setCublasLtMaxWorkspaceBytes() for semantics.
42 std::size_t cublaslt_max_workspace_bytes_ = kCublasLtMaxWorkspaceBytes;
44 PinnedHostBuffer pinned_info_{sizeof(int)}; // pinned host memory for async D2H of info word
45 bool info_synced_ = true;
46
47 int& info_word() { return *static_cast<int*>(pinned_info_.get()); }
48 int info_word() const { return *static_cast<const int*>(pinned_info_.get()); }
49
50 cudaStream_t stream() const { return stream_.get(); }
51 cusolverDnHandle_t cusolverHandle() const { return cusolver_.get(); }
52 cublasHandle_t cublasHandle() const { return cublas_.get(); }
53
54 GpuSolverContext() {
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));
66 ensure_scratch(0);
67 }
68
74 explicit GpuSolverContext(Context& ctx)
75 : bound_ctx_(&ctx),
76 stream_(ctx.stream(), CudaStreamDeleter{/*owns=*/false}),
77 cusolver_(ctx.cusolverHandle(), CusolverHandleDeleter{/*owns=*/false}),
78 cublas_(ctx.cublasHandle(), CublasHandleDeleter{/*owns=*/false}) {
79 ensure_scratch(0);
80 }
81
82 ~GpuSolverContext() = default;
83 GpuSolverContext(GpuSolverContext&& o) noexcept = default;
84
85 GpuSolverContext& operator=(GpuSolverContext&& o) noexcept {
86 if (this != &o) {
87 // A pending info copy may still write pinned_info_, whose cudaFreeHost deleter is not stream-ordered.
88 if (!info_synced_ && pinned_info_) (void)cudaStreamSynchronize(stream());
89 // Release plan-cache descriptors before the moves below replace the cuBLASLt handle they were built with.
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_;
102 info_ = o.info_;
103 pinned_info_ = std::move(o.pinned_info_);
104 info_synced_ = o.info_synced_;
105 o.bound_ctx_ = nullptr;
106 }
107 return *this;
108 }
109
112 cublasLtHandle_t cublasLtHandle() {
113 if (bound_ctx_) return bound_ctx_->cublasLtHandle();
114 if (!cublas_lt_) {
115 cublasLtHandle_t h = nullptr;
116 EIGEN_CUBLAS_CHECK(cublasLtCreate(&h));
117 cublas_lt_ = UniqueCublasLtHandle(h);
118 }
119 return cublas_lt_.get();
120 }
121
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_;
128 }
129
130 GpuSolverContext(const GpuSolverContext&) = delete;
131 GpuSolverContext& operator=(const GpuSolverContext&) = delete;
132
133 // Scratch layout: [ workspace (aligned) | info_word (sizeof(int)) ].
134 // Workspace size is rounded up to 16 bytes so the info word lands aligned.
135 static constexpr size_t kInfoBytes = sizeof(int);
136 static constexpr size_t kScratchAlign = 16;
137
138 static size_t scratchBytesFor(size_t workspace_bytes) {
139 workspace_bytes = (workspace_bytes + kScratchAlign - 1) & ~(kScratchAlign - 1);
140 return workspace_bytes + kInfoBytes;
141 }
142
143 // Ensure d_scratch_ holds at least `workspace_bytes` of scratch plus the trailing
144 // info word. Grows but never shrinks. Syncs the stream before reallocating to
145 // avoid freeing memory that async kernels may still be using.
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);
151 }
152 }
153
154 void* scratch_workspace() const { return d_scratch_.get(); }
155
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);
159 }
160
161 // Mark a factorization as pending: its info word is not yet available.
162 void mark_pending() {
163 info_synced_ = false;
164 info_ = InvalidInput;
165 }
166
167 // Common compute() prologue: reset info state. Returns false for the empty
168 // (n == 0) case, which is trivially successful — the caller returns early.
169 bool begin_compute(bool nonempty) {
170 info_ = InvalidInput;
171 if (!nonempty) {
172 info_ = Success;
173 info_synced_ = true;
174 return false;
175 }
176 return true;
177 }
178
179 // Common factorize() epilogue: enqueue the async D2H copy of the info word
180 // into pinned host memory. Read later by the lazy sync_info().
181 void enqueue_info_copy() {
182 EIGEN_CUDA_RUNTIME_CHECK(
183 cudaMemcpyAsync(&info_word(), scratch_info(), sizeof(int), cudaMemcpyDeviceToHost, stream()));
184 }
185
186 // Synchronize the stream and interpret the info word; no-op once synced.
187 void sync_info() {
188 if (!info_synced_) {
189 EIGEN_CUDA_RUNTIME_CHECK(cudaStreamSynchronize(stream()));
190 info_ = (info_word() == 0) ? Success : NumericalIssue;
191 info_synced_ = true;
192 }
193 }
194
195 ComputationInfo info() {
196 sync_info();
197 return info_;
198 }
199
200 // Blocking download of solver-owned device data. It waits only for the solver's stream; cudaMemcpy would run on
201 // the legacy default stream and wait for every blocking stream on the device.
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()));
206 }
207};
208
209} // namespace internal
210} // namespace gpu
211} // namespace Eigen
212
213#endif // EIGEN_GPU_SOLVER_CONTEXT_H
ComputationInfo
NumericalIssue
Namespace containing all symbols from the Eigen library.