Skip to main content

tenferro_gpu/
lib.rs

1//! GPU backend implementations for tenferro tensors.
2//!
3//! # Examples
4//!
5//! ```rust
6//! #[cfg(feature = "cuda")]
7//! use tenferro_gpu::{cuda::cuda_devices, cuda::CudaBackend, cuda::CudaDeviceError};
8//!
9//! #[cfg(feature = "cuda")]
10//! fn first_cuda_backend() -> Result<Option<CudaBackend>, CudaDeviceError> {
11//!     let devices = cuda_devices()?;
12//!     let Some(device) = devices.first() else {
13//!         return Ok(None);
14//!     };
15//!     Ok(Some(CudaBackend::new(device.id())?))
16//! }
17//!
18//! // This ordinary doctest checks the discovery-based selection API without
19//! // requiring CUDA hardware at test time.
20//! #[cfg(feature = "cuda")]
21//! let _example: fn() -> Result<Option<CudaBackend>, CudaDeviceError> = first_cuda_backend;
22//! ```
23// A misaligned pointer handed to a CUDA library is undefined behaviour and
24// fails only on some library versions: `cuDoubleComplex` is `double2`, which
25// CUDA declares `__align__(16)`, while `num_complex::Complex64` is 8-aligned,
26// so casting `&Complex64` to `*const cuDoubleComplex` produced a pointer
27// cuBLAS >= 12.9 faults on (issue #1870). Deny the whole cast class at this
28// FFI boundary rather than re-auditing it by hand.
29#![deny(clippy::cast_ptr_alignment)]
30
31#[cfg(feature = "cuda")]
32use std::any::Any;
33
34#[cfg(feature = "cuda")]
35mod cubecl;
36#[cfg(any(feature = "cuda", feature = "webgpu"))]
37mod event_domain_admission;
38#[cfg(any(feature = "cuda", feature = "webgpu"))]
39mod event_retirement;
40#[cfg(any(feature = "cuda", feature = "webgpu"))]
41mod kernels;
42#[cfg(any(feature = "cuda", feature = "webgpu"))]
43mod native_permutation;
44#[cfg(feature = "webgpu")]
45pub mod webgpu;
46
47/// CUDA provider namespace.
48#[cfg(feature = "cuda")]
49pub mod cuda {
50    pub use super::cubecl::{
51        cuda_capabilities, cuda_devices, cuda_runtime_engine_registration,
52        cuda_runtime_hardware_class, download_tensor, gpu_available, upload_tensor,
53        with_cuda_exec_session, CudaBackend, CudaComputeCapability, CudaDeviceError, CudaDeviceId,
54        CudaDeviceInfo, CudaDeviceUuid, CudaExecSession, CudaExtensionCache,
55        CudaExtensionCacheGuard, CudaRuntime, CudaRuntimeIdentity, CutensorWorkspaceStats,
56        GpuExtensionCapability, WorkspaceRetirementStats,
57    };
58
59    /// Public tenferro-wide CubeCL session (issue #1597).
60    ///
61    /// Exposes a narrow prelude of the CubeCL types needed to write and launch
62    /// `#[cube]` kernels against tenferro's GPU runtime. This module does not
63    /// re-export the whole of `cubecl`; downstream crates declare the framework
64    /// `t4a-cubecl` package explicitly.
65    pub mod cubecl {
66        pub use super::super::cubecl::session_cubecl::Session;
67        // Narrow prelude: only the types needed to describe a CubeCL launch.
68        pub use ::cubecl::prelude::{ArrayArg, CubeCount, CubeDim, TensorBinding};
69    }
70
71    /// The `cudarc` crate tenferro-gpu is built against (issue #1940).
72    ///
73    /// Downstream code that calls a CUDA vendor library (cuBLAS, cuSOLVER, ...)
74    /// on tenferro buffers through [`raw::Session`] should take its bindings
75    /// from here, so they always match tenferro's `cudarc` version and CUDA
76    /// version selection. The enabled features are `driver`, `runtime`,
77    /// `nvrtc`, `cublas`, `dynamic-loading` and `cuda-12080`; a downstream crate
78    /// that needs another `cudarc` module (for example `cusolver`) adds its own
79    /// `cudarc` dependency of the same `0.19` line with that feature, and Cargo
80    /// unifies the two into one crate. The re-export is not a stable API of its
81    /// own: it moves with tenferro's `cudarc` requirement.
82    ///
83    /// # Examples
84    ///
85    /// ```
86    /// use tenferro_gpu::cuda::cudarc::cublas::sys::cublasOperation_t;
87    ///
88    /// let no_transpose = cublasOperation_t::CUBLAS_OP_N;
89    /// assert_ne!(no_transpose, cublasOperation_t::CUBLAS_OP_T);
90    /// ```
91    pub use ::cudarc;
92
93    /// Type-safe raw CUDA extension session (issue #1597).
94    pub mod raw {
95        pub use super::super::cubecl::raw::{
96            CudaResourceGuard, DeviceBytes, Function, KernelArg, LaunchConfig, Module,
97            NvrtcOptions, Session, StreamRef, TensorMut, TensorRef,
98        };
99    }
100}
101
102/// Apple shared-allocation provider namespace.
103#[cfg(feature = "webgpu")]
104pub mod apple {
105    pub use super::webgpu::{AppleContext, AppleTransferStats};
106}
107
108#[cfg(any(feature = "cuda", feature = "webgpu"))]
109use tenferro_tensor::*;
110
111#[cfg(feature = "cuda")]
112pub(crate) mod backend {
113    pub use tenferro_tensor::backend::*;
114}
115
116#[cfg(feature = "cuda")]
117pub(crate) mod config {
118    pub use tenferro_tensor::config::*;
119}
120
121#[cfg(feature = "cuda")]
122pub(crate) mod types {
123    pub(crate) use crate::CubeclBuffer;
124    pub use tenferro_tensor::types::*;
125}
126
127/// Scalar-independent CubeCL allocation stored behind tensor backend-buffer
128/// trait objects; dtype is carried by the borrowed tensor descriptor.
129#[cfg(feature = "cuda")]
130pub(crate) struct CubeclBuffer {
131    handle: cubecl_runtime::server::Handle,
132    byte_len: usize,
133    device_ordinal: usize,
134    allocation_domain: AllocationDomainId,
135    allocation_id: AllocationId,
136    // Memoized device address resolved by the first raw-FFI access.
137    // Zero means "not memoized": no CUDA allocation lives at the null address.
138    //
139    // INVARIANT: in pinned CubeCL rev a2adda17, a retained handle's memory
140    // slice keeps its storage offset (pool coalescing merges only free
141    // slices) and its backing storage is never deallocated while any of its
142    // slices is live, so the resolved address is stable for this buffer's
143    // lifetime. Raw-FFI callers must still route cross-stream accesses
144    // through `get_resource` for CubeCL's stream alignment.
145    //
146    // INVARIANT (issue #1868): the memoized address is cleared whenever a
147    // CubeCL kernel is queued to write this buffer. Resolving the address
148    // through `get_resource` is a blocking server round trip, which is also
149    // what pushes queued kernels onto the CUstream; a raw vendor call reached
150    // through a memoized address skips that round trip and could otherwise be
151    // issued ahead of a kernel that precedes it in program order. Clearing on
152    // write keeps the fast path for read-only reuse and forces exactly one
153    // round trip after each write.
154    device_addr: std::sync::atomic::AtomicU64,
155}
156
157#[cfg(feature = "cuda")]
158static NEXT_CUDA_ALLOCATION_ID: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(1);
159
160#[cfg(feature = "cuda")]
161impl std::fmt::Debug for CubeclBuffer {
162    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
163        f.debug_struct("CubeclBuffer")
164            .field("byte_len", &self.byte_len)
165            .field("device_ordinal", &self.device_ordinal)
166            .field("allocation_domain", &self.allocation_domain)
167            .field("allocation_id", &self.allocation_id)
168            .finish()
169    }
170}
171
172#[cfg(feature = "cuda")]
173impl CubeclBuffer {
174    pub(crate) fn new(
175        handle: cubecl_runtime::server::Handle,
176        byte_len: usize,
177        device_ordinal: usize,
178        allocation_domain: AllocationDomainId,
179    ) -> Self {
180        Self {
181            handle,
182            byte_len,
183            device_ordinal,
184            allocation_domain,
185            allocation_id: AllocationId::from_backend_id(
186                NEXT_CUDA_ALLOCATION_ID.fetch_add(1, std::sync::atomic::Ordering::Relaxed),
187            ),
188            device_addr: std::sync::atomic::AtomicU64::new(0),
189        }
190    }
191
192    pub(crate) fn handle(&self) -> &cubecl_runtime::server::Handle {
193        &self.handle
194    }
195
196    /// Return the memoized device address, if one is currently valid.
197    pub(crate) fn cached_device_addr(&self) -> Option<u64> {
198        match self.device_addr.load(std::sync::atomic::Ordering::Acquire) {
199            0 => None,
200            addr => Some(addr),
201        }
202    }
203
204    /// Memoize the device address resolved through `get_resource` for this
205    /// buffer's handle.
206    pub(crate) fn memoize_device_addr(&self, addr: u64) {
207        self.device_addr
208            .store(addr, std::sync::atomic::Ordering::Release);
209    }
210
211    /// Drop the memoized address because a CubeCL kernel was queued to write
212    /// this buffer.
213    ///
214    /// The next raw-FFI access then resolves through `get_resource`, whose
215    /// blocking server round trip also pushes the queued kernel onto the
216    /// CUstream, so the vendor call cannot overtake it. See the
217    /// `device_addr` invariant and issue #1868.
218    pub(crate) fn invalidate_device_addr(&self) {
219        self.device_addr
220            .store(0, std::sync::atomic::Ordering::Release);
221    }
222
223    pub(crate) fn element_len<T: 'static>(&self) -> usize {
224        let element_size = std::mem::size_of::<T>();
225        debug_assert!(element_size != 0 && self.byte_len.is_multiple_of(element_size));
226        self.byte_len / element_size
227    }
228
229    pub(crate) fn device_ordinal(&self) -> usize {
230        self.device_ordinal
231    }
232
233    pub(crate) fn allocation_domain(&self) -> AllocationDomainId {
234        self.allocation_domain
235    }
236}
237
238#[cfg(feature = "cuda")]
239impl<T: Send + Sync + 'static> BackendStorage<T> for CubeclBuffer {
240    fn backend_family(&self) -> &'static str {
241        "cubecl"
242    }
243
244    fn len(&self) -> usize {
245        self.element_len::<T>()
246    }
247
248    fn allocation_domain(&self) -> Option<AllocationDomainId> {
249        Some(self.allocation_domain)
250    }
251
252    fn allocation_id(&self) -> Option<AllocationId> {
253        Some(self.allocation_id)
254    }
255
256    fn prepare_device_access(
257        &self,
258        request: DeviceAccessRequest<'_>,
259    ) -> std::result::Result<Box<dyn PreparedDeviceAccess>, DeviceAccessError> {
260        Ok(Box::new(crate::cubecl::dispatch::prepare_cubecl_access(
261            self, request,
262        )?))
263    }
264
265    fn as_any(&self) -> &dyn Any {
266        self
267    }
268}