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}