Skip to main content

tenferro_gpu/cubecl/
session_cubecl.rs

1//! Public tenferro-wide CubeCL session (issue #1597).
2//!
3//! [`Session`] is the promoted, public view over the existing owner-scoped
4//! CubeCL integration helpers in [`super::interop`]. While a session is alive,
5//! the wrapped `ComputeClient` is the exact tenferro CubeCL client bound to the
6//! tenferro `CudaRuntime`; work enqueued through it is flushed at session exit.
7//!
8//! The prelude is intentionally narrow: only the types needed to write and
9//! launch CubeCL kernels are re-exported. Downstream crates must depend on the
10//! framework's `t4a-cubecl` package (see the design doc) to name `#[cube]`
11//! kernels; this module does **not** re-export the whole of `cubecl`.
12
13use std::marker::PhantomData;
14use std::rc::Rc;
15
16use cubecl::client::ComputeClient;
17use cubecl::prelude::{ArrayArg, CubeCount, CubeDim, CubeElement, CubePrimitive, TensorBinding};
18use cubecl_cuda::CudaRuntime as CubeclCudaRuntime;
19
20use tenferro_tensor::{TensorRank, TensorScalar, TensorWrite, TypedTensor};
21
22use super::runtime::CudaRuntime;
23
24/// Public CubeCL extension session.
25///
26/// Not constructible by users; obtained only from
27/// [`CudaExecSession::with_cubecl`](super::exec_session::CudaExecSession::with_cubecl).
28/// The session borrows the exact tenferro CubeCL client for the request scope
29/// and is `!Send + !Sync` by construction, so the execution authority cannot
30/// migrate to another thread.
31pub struct Session<'s> {
32    runtime: CudaRuntime,
33    _scope: PhantomData<&'s ()>,
34    _not_send_sync: PhantomData<Rc<()>>,
35}
36
37impl<'s> Session<'s> {
38    /// Wrap the runtime for one scoped callback.
39    ///
40    /// # Safety
41    ///
42    /// Caller must ensure the callback runs on a thread that owns the tenferro
43    /// primary context activation for `runtime`, and that the returned borrow
44    /// does not outlive the current context.
45    pub(crate) unsafe fn new(runtime: CudaRuntime) -> Self {
46        Self {
47            runtime,
48            _scope: PhantomData,
49            _not_send_sync: PhantomData,
50        }
51    }
52
53    /// Borrow the tenferro CubeCL client.
54    pub fn client(&self) -> &ComputeClient<CubeclCudaRuntime> {
55        self.runtime.client()
56    }
57
58    /// Build a CubeCL tensor binding for a GPU-backed tensor.
59    ///
60    /// # Errors
61    ///
62    /// Returns [`crate::Error::RuntimeState`] when the tensor is not resident
63    /// on this session's runtime/device, or [`crate::Error::Validation`] when
64    /// its layout cannot be bound.
65    pub fn tensor_binding<T>(
66        &self,
67        tensor: &TypedTensor<T, impl TensorRank>,
68        op: &'static str,
69    ) -> crate::Result<TensorBinding<CubeclCudaRuntime>>
70    where
71        T: CubeElement + TensorScalar + Clone,
72    {
73        super::dispatch::ensure_resident_on_runtime(&self.runtime, tensor, op)?;
74        super::interop::typed_tensor_binding(tensor, op)
75    }
76
77    /// Build a CubeCL array argument for a GPU-backed tensor.
78    ///
79    /// # Errors
80    ///
81    /// Returns [`crate::Error::RuntimeState`] when the tensor is not resident
82    /// on this session's runtime/device, or [`crate::Error::Validation`] when
83    /// its layout cannot be bound.
84    pub fn array_arg<T>(
85        &self,
86        tensor: &TypedTensor<T, impl TensorRank>,
87        op: &'static str,
88    ) -> crate::Result<ArrayArg<CubeclCudaRuntime>>
89    where
90        T: CubeElement + TensorScalar + Clone,
91    {
92        super::dispatch::ensure_resident_on_runtime(&self.runtime, tensor, op)?;
93        super::interop::typed_tensor_array_arg(tensor, op)
94    }
95
96    /// Allocate a dense GPU tensor on the session's device.
97    ///
98    /// # Errors
99    ///
100    /// Returns [`crate::Error::Validation`] when the shape product overflows,
101    /// or [`crate::Error::BackendSource`] when allocation fails.
102    pub fn alloc_output<T>(&self, shape: &[usize]) -> crate::Result<TypedTensor<T>>
103    where
104        T: CubeElement + TensorScalar + Clone + Send + Sync + 'static,
105    {
106        super::interop::alloc_output(&self.runtime, shape)
107    }
108
109    /// Allocate a dense GPU tensor zero-filled with the session's device.
110    ///
111    /// Reuses the backend's fill-zero structural kernel; it never uploads a
112    /// host tensor or exposes a device pointer to the caller.
113    ///
114    /// # Errors
115    ///
116    /// Returns [`crate::Error::Validation`] when the shape product, output
117    /// byte length, or launch count overflows, [`crate::Error::RuntimeState`]
118    /// when the output is not resident, or [`crate::Error::BackendSource`]
119    /// when allocation or backend resource inspection fails.
120    ///
121    /// # Examples
122    ///
123    /// ```
124    /// use tenferro_gpu::cuda::cubecl::Session;
125    ///
126    /// fn check(session: &Session<'_>) -> tenferro_tensor::Result<()> {
127    ///     let _ = session.alloc_zero_output::<f32>(&[4])?;
128    ///     Ok(())
129    /// }
130    /// ```
131    pub fn alloc_zero_output<T>(&self, shape: &[usize]) -> crate::Result<TypedTensor<T>>
132    where
133        T: CubeElement + CubePrimitive + TensorScalar + Clone + Send + Sync + 'static,
134    {
135        super::interop::alloc_zero_output(&self.runtime, shape)
136    }
137
138    /// Overwrite a caller-owned CUDA destination with an exact `+0.0`.
139    ///
140    /// This is the `beta = 0` reset for re-executing an accumulate-form
141    /// operation into storage the caller keeps: the previous contents are never
142    /// read, so a stale `NaN` or `Inf` cannot survive and `-0.0` is normalized
143    /// to `+0.0`, which `0 * y` cannot promise. A compact destination
144    /// (including a nonzero-offset region view) is filled by one stream-ordered
145    /// memset, so nothing is allocated and no kernel is compiled; a strided
146    /// region view runs one native fill kernel over its own coordinates and
147    /// leaves every element outside the region untouched.
148    ///
149    /// Supports `F32`, `F64`, `I32`, `I64`, `C32`, `C64`, and compact `Bool`
150    /// destinations.
151    ///
152    /// # Errors
153    ///
154    /// Returns [`crate::Error::RuntimeState`] when the destination is not
155    /// resident on this session's runtime, [`crate::Error::Validation`] when
156    /// the destination layout or its byte span cannot be represented,
157    /// [`crate::Error::Unsupported`] for a strided `Bool` destination, and
158    /// [`crate::Error::BackendSource`] when the fill cannot be enqueued.
159    ///
160    /// # Examples
161    ///
162    /// ```
163    /// use tenferro_gpu::cuda::cubecl::Session;
164    /// use tenferro_tensor::{Tensor, TensorWrite};
165    ///
166    /// fn reset(session: &Session<'_>) -> tenferro_tensor::Result<()> {
167    ///     let destination = session.alloc_output::<f32>(&[4])?;
168    ///     let mut destination = Tensor::from_typed::<f32>(destination);
169    ///     session.fill_zero_write(TensorWrite::from_tensor(&mut destination))
170    /// }
171    /// ```
172    pub fn fill_zero_write(&self, output: TensorWrite<'_>) -> crate::Result<()> {
173        super::interop::fill_zero_write(&self.runtime, output)
174    }
175
176    /// Scale a mutable CUDA tensor in place by a real factor.
177    ///
178    /// Supports F32, F64, C32, and C64 payloads; the factor is interpreted as
179    /// a real scalar for complex payloads. This is the op-family scale
180    /// primitive used for normalization after a vendor transform; it never
181    /// exposes a device pointer to the caller.
182    ///
183    /// # Errors
184    ///
185    /// Returns [`crate::Error::RuntimeState`] when the tensor is not resident
186    /// on this session's runtime, a typed layout error for a non-zero-offset
187    /// discontinuous view, or the typed unsupported-dtype error for other
188    /// payloads.
189    pub fn scale_tensor_write(&self, output: TensorWrite<'_>, factor: f64) -> crate::Result<()> {
190        super::interop::scale_tensor_write(&self.runtime, output, factor)
191    }
192
193    /// Return the cube count for a one-dimensional kernel domain of `len`.
194    ///
195    /// # Errors
196    ///
197    /// Returns [`crate::Error::Validation`] carrying
198    /// `ValidationError::InvalidArgument` when the one-dimensional launch for
199    /// `len` elements would require more than `u32::MAX` workgroups.
200    pub fn cube_count_1d(&self, len: usize) -> crate::Result<CubeCount> {
201        super::interop::cube_count_for_len(len)
202    }
203
204    /// Return the standard one-dimensional CubeCL launch dimension.
205    pub fn cube_dim_1d(&self) -> CubeDim {
206        super::interop::cube_dim_1d()
207    }
208}