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;