Skip to main content

tenferro_gpu/cubecl/
runtime.rs

1//! CubeCL CUDA runtime initialization and synchronization.
2
3use std::fmt;
4use std::hash::{Hash, Hasher};
5use std::io::Write;
6use std::sync::{Arc, Mutex, OnceLock};
7
8use cubecl::client::ComputeClient;
9use cubecl::stream_id::StreamId;
10use cubecl::Runtime;
11use cubecl_cuda::{CudaDevice, CudaRuntime as CubeclCudaRuntime};
12use cubecl_runtime::config::{CubeClRuntimeConfig, RuntimeConfig};
13use cudarc::cublas::sys as cublas_sys;
14use cudarc::driver::result::DriverError;
15use cudarc::driver::sys::{CUcontext, CUdevice, CUresult};
16use cudarc::runtime::{result as cuda_result, sys as cuda_sys, sys::cudaStream_t};
17use tenferro_tensor::AllocationDomainId;
18
19use super::device::{
20    cuda_devices, unavailable_device_error, CudaDeviceError, CudaDeviceId, CudaDeviceInfo,
21};
22use super::identity::GpuExtensionCapability;
23
24/// Returns `true` if a CUDA device can initialize a CubeCL runtime.
25///
26/// Use this in test helpers to skip GPU tests on machines without hardware.
27pub fn gpu_available() -> bool {
28    let library_present = std::panic::catch_unwind(|| {
29        // SAFETY: `is_culib_present` only probes candidate library names and
30        // does not call CUDA function pointers or retain a library handle.
31        unsafe { cudarc::driver::sys::is_culib_present() }
32    })
33    .unwrap_or(false);
34    if !library_present {
35        return false;
36    }
37    let Ok(devices) = cuda_devices() else {
38        return false;
39    };
40    let Some(device_id) = devices.first().map(|device| device.id()) else {
41        return false;
42    };
43    std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
44        let Ok(runtime) = CudaRuntime::new(device_id) else {
45            return false;
46        };
47        runtime.synchronize().is_ok()
48    }))
49    .unwrap_or(false)
50}
51
52/// RAII guard that attempts to restore the thread's previous CUDA device and
53/// current context when dropped.
54///
55/// Used by the `with_raw` enter/exit protocol: the guard is created after the
56/// calling thread's device/context are saved and the tenferro primary context
57/// is activated. Drop attempts best-effort restoration of the saved state on
58/// normal return, `Err`, and unwind; a restoration failure is logged to
59/// stderr (non-panicking) and never returned.
60pub(crate) struct RawContextRestore {
61    saved_device: Result<i32, cudarc::runtime::result::RuntimeError>,
62    saved_context: Result<Option<CUcontext>, cudarc::driver::result::DriverError>,
63    op: &'static str,
64}
65
66impl RawContextRestore {
67    /// Save the current device/context, then activate `device`/`context`.
68    pub(crate) fn enter(op: &'static str, device: i32, context: CUcontext) -> crate::Result<Self> {
69        let saved_device = cudarc::runtime::result::device::get();
70        let saved_context = cudarc::driver::result::ctx::get_current();
71        cudarc::runtime::result::device::set(device)
72            .map_err(|err| crate::Error::backend_source(op, err))?;
73        if let Err(err) = unsafe { cudarc::driver::result::ctx::set_current(context) } {
74            // Roll the device and context back so a partial activation failure
75            // cannot leave the caller's thread on a different device or with a
76            // different current context (setting the device can implicitly
77            // change the thread's current context to the new primary).
78            if let Ok(previous_device) = saved_device {
79                let _ = cudarc::runtime::result::device::set(previous_device);
80            }
81            match saved_context {
82                Ok(Some(previous)) => {
83                    let _ = unsafe { cudarc::driver::result::ctx::set_current(previous) };
84                }
85                Ok(None) => {
86                    let _ =
87                        unsafe { cudarc::driver::result::ctx::set_current(std::ptr::null_mut()) };
88                }
89                Err(_) => {}
90            }
91            return Err(crate::Error::backend_source(op, err));
92        }
93        Ok(Self {
94            saved_device,
95            saved_context,
96            op,
97        })
98    }
99
100    fn restore(&self) {
101        let mut stderr = std::io::stderr();
102        if let Ok(device) = self.saved_device {
103            if let Err(err) = cudarc::runtime::result::device::set(device) {
104                let _ = writeln!(
105                    stderr,
106                    "tenferro-gpu: failed to restore CUDA device during {}: {err:?}",
107                    self.op
108                );
109            }
110        }
111        match self.saved_context {
112            Ok(Some(context)) => {
113                if let Err(err) = unsafe { cudarc::driver::result::ctx::set_current(context) } {
114                    let _ = writeln!(
115                        stderr,
116                        "tenferro-gpu: failed to restore CUDA context during {}: {err:?}",
117                        self.op
118                    );
119                }
120            }
121            // The thread had no current context before the guard; restore that
122            // state instead of leaving the tenferro primary context current.
123            Ok(None) => {
124                if let Err(err) =
125                    unsafe { cudarc::driver::result::ctx::set_current(std::ptr::null_mut()) }
126                {
127                    let _ = writeln!(
128                        stderr,
129                        "tenferro-gpu: failed to clear CUDA context during {}: {err:?}",
130                        self.op
131                    );
132                }
133            }
134            // The saved-context query itself failed; nothing can be restored.
135            Err(_) => {}
136        }
137    }
138}
139
140impl Drop for RawContextRestore {
141    fn drop(&mut self) {
142        self.restore();
143    }
144}
145
146/// Opaque identity of one exact CUDA runtime instance.
147///
148/// Cloning the identity preserves the underlying executable runtime witness;
149/// constructing another runtime, even for the same device ordinal, produces a
150/// distinct identity. The cache key intentionally carries no provider or
151/// device identifier and grants no execution authority.
152#[derive(Clone, Debug)]
153pub struct CudaRuntimeIdentity {
154    marker: Arc<u8>,
155}
156
157impl CudaRuntimeIdentity {
158    fn fresh() -> Self {
159        Self {
160            marker: Arc::new(0),
161        }
162    }
163}
164
165impl PartialEq for CudaRuntimeIdentity {
166    fn eq(&self, other: &Self) -> bool {
167        Arc::ptr_eq(&self.marker, &other.marker)
168    }
169}
170
171impl Eq for CudaRuntimeIdentity {}
172
173impl Hash for CudaRuntimeIdentity {
174    fn hash<H: Hasher>(&self, state: &mut H) {
175        // INVARIANT: `marker` is retained by every clone of this identity, so
176        // its Arc allocation address is move/clone-invariant while witnessed.
177        state.write_usize(Arc::as_ptr(&self.marker) as usize);
178    }
179}
180
181/// CubeCL CUDA runtime wrapper.
182///
183/// # Examples
184///
185/// ```
186/// use tenferro_gpu::cuda::CudaRuntime;
187///
188/// let _ctor: fn(tenferro_gpu::cuda::CudaDeviceId) ->
189///     Result<CudaRuntime, tenferro_gpu::cuda::CudaDeviceError> = CudaRuntime::new;
190/// let _sync: fn(&CudaRuntime) -> tenferro_tensor::Result<()> =
191///     CudaRuntime::synchronize;
192/// ```
193#[derive(Clone)]
194pub struct CudaRuntime {
195    inner: Arc<CudaRuntimeState>,
196}
197
198pub(crate) struct CudaRuntimeState {
199    client: ComputeClient<CubeclCudaRuntime>,
200    device_id: CudaDeviceId,
201    device_ordinal: usize,
202    device_info: CudaDeviceInfo,
203    primary_context: CudaPrimaryContext,
204    identity: CudaRuntimeIdentity,
205    allocation_domain: AllocationDomainId,
206    // Memoized raw CUDA stream handles keyed by the bounded CubeCL stream-pool
207    // slot, not the process-global and monotonically increasing `StreamId`.
208    //
209    // INVARIANT: in pinned CubeCL rev a2adda17, the CUDA server maps each
210    // `StreamId` to a fixed `StreamPool` slot whose `CUstream` is created once
211    // and never destroyed or replaced while the server is alive, and the
212    // server outlives the `ComputeClient` clone owned by this state. The table
213    // is owned by this runtime object (not thread-local/global), holds one
214    // entry per CubeCL stream slot, and is dropped with the runtime.
215    raw_streams: Box<[OnceLock<u64>]>,
216    // One lazily created cuBLAS handle per bounded CubeCL stream-pool slot.
217    // Each slot lock covers pointer-mode selection and the enqueue itself:
218    // distinct `StreamId`s can map to the same physical stream, and cuBLAS
219    // handle configuration is mutable.
220    cublas_handles: Box<[Mutex<Option<CublasStreamHandle>>]>,
221    // Lazily allocated pinned-host staging slot for single-scalar downloads
222    // (`cudaHostAlloc`, `PINNED_SCALAR_BYTES` bytes). Freed in `Drop` with
223    // `cudaFreeHost` while the primary context is still retained.
224    pinned_scalar: Mutex<PinnedScalarSlot>,
225    // Vendor workspaces whose CubeCL handle may only return to the pool after
226    // their stream reaches the event recorded at retirement time.
227    workspace_retirements: Mutex<super::workspace_retirement::WorkspaceRetirementQueue>,
228}
229
230/// One cached cuBLAS handle bound to a fixed CUDA stream.
231struct CublasStreamHandle(cublas_sys::cublasHandle_t);
232
233/// Pinned-host staging slot; `ptr` is null until the first scalar download.
234struct PinnedScalarSlot {
235    ptr: *mut std::ffi::c_void,
236}
237
238/// Size of the runtime-owned pinned staging slot: large enough for the widest
239/// supported scalar (`Complex64`, 16 bytes).
240pub(crate) const PINNED_SCALAR_BYTES: usize = 16;
241
242// SAFETY: `CudaRuntimeState` owns a retained CUDA primary context and a CubeCL
243// client for one device ordinal. Methods set the context current before raw CUDA
244// calls, and backend/executor layers serialize mutating tensor execution. The
245// raw cuBLAS handles and the pinned staging pointer are plain CUDA resource
246// addresses owned by this state and released in `Drop`.
247unsafe impl Send for CudaRuntimeState {}
248// SAFETY: Shared state access exposes immutable runtime handles; synchronization
249// and stream queries use explicit CUDA/CubeCL handles and do not mutate Rust
250// aliasing-visible fields. cuBLAS handles are locked per bounded physical
251// stream slot, and the pinned staging slot is used only while its mutex is held.
252unsafe impl Sync for CudaRuntimeState {}
253
254impl fmt::Debug for CudaRuntime {
255    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
256        f.debug_struct("CudaRuntime")
257            .field("device_id", &self.inner.device_id)
258            .finish_non_exhaustive()
259    }
260}
261
262struct CudaPrimaryContext {
263    cuda_device: CUdevice,
264    cuda_context: CUcontext,
265}
266
267impl CudaPrimaryContext {
268    fn retain(cuda_device: CUdevice) -> crate::Result<Self> {
269        let cuda_context = unsafe { cudarc::driver::result::primary_ctx::retain(cuda_device) }
270            .map_err(|err| crate::Error::backend_source("cubecl_runtime_init", err))?;
271        Ok(Self {
272            cuda_device,
273            cuda_context,
274        })
275    }
276
277    fn context(&self) -> CUcontext {
278        self.cuda_context
279    }
280}
281
282impl Drop for CudaPrimaryContext {
283    fn drop(&mut self) {
284        if let Err(err) = unsafe { cudarc::driver::result::primary_ctx::release(self.cuda_device) }
285        {
286            report_cuda_primary_context_release_error(&err);
287        }
288    }
289}
290
291#[cold]
292fn report_cuda_primary_context_release_error(err: &impl fmt::Debug) {
293    eprintln!("tenferro-gpu: failed to release CUDA primary context during Drop: {err:?}");
294}
295
296#[cold]
297fn report_cuda_runtime_drop_error(err: &crate::Error) {
298    eprintln!("tenferro-gpu: failed to synchronize CUDA runtime during Drop: {err}");
299}
300
301impl CudaRuntime {
302    /// Initialize the CubeCL CUDA runtime on the caller-selected device.
303    ///
304    /// # Examples
305    ///
306    /// ```
307    /// use tenferro_gpu::{cuda::CudaDeviceError, cuda::CudaDeviceId, cuda::CudaRuntime};
308    ///
309    /// let _ctor: fn(CudaDeviceId) -> Result<CudaRuntime, CudaDeviceError> = CudaRuntime::new;
310    /// ```
311    ///
312    /// # Errors
313    ///
314    /// Returns [`CudaDeviceError::Discovery`] when fallback discovery for an
315    /// invalid selected ordinal fails, [`CudaDeviceError::Unavailable`] when
316    /// that ordinal is not available, or [`CudaDeviceError::Initialization`]
317    /// when CUDA driver, runtime, context, or CubeCL client initialization
318    /// fails.
319    pub fn new(device_id: CudaDeviceId) -> Result<Self, CudaDeviceError> {
320        let device_ordinal = usize::try_from(device_id.ordinal()).map_err(|source| {
321            cuda_initialization_error(device_id, "convert_device_ordinal", source)
322        })?;
323        let cuda_ordinal = i32::try_from(device_id.ordinal()).map_err(|source| {
324            cuda_initialization_error(device_id, "convert_cuda_ordinal", source)
325        })?;
326        cudarc::driver::result::init()
327            .map_err(|source| cuda_initialization_error(device_id, "initialize_driver", source))?;
328        let cuda_device = match cudarc::driver::result::device::get(cuda_ordinal) {
329            Ok(cuda_device) => cuda_device,
330            Err(source) if is_invalid_device_lookup(source) => {
331                return Err(unavailable_device_error(device_id, cuda_devices()?));
332            }
333            Err(source) => {
334                return Err(cuda_initialization_error(device_id, "get_device", source));
335            }
336        };
337        let primary_context = CudaPrimaryContext::retain(cuda_device).map_err(|source| {
338            cuda_initialization_error(device_id, "retain_primary_context", source)
339        })?;
340        unsafe { cudarc::driver::result::ctx::set_current(primary_context.context()) }.map_err(
341            |source| cuda_initialization_error(device_id, "set_current_context", source),
342        )?;
343        cudarc::runtime::result::device::set(cuda_ordinal)
344            .map_err(|source| cuda_initialization_error(device_id, "set_device", source))?;
345        let device = CudaDevice::new(device_ordinal);
346        let client = CubeclCudaRuntime::client(&device);
347        let discovered = cuda_devices()?;
348        let device_info = discovered
349            .iter()
350            .find(|info| info.id() == device_id)
351            .cloned()
352            .ok_or_else(|| unavailable_device_error(device_id, discovered))?;
353        Ok(Self {
354            inner: Arc::new(CudaRuntimeState {
355                client,
356                device_id,
357                device_ordinal,
358                device_info,
359                primary_context,
360                identity: CudaRuntimeIdentity::fresh(),
361                allocation_domain: AllocationDomainId::fresh(),
362                raw_streams: (0..cubecl_stream_slots())
363                    .map(|_| OnceLock::new())
364                    .collect(),
365                cublas_handles: (0..cubecl_stream_slots())
366                    .map(|_| Mutex::new(None))
367                    .collect(),
368                pinned_scalar: Mutex::new(PinnedScalarSlot {
369                    ptr: std::ptr::null_mut(),
370                }),
371                workspace_retirements: Mutex::new(Default::default()),
372            }),
373        })
374    }
375
376    pub(crate) fn client(&self) -> &ComputeClient<CubeclCudaRuntime> {
377        &self.inner.client
378    }
379
380    /// Return the caller-selected CUDA device identity that this runtime targets.
381    ///
382    /// # Examples
383    ///
384    /// ```
385    /// use tenferro_gpu::{cuda::CudaDeviceId, cuda::CudaRuntime};
386    ///
387    /// let _device_id: fn(&CudaRuntime) -> CudaDeviceId = CudaRuntime::device_id;
388    /// ```
389    pub fn device_id(&self) -> CudaDeviceId {
390        self.inner.device_id
391    }
392
393    /// Return immutable metadata for the runtime's device.
394    ///
395    /// # Examples
396    ///
397    /// ```
398    /// use tenferro_gpu::cuda::CudaRuntime;
399    ///
400    /// let _info: fn(&CudaRuntime) -> &tenferro_gpu::cuda::CudaDeviceInfo =
401    ///     CudaRuntime::device_info;
402    /// ```
403    pub fn device_info(&self) -> &CudaDeviceInfo {
404        &self.inner.device_info
405    }
406
407    /// Return the allocation ownership domain of this runtime.
408    ///
409    /// # Examples
410    ///
411    /// ```
412    /// use tenferro_gpu::cuda::CudaRuntime;
413    ///
414    /// let _domain: fn(&CudaRuntime) -> tenferro_tensor::AllocationDomainId =
415    ///     CudaRuntime::allocation_domain;
416    /// ```
417    pub fn allocation_domain(&self) -> AllocationDomainId {
418        self.inner.allocation_domain
419    }
420
421    /// Report whether this CUDA session supports a GPU extension capability.
422    ///
423    /// The CUDA provider supports the full extension vocabulary: external
424    /// CubeCL kernels, native module loading, runtime compilation (NVRTC), raw
425    /// stream borrowing, and same-device copy. `PeerCopy` is reported as a
426    /// directional query; availability is hardware-dependent and is checked
427    /// per source/destination pair rather than here.
428    ///
429    /// # Examples
430    ///
431    /// ```
432    /// use tenferro_gpu::cuda::GpuExtensionCapability;
433    /// use tenferro_gpu::cuda::CudaRuntime;
434    ///
435    /// let _supports: fn(&CudaRuntime, GpuExtensionCapability) -> bool =
436    ///     CudaRuntime::supports_extension;
437    /// ```
438    pub fn supports_extension(&self, capability: GpuExtensionCapability) -> bool {
439        capabilities_for_device(capability)
440    }
441
442    pub(crate) fn device_ordinal(&self) -> usize {
443        self.inner.device_ordinal
444    }
445
446    pub(crate) fn primary_context(&self) -> CUcontext {
447        self.inner.primary_context.context()
448    }
449
450    /// Run `f` with the tenferro primary context current on this thread.
451    ///
452    /// Saves the calling thread's current CUDA device/context, activates the
453    /// tenferro primary context for the duration of `f`, and attempts to
454    /// restore the saved state on every exit path (normal return, `Err`, and
455    /// unwind). Restoration is best-effort: a failure to restore the
456    /// caller's previous device/context is logged to stderr rather than
457    /// returned. This is the scoped context authority used by vendor-library
458    /// lifecycle paths (plan creation/retirement) that run outside a
459    /// raw-session callback.
460    ///
461    /// # Errors
462    ///
463    /// Returns [`crate::Error::BackendSource`] when the tenferro primary
464    /// context cannot be activated (device or context driver failure); a
465    /// partial activation is best-effort rolled back before the error is
466    /// returned (rollback failures are discarded).
467    ///
468    /// # Examples
469    ///
470    /// ```
471    /// use tenferro_gpu::cuda::CudaRuntime;
472    ///
473    /// let _check: fn(&CudaRuntime) -> tenferro_tensor::Result<u64> = |rt| {
474    ///     rt.with_current_context("test.context", || 7)
475    /// };
476    /// ```
477    pub fn with_current_context<R>(
478        &self,
479        op: &'static str,
480        f: impl FnOnce() -> R,
481    ) -> crate::Result<R> {
482        let device_ordinal = i32::try_from(self.device_ordinal())
483            .map_err(|source| crate::Error::backend_source(op, source))?;
484        let _guard = RawContextRestore::enter(op, device_ordinal, self.primary_context())?;
485        Ok(f())
486    }
487
488    /// Flush pending CubeCL work on the current stream.
489    ///
490    /// Used by the raw-session enter protocol so raw library calls observe
491    /// previously enqueued CubeCL work.
492    pub(crate) fn flush_cubecl(&self, op: &'static str) -> crate::Result<()> {
493        self.inner.flush_cubecl(op)
494    }
495
496    /// Return the opaque identity of this exact executable runtime instance.
497    ///
498    /// # Examples
499    ///
500    /// ```
501    /// use tenferro_gpu::cuda::CudaRuntime;
502    ///
503    /// let _identity: fn(&CudaRuntime) -> tenferro_gpu::cuda::CudaRuntimeIdentity =
504    ///     CudaRuntime::runtime_identity;
505    /// ```
506    pub fn runtime_identity(&self) -> CudaRuntimeIdentity {
507        self.inner.identity.clone()
508    }
509
510    pub(crate) fn allocation_domain_id(&self) -> AllocationDomainId {
511        self.inner.allocation_domain
512    }
513
514    pub(crate) fn set_current_cuda_context(&self, op: &'static str) -> crate::Result<()> {
515        self.inner.set_current_cuda_context(op)
516    }
517
518    pub(crate) fn workspace_retirements(
519        &self,
520    ) -> &Mutex<super::workspace_retirement::WorkspaceRetirementQueue> {
521        &self.inner.workspace_retirements
522    }
523
524    pub(crate) fn state(&self) -> &CudaRuntimeState {
525        &self.inner
526    }
527
528    pub(crate) fn raw_cuda_stream(&self) -> crate::Result<u64> {
529        self.inner.raw_cuda_stream()
530    }
531
532    /// Run one cuBLAS enqueue with the handle for the current CubeCL stream.
533    ///
534    /// The caller must have the tenferro primary context current on this
535    /// thread (see [`CudaRuntimeState::set_current_cuda_context`]).
536    pub(crate) fn with_cublas_handle<R>(
537        &self,
538        op: &'static str,
539        pointer_mode: cublas_sys::cublasPointerMode_t,
540        cross_stream_handles: Vec<cubecl_runtime::server::Handle>,
541        execute: impl FnOnce(cublas_sys::cublasHandle_t) -> crate::Result<R>,
542    ) -> crate::Result<R> {
543        self.inner
544            .with_cublas_handle(op, pointer_mode, cross_stream_handles, execute)
545    }
546
547    pub(crate) fn finish_vendor_enqueue<R>(
548        &self,
549        op: &'static str,
550        cross_stream_handles: Vec<cubecl_runtime::server::Handle>,
551        result: crate::Result<R>,
552    ) -> crate::Result<R> {
553        self.inner
554            .finish_vendor_enqueue(op, cross_stream_handles, result)
555    }
556
557    pub(crate) fn stream_slot(&self) -> usize {
558        self.inner.stream_slot()
559    }
560
561    pub(crate) fn stream_slot_count(&self) -> usize {
562        self.inner.raw_streams.len()
563    }
564
565    pub(crate) fn is_current_stream_slot(&self, handle: &cubecl_runtime::server::Handle) -> bool {
566        self.inner.stream_slot_for(handle.stream) == self.inner.stream_slot()
567    }
568
569    /// Download up to [`PINNED_SCALAR_BYTES`] bytes from a device address
570    /// through the runtime-owned pinned staging slot.
571    ///
572    /// Enqueues an async device-to-host copy on the current thread's CubeCL
573    /// stream and synchronizes only that stream, so previously enqueued work
574    /// on the stream is observed without a device-wide barrier.
575    pub(crate) fn download_scalar_bytes(
576        &self,
577        device_addr: u64,
578        out: &mut [u8],
579        op: &'static str,
580        retained: cubecl_runtime::server::Handle,
581    ) -> crate::Result<()> {
582        self.inner
583            .download_scalar_bytes(device_addr, out, op, retained)
584    }
585
586    /// Resolve `handle`'s device address on the current stream, then flush
587    /// pending CubeCL work, in one device-thread hand-off.
588    ///
589    /// This is exactly `client().get_resource(handle)` followed by
590    /// `flush_cubecl`, in that order and on the same stream, so the ordering a
591    /// raw vendor enqueue relies on is unchanged; it only saves the second
592    /// blocking round trip (#1887). It does not publish a write (that remains
593    /// tensor4all/cubecl#16).
594    pub(crate) fn resolve_and_flush(
595        &self,
596        handle: cubecl_runtime::server::Handle,
597        op: &'static str,
598    ) -> crate::Result<u64> {
599        use cubecl_runtime::server::ComputeServer;
600
601        let stream_id = StreamId::current();
602        let binding = handle.binding();
603        let resolved = self
604            .inner
605            .client
606            .with_server(move |server| {
607                let resource = server.get_resource(binding, stream_id);
608                server.flush(stream_id).map(|()| resource)
609            })
610            .ok_or_else(|| crate::Error::runtime_state(op, "CubeCL server is unavailable"))?
611            .map_err(|err| crate::Error::backend_source(op, err))?
612            .map_err(|err| crate::Error::backend_source(op, err))?;
613        Ok(resolved.resource().ptr)
614    }
615
616    /// Copy `len` bytes of a CubeCL allocation into caller-owned host memory.
617    ///
618    /// See `CudaRuntimeState::download_into_host` for the ordering and
619    /// failure contract.
620    ///
621    /// # Safety
622    ///
623    /// `dst` must be valid for writes of `len` bytes and stay allocated until
624    /// this call returns. On error the caller must not free or reuse `dst`
625    /// (the device may still be writing it): leak it instead.
626    pub(crate) unsafe fn download_into_host(
627        &self,
628        handle: cubecl_runtime::server::Handle,
629        dst: *mut u8,
630        len: usize,
631        op: &'static str,
632    ) -> crate::Result<()> {
633        // SAFETY: forwarded caller contract.
634        unsafe { self.inner.download_into_host(handle, dst, len, op) }
635    }
636
637    /// Block the current thread until work submitted to the current CUDA stream completes.
638    ///
639    /// # Examples
640    ///
641    /// ```
642    /// use tenferro_gpu::cuda::CudaRuntime;
643    ///
644    /// let _sync: fn(&CudaRuntime) -> tenferro_tensor::Result<()> =
645    ///     CudaRuntime::synchronize;
646    /// ```
647    ///
648    /// # Errors
649    ///
650    /// Returns [`crate::Error::RuntimeState`] when CubeCL cannot expose the
651    /// current stream, or [`crate::Error::BackendSource`] when CUDA context or
652    /// stream synchronization fails.
653    pub fn synchronize(&self) -> crate::Result<()> {
654        self.inner.synchronize()
655    }
656}
657
658impl CudaRuntimeState {
659    fn stream_slot(&self) -> usize {
660        self.stream_slot_for(StreamId::current())
661    }
662
663    fn stream_slot_for(&self, stream_id: StreamId) -> usize {
664        stream_id.value as usize % self.raw_streams.len()
665    }
666
667    pub(crate) fn set_current_cuda_context(&self, op: &'static str) -> crate::Result<()> {
668        // Fast path: the tenferro primary context is already current on this
669        // thread. `cuCtxGetCurrent` only reads driver thread state, so this
670        // skips the per-op `cudaSetDevice` + `cuCtxSetCurrent` round trips.
671        // Runtime-API calls made afterwards operate on the current driver
672        // context, so no separate runtime-API device activation is needed.
673        if let Ok(Some(current)) = cudarc::driver::result::ctx::get_current() {
674            if current == self.primary_context.context() {
675                return Ok(());
676            }
677        }
678        // INVARIANT: CUDA ordinals are device identifiers; bad ordinals are
679        // reported by CUDA instead of indexing memory in tenferro.
680        let device_ordinal = i32::try_from(self.device_id.ordinal())
681            .map_err(|source| crate::Error::backend_source(op, source))?;
682        cudarc::runtime::result::device::set(device_ordinal)
683            .map_err(|err| crate::Error::backend_source(op, err))?;
684        unsafe { cudarc::driver::result::ctx::set_current(self.primary_context.context()) }
685            .map_err(|err| crate::Error::backend_source(op, err))
686    }
687
688    fn raw_cuda_stream(&self) -> crate::Result<u64> {
689        let stream_id = StreamId::current();
690        let slot = self.stream_slot();
691        if let Some(&stream) = self.raw_streams[slot].get() {
692            return Ok(stream);
693        }
694        let stream = self
695            .client
696            .with_server(move |server| {
697                server
698                    .raw_stream(stream_id)
699                    .map(|stream| stream as u64)
700                    .map_err(|err| crate::Error::backend_source("raw_cuda_stream", err))
701            })
702            .ok_or_else(|| {
703                crate::Error::runtime_state("raw_cuda_stream", "CubeCL server is unavailable")
704            })??;
705        Ok(*self.raw_streams[slot].get_or_init(|| stream))
706    }
707
708    fn flush_cubecl(&self, op: &'static str) -> crate::Result<()> {
709        self.client
710            .flush()
711            .map_err(|err| crate::Error::backend_source(op, err))
712    }
713
714    fn synchronize(&self) -> crate::Result<()> {
715        const OP: &str = "cubecl_runtime_synchronize";
716        // A cached raw stream does not drain CubeCL's host-side launch queue.
717        self.flush_cubecl(OP)?;
718        let stream = self.raw_cuda_stream()?;
719        self.synchronize_raw_stream(stream, OP)?;
720        // An explicit barrier also resolves deferred workspace retirements, so
721        // callers that synchronize observe released workspaces.
722        self.workspace_retirements
723            .lock()
724            .unwrap_or_else(|error| error.into_inner())
725            .drain_blocking(self);
726        Ok(())
727    }
728
729    pub(crate) fn synchronize_raw_stream(
730        &self,
731        stream: u64,
732        op: &'static str,
733    ) -> crate::Result<()> {
734        self.set_current_cuda_context(op)?;
735        unsafe { cuda_result::stream::synchronize(stream as usize as cudaStream_t) }
736            .map_err(|err| crate::Error::backend_source(op, err))
737    }
738
739    fn retire_initialized_streams(&self) -> bool {
740        const OP: &str = "cuda_runtime_drop";
741        if let Err(error) = self.set_current_cuda_context(OP) {
742            report_cuda_runtime_drop_error(&error);
743            return false;
744        }
745        let mut retired = true;
746        for stream in &self.raw_streams {
747            let Some(&stream) = stream.get() else {
748                continue;
749            };
750            // SAFETY: every initialized entry is a CubeCL-owned stream that
751            // remains live until this runtime state and its client are dropped.
752            if let Err(source) =
753                unsafe { cuda_result::stream::synchronize(stream as usize as cudaStream_t) }
754            {
755                retired = false;
756                report_cuda_runtime_drop_error(&crate::Error::backend_source(OP, source));
757            }
758        }
759        retired
760    }
761
762    fn with_cublas_handle<R>(
763        &self,
764        op: &'static str,
765        pointer_mode: cublas_sys::cublasPointerMode_t,
766        cross_stream_handles: Vec<cubecl_runtime::server::Handle>,
767        execute: impl FnOnce(cublas_sys::cublasHandle_t) -> crate::Result<R>,
768    ) -> crate::Result<R> {
769        let poisoned = || crate::Error::runtime_state(op, "cuBLAS handle cache lock poisoned");
770        let slot = self.stream_slot();
771        let mut cached = self.cublas_handles[slot].lock().map_err(|_| poisoned())?;
772        let handle = match *cached {
773            Some(ref handle) => handle.0,
774            None => {
775                if !cublas_library_present() {
776                    return Err(crate::Error::io_source(op, CublasLibraryMissing));
777                }
778                let stream = self.raw_cuda_stream()? as usize as cublas_sys::cudaStream_t;
779                let mut raw = std::ptr::null_mut();
780                // SAFETY: the caller holds the tenferro primary context current;
781                // the handle is created on this device and bound to this stream.
782                check_cublas(op, "cublasCreate", unsafe {
783                    cublas_sys::cublasCreate_v2(&mut raw)
784                })?;
785                // SAFETY: `raw` was just created and `stream` is owned by this
786                // runtime for the lifetime of the cached handle.
787                if let Err(err) = check_cublas(op, "cublasSetStream", unsafe {
788                    cublas_sys::cublasSetStream_v2(raw, stream)
789                }) {
790                    // SAFETY: `raw` is live and not stored anywhere else.
791                    let _ = unsafe { cublas_sys::cublasDestroy_v2(raw) };
792                    return Err(err);
793                }
794                *cached = Some(CublasStreamHandle(raw));
795                raw
796            }
797        };
798        // SAFETY: the per-stream slot lock is held across configuration and
799        // enqueue, so no caller can race this mutable handle state.
800        check_cublas(op, "cublasSetPointerMode", unsafe {
801            cublas_sys::cublasSetPointerMode_v2(handle, pointer_mode)
802        })?;
803        let result = execute(handle);
804        self.finish_vendor_enqueue(op, cross_stream_handles, result)
805    }
806
807    fn finish_vendor_enqueue<R>(
808        &self,
809        op: &'static str,
810        cross_stream_handles: Vec<cubecl_runtime::server::Handle>,
811        result: crate::Result<R>,
812    ) -> crate::Result<R> {
813        if cross_stream_handles.is_empty() {
814            return result;
815        }
816        let retirement = self.synchronize();
817        match (result, retirement) {
818            (Ok(value), Ok(())) => Ok(value),
819            (Err(error), Ok(())) => Err(error),
820            (Ok(_), Err(retirement)) => {
821                // No completion barrier was proven. Retain the foreign-stream
822                // allocations so their owners cannot reclaim or mutate them.
823                std::mem::forget(cross_stream_handles);
824                Err(crate::Error::backend_source(op, retirement))
825            }
826            (Err(error), Err(_retirement)) => {
827                std::mem::forget(cross_stream_handles);
828                Err(error)
829            }
830        }
831    }
832
833    fn download_scalar_bytes(
834        &self,
835        device_addr: u64,
836        out: &mut [u8],
837        op: &'static str,
838        retained: cubecl_runtime::server::Handle,
839    ) -> crate::Result<()> {
840        if out.len() > PINNED_SCALAR_BYTES {
841            return Err(crate::Error::Internal(format!(
842                "pinned scalar staging supports at most {PINNED_SCALAR_BYTES} bytes, got {}",
843                out.len()
844            )));
845        }
846        // Reused output addresses can already be cached: pointer lookup is not
847        // a queue barrier. Submit pending kernels before the raw D2H copy.
848        self.flush_cubecl(op)?;
849        self.set_current_cuda_context(op)?;
850        let stream = self.raw_cuda_stream()? as usize as cudaStream_t;
851        // ponytail: one shared staging slot serializes concurrent scalar
852        // downloads per runtime; add per-thread slots if that lock contends.
853        let mut slot = self
854            .pinned_scalar
855            .lock()
856            .map_err(|_| crate::Error::runtime_state(op, "pinned scalar staging lock poisoned"))?;
857        if slot.ptr.is_null() {
858            let mut ptr = std::ptr::null_mut();
859            // SAFETY: the primary context is current; the allocation is freed
860            // in this state's `Drop` with `cudaFreeHost`.
861            unsafe {
862                cuda_sys::cudaHostAlloc(
863                    &mut ptr,
864                    PINNED_SCALAR_BYTES,
865                    cuda_sys::cudaHostAllocDefault,
866                )
867            }
868            .result()
869            .map_err(|err| crate::Error::backend_source(op, err))?;
870            slot.ptr = ptr;
871        }
872        let src = super::interop::cuda_device_ptr_from_addr(device_addr, op)?;
873        // SAFETY: `slot.ptr` is a live pinned allocation of PINNED_SCALAR_BYTES
874        // bytes, `out.len()` is validated above, and the mutex guard keeps the
875        // slot exclusive until the copy below is known to have completed. Every
876        // exit that cannot prove completion abandons the slot instead of
877        // returning it, so exclusivity never rests on an unproven barrier.
878        let staging = unsafe { std::slice::from_raw_parts_mut(slot.ptr.cast::<u8>(), out.len()) };
879        // Neither submitting the copy nor waiting on it proves the device is
880        // done with `staging` and `retained` once it reports an error: an async
881        // CUDA call can surface a failure from an earlier launch on the stream,
882        // so a non-success return says nothing about what is still running.
883        // Both paths therefore leak the source allocation and abandon the
884        // staging slot rather than let a later download reuse a destination the
885        // device may still write, or let `Drop` `cudaFreeHost` it. Each failure
886        // leaks one slot and one handle; the next call allocates fresh ones.
887        // SAFETY: `src` is a residency-checked device address owned by this
888        // runtime, the copy length equals the destination slice length, and
889        // `stream` is the memoized CubeCL stream the copy is enqueued on.
890        let completed = unsafe { cuda_result::memcpy_dtoh_async(staging, src, stream) }
891            .and_then(|()| unsafe { cuda_result::stream::synchronize(stream) });
892        if let Err(err) = completed {
893            std::mem::forget(retained);
894            slot.ptr = std::ptr::null_mut();
895            return Err(crate::Error::backend_source(op, err));
896        }
897        out.copy_from_slice(staging);
898        Ok(())
899    }
900
901    /// Copy `len` bytes from the start of `handle` straight into host memory.
902    ///
903    /// This replaces `read_one` plus a host copy for whole-tensor downloads
904    /// (#2009): CubeCL's read stages into its own pinned (or, above 100 MB,
905    /// freshly zeroed pageable) buffer, and the result then had to be copied
906    /// again into the tensor's `Vec<T>`. Here the driver copies directly into
907    /// the destination the caller will own.
908    ///
909    /// Ordering: pending CubeCL launches are flushed, and `get_resource`
910    /// resolves the allocation on the current stream exactly as `read_one`
911    /// does, so writes queued on other CubeCL streams are waited for on this
912    /// stream. The copy is enqueued on that stream and the stream is
913    /// synchronized before returning, so `dst` is complete on `Ok`.
914    ///
915    /// Failure: an error from the copy or the barrier does not prove the
916    /// device is done with either side, so the source allocation is leaked and
917    /// the caller must leak `dst` as well (see the safety contract).
918    ///
919    /// # Safety
920    ///
921    /// `dst` must be valid for writes of `len` bytes until this returns, and
922    /// must not be freed or reused after an error.
923    unsafe fn download_into_host(
924        &self,
925        handle: cubecl_runtime::server::Handle,
926        dst: *mut u8,
927        len: usize,
928        op: &'static str,
929    ) -> crate::Result<()> {
930        if len == 0 {
931            return Ok(());
932        }
933        self.flush_cubecl(op)?;
934        let resource = self
935            .client
936            .get_resource(handle)
937            .map_err(|err| crate::Error::backend_source(op, err))?;
938        let available = usize::try_from(resource.resource().size).unwrap_or(usize::MAX);
939        if available < len {
940            return Err(crate::Error::Internal(format!(
941                "{op}: download of {len} bytes exceeds the {available}-byte allocation"
942            )));
943        }
944        let src = resource.resource().ptr;
945        self.set_current_cuda_context(op)?;
946        let stream = self.raw_cuda_stream()? as usize as cudarc::driver::sys::CUstream;
947        // SAFETY: `src` is the device address of a live CubeCL allocation of at
948        // least `len` bytes (checked above) kept alive by `resource`; `dst` is
949        // valid for `len` byte writes per the caller contract; `stream` is the
950        // current CubeCL stream that `get_resource` ordered the allocation on.
951        let completed = unsafe {
952            cudarc::driver::sys::cuMemcpyDtoHAsync_v2(dst.cast(), src, len, stream).result()
953        }
954        .and_then(|()| unsafe { cudarc::driver::result::stream::synchronize(stream) });
955        if let Err(err) = completed {
956            std::mem::forget(resource);
957            return Err(crate::Error::backend_source(op, err));
958        }
959        Ok(())
960    }
961
962    /// Destroy cached cuBLAS handles and free the pinned staging slot.
963    ///
964    /// Called from `Drop` after all initialized streams have retired, which
965    /// leaves the primary context current on the dropping thread.
966    fn release_cuda_library_resources(&mut self) {
967        for cached in &self.cublas_handles {
968            if let Ok(mut handle) = cached.lock() {
969                let Some(handle) = handle.take() else {
970                    continue;
971                };
972                // SAFETY: each stored handle is live and no longer reachable.
973                if let Err(err) = unsafe { cublas_sys::cublasDestroy_v2(handle.0) }.result() {
974                    report_cuda_resource_release_error("cuBLAS handle", &err);
975                }
976            }
977        }
978        if let Ok(mut slot) = self.pinned_scalar.lock() {
979            if !slot.ptr.is_null() {
980                // SAFETY: the slot owns exactly one live cudaHostAlloc allocation.
981                if let Err(err) = unsafe { cuda_sys::cudaFreeHost(slot.ptr) }.result() {
982                    report_cuda_resource_release_error("pinned scalar staging", &err);
983                }
984                slot.ptr = std::ptr::null_mut();
985            }
986        }
987    }
988}
989
990fn cubecl_stream_slots() -> usize {
991    usize::from(CubeClRuntimeConfig::get().streaming.max_streams.max(1))
992}
993
994/// Typed load failure for the dynamically loaded cuBLAS library.
995#[derive(Debug, thiserror::Error)]
996#[error(
997    "cuBLAS shared library not found; ensure `LD_LIBRARY_PATH` includes the CUDA toolkit library directory"
998)]
999struct CublasLibraryMissing;
1000
1001/// Report whether the cuBLAS shared library can be dynamically loaded.
1002fn cublas_library_present() -> bool {
1003    use std::sync::OnceLock;
1004    static PRESENT: OnceLock<bool> = OnceLock::new();
1005    // SAFETY: `is_culib_present` only probes candidate library names and does
1006    // not call cuBLAS function pointers or retain a library handle.
1007    *PRESENT.get_or_init(|| unsafe { cublas_sys::is_culib_present() })
1008}
1009
1010/// Map a non-success cuBLAS status to a typed provider error.
1011pub(super) fn check_cublas(
1012    op: &'static str,
1013    call: &'static str,
1014    status: cublas_sys::cublasStatus_t,
1015) -> crate::Result<()> {
1016    if matches!(status, cublas_sys::cublasStatus_t::CUBLAS_STATUS_SUCCESS) {
1017        Ok(())
1018    } else {
1019        Err(super::error::provider_status(
1020            op,
1021            "cuBLAS",
1022            call,
1023            status as i32,
1024        ))
1025    }
1026}
1027
1028#[cold]
1029fn report_cuda_resource_release_error(what: &'static str, err: &impl fmt::Debug) {
1030    eprintln!("tenferro-gpu: failed to release {what} during Drop: {err:?}");
1031}
1032
1033fn is_invalid_device_lookup(source: DriverError) -> bool {
1034    source.0 == CUresult::CUDA_ERROR_INVALID_DEVICE
1035}
1036
1037/// CUDA provider support for the shared GPU extension vocabulary.
1038///
1039/// See [`GpuExtensionCapability`](super::identity::GpuExtensionCapability) for
1040/// the vocabulary. `PeerCopy` is hardware/topology dependent and is therefore
1041/// reported false at the provider level; the directional query in the explicit
1042/// multi-GPU copy API decides availability per source/destination pair.
1043pub(crate) fn capabilities_for_device(capability: GpuExtensionCapability) -> bool {
1044    !matches!(capability, GpuExtensionCapability::PeerCopy)
1045}
1046
1047fn cuda_initialization_error<E>(
1048    device: CudaDeviceId,
1049    operation: &'static str,
1050    source: E,
1051) -> CudaDeviceError
1052where
1053    E: std::error::Error + Send + Sync + 'static,
1054{
1055    CudaDeviceError::Initialization {
1056        device,
1057        operation,
1058        source: Box::new(source),
1059    }
1060}
1061
1062impl Drop for CudaRuntimeState {
1063    fn drop(&mut self) {
1064        // Drop cannot surface errors, but the runtime must not release the
1065        // primary context while queued kernels on any initialized slot may
1066        // still reference it.
1067        if self.retire_initialized_streams() {
1068            // Every stream is retired, so queued workspace retirements are
1069            // complete and their handles can return to the pool.
1070            self.workspace_retirements
1071                .lock()
1072                .unwrap_or_else(|error| error.into_inner())
1073                .drain_blocking(self);
1074            // Retirement left the primary context current; release CUDA
1075            // library resources before the retained primary context drops.
1076            self.release_cuda_library_resources();
1077        }
1078        // On retirement failure, raw library resources intentionally leak:
1079        // their pointer-only owners have no Drop implementation, so Rust does
1080        // not reclaim resources that may still be in use asynchronously.
1081    }
1082}
1083
1084#[cfg(test)]
1085mod tests;