Skip to main content

tenferro_gpu/cubecl/
exec_session.rs

1use cubecl::prelude::{CubeElement, CubePrimitive};
2use num_complex::{Complex32, Complex64};
3use std::marker::PhantomData;
4use std::rc::Rc;
5use tenferro_tensor::backend::{
6    BackendSession, BackendSessionHost, ElementwiseFusionPlan, ElementwiseReadOp, SessionCachedDot,
7    TensorAnalytic, TensorBuffer, TensorDeviceTransfer, TensorDot, TensorElementwise, TensorFusion,
8    TensorIndexing, TensorReduction, TensorStructural,
9};
10use tenferro_tensor::config::{
11    CompareDir, DotGeneralConfig, GatherConfig, PadConfig, ScatterConfig, SliceConfig,
12};
13use tenferro_tensor::DType;
14use tenferro_tensor::{
15    with_session_entry_guard, TensorRank, TensorScalar, TensorViewCanonicalization,
16    TypedTensorView, TypedTensorViewMut,
17};
18use tenferro_tensor::{DotGeneralAccumulation, Tensor, TensorRead, TensorWrite, TypedTensor};
19
20use super::identity::GpuExtensionCapability;
21use super::{gemm, ops, runtime::RawContextRestore};
22use super::{
23    raw, session_cubecl, CudaBackend, CudaDeviceInfo, CudaExtensionCache, CudaRuntime,
24    CudaRuntimeIdentity,
25};
26
27/// Best-effort exit flush for a `with_cubecl` session.
28///
29/// Flushes once eagerly (returned to the caller as an error if it fails) and
30/// once more on `Drop` so a panic/unwind path still drains pending CubeCL
31/// work.
32struct CubeclExitFlush<'a> {
33    op: &'static str,
34    client: &'a cubecl::client::ComputeClient<cubecl_cuda::CudaRuntime>,
35    flushed: bool,
36}
37
38impl<'a> CubeclExitFlush<'a> {
39    fn new(
40        op: &'static str,
41        client: &'a cubecl::client::ComputeClient<cubecl_cuda::CudaRuntime>,
42    ) -> Self {
43        Self {
44            op,
45            client,
46            flushed: false,
47        }
48    }
49
50    /// Flush now and return the typed result.
51    fn flush_now(&mut self) -> crate::Result<()> {
52        self.client
53            .flush()
54            .map_err(|err| crate::Error::backend_source(self.op, err))?;
55        self.flushed = true;
56        Ok(())
57    }
58}
59
60impl Drop for CubeclExitFlush<'_> {
61    fn drop(&mut self) {
62        if !self.flushed {
63            let _ = self.client.flush();
64        }
65    }
66}
67
68/// Native-session marker for [`CudaExecSession`]; private to this crate so no other
69/// crate can create a token that claims to be this session.
70pub(super) struct CudaExecSessionMarker;
71
72/// Borrowed CUDA execution capability.
73///
74/// This is the single public execution-authority boundary for CUDA kernel
75/// extensions (issue #1597). External operation crates obtain it through
76/// [`with_cuda_exec_session`] and then borrow backend/device-scoped extension
77/// sessions via [`CudaExecSession::with_cubecl`] and
78/// [`CudaExecSession::with_raw`].
79///
80/// The session is not constructible by users and is `!Send + !Sync`: it
81/// carries thread-local execution capability. Success of an enrolled operation
82/// means the work was enqueued; only [`CudaExecSession::synchronize`] is a
83/// host barrier.
84///
85/// The backend owner is not an operation route, so an operation bound does not
86/// hold for it:
87///
88/// ```compile_fail
89/// fn requires_elementwise<B: tenferro_tensor::TensorElementwise>() {}
90/// requires_elementwise::<tenferro_gpu::cuda::CudaBackend>();
91/// ```
92#[derive(Debug)]
93pub struct CudaExecSession<'a> {
94    backend: &'a mut CudaBackend,
95    _not_send_sync: PhantomData<Rc<()>>,
96}
97
98/// The typed tensor behind `input`, or the refusal this method produces for one.
99///
100/// Callers reach this from a match on `input.dtype()`, so `None` means the tag table
101/// and the runtime dtype disagree rather than a caller mistake.
102fn gpu_resident_typed<'a, T: TensorScalar>(
103    op: &'static str,
104    input: &'a Tensor,
105) -> crate::Result<&'a TypedTensor<T>> {
106    input.as_typed::<T>().ok_or_else(|| {
107        crate::Error::unsupported(
108            op,
109            "an externally defined payload is not supported by this GPU operation",
110        )
111    })
112}
113
114impl CudaExecSession<'_> {
115    /// Borrow the provider runtime without exposing the backend.
116    pub fn runtime(&self) -> &CudaRuntime {
117        self.backend.runtime()
118    }
119
120    /// Return the identity of the borrowed provider runtime.
121    pub fn runtime_identity(&self) -> CudaRuntimeIdentity {
122        self.backend.runtime_identity()
123    }
124
125    /// Report whether this session supports a GPU extension capability.
126    ///
127    /// # Examples
128    ///
129    /// ```
130    /// use tenferro_gpu::cuda::{CudaExecSession, GpuExtensionCapability};
131    ///
132    /// // Method-call check only: `CudaExecSession` is not user-constructible, so
133    /// // the example asserts the method is callable from an external crate.
134    /// fn check(session: &CudaExecSession<'_>, capability: GpuExtensionCapability) -> bool {
135    ///     session.supports(capability)
136    /// }
137    /// let _ = check;
138    /// ```
139    pub fn supports(&self, capability: GpuExtensionCapability) -> bool {
140        self.backend.runtime().supports_extension(capability)
141    }
142
143    /// Borrow immutable metadata for the session's device.
144    ///
145    /// # Examples
146    ///
147    /// ```
148    /// use tenferro_gpu::cuda::CudaExecSession;
149    ///
150    /// // Method-call check only: `CudaExecSession` is not user-constructible, so
151    /// // the example asserts the method is callable from an external crate.
152    /// fn check(session: &CudaExecSession<'_>) {
153    ///     let _ = session.device_info();
154    /// }
155    /// let _ = check;
156    /// ```
157    pub fn device_info(&self) -> &CudaDeviceInfo {
158        self.backend.runtime().device_info()
159    }
160
161    /// Return the allocation ownership domain of this session.
162    ///
163    /// # Examples
164    ///
165    /// ```
166    /// use tenferro_gpu::cuda::CudaExecSession;
167    ///
168    /// // Method-call check only: `CudaExecSession` is not user-constructible, so
169    /// // the example asserts the method is callable from an external crate.
170    /// fn check(session: &CudaExecSession<'_>) -> tenferro_tensor::AllocationDomainId {
171    ///     session.allocation_domain()
172    /// }
173    /// let _ = check;
174    /// ```
175    pub fn allocation_domain(&self) -> tenferro_tensor::AllocationDomainId {
176        self.backend.runtime().allocation_domain()
177    }
178
179    /// Validate that a dense GPU tensor is resident on this exact session:
180    /// CubeCL-backed, same allocation domain, and placed on this runtime's
181    /// CUDA device. Rejects host tensors, foreign-backend buffers, and
182    /// foreign-runtime/device tensors without an implicit transfer.
183    ///
184    /// This is the credentialed public-seam residency guard for extension
185    /// crates that receive a session but must validate inputs before entering
186    /// a `with_raw`/`with_cubecl` sub-session.
187    ///
188    /// # Errors
189    ///
190    /// Returns [`crate::Error::RuntimeState`] when the tensor is not resident
191    /// on this exact session runtime/device.
192    ///
193    /// # Examples
194    ///
195    /// ```
196    /// use tenferro_gpu::cuda::CudaExecSession;
197    ///
198    /// // Method-call check only: `CudaExecSession` is not user-constructible.
199    /// fn check(session: &CudaExecSession<'_>, tensor: &tenferro_tensor::Tensor) -> tenferro_tensor::Result<()> {
200    ///     session.ensure_gpu_resident(tensor, "test.ensure_gpu_resident")
201    /// }
202    /// let _ = check;
203    /// ```
204    pub fn ensure_gpu_resident(&self, input: &Tensor, op: &'static str) -> crate::Result<()> {
205        match input.dtype() {
206            DType::F32 => super::dispatch::ensure_resident_on_runtime(
207                self.runtime(),
208                gpu_resident_typed::<f32>("ensure_gpu_resident", input)?,
209                op,
210            ),
211            DType::F64 => super::dispatch::ensure_resident_on_runtime(
212                self.runtime(),
213                gpu_resident_typed::<f64>("ensure_gpu_resident", input)?,
214                op,
215            ),
216            DType::I32 => super::dispatch::ensure_resident_on_runtime(
217                self.runtime(),
218                gpu_resident_typed::<i32>("ensure_gpu_resident", input)?,
219                op,
220            ),
221            DType::I64 => super::dispatch::ensure_resident_on_runtime(
222                self.runtime(),
223                gpu_resident_typed::<i64>("ensure_gpu_resident", input)?,
224                op,
225            ),
226            DType::Bool => super::dispatch::ensure_resident_on_runtime(
227                self.runtime(),
228                gpu_resident_typed::<bool>("ensure_gpu_resident", input)?,
229                op,
230            ),
231            DType::C32 => super::dispatch::ensure_resident_on_runtime(
232                self.runtime(),
233                gpu_resident_typed::<Complex32>("ensure_gpu_resident", input)?,
234                op,
235            ),
236            DType::C64 => super::dispatch::ensure_resident_on_runtime(
237                self.runtime(),
238                gpu_resident_typed::<Complex64>("ensure_gpu_resident", input)?,
239                op,
240            ),
241            // A caller-owned payload has no GPU implementation for this operation.
242            DType::External(_) => Err(crate::Error::unsupported(
243                "ensure_gpu_resident",
244                "an externally defined payload is not supported by this GPU operation",
245            )),
246        }
247    }
248
249    /// Block the host until work enqueued on the session's stream completes.
250    ///
251    /// This is the only host barrier on the success path; ordinary successful
252    /// session operations only enqueue.
253    ///
254    /// # Errors
255    ///
256    /// Returns [`crate::Error::BackendSource`] when CUDA stream
257    /// synchronization fails.
258    pub fn synchronize(&mut self) -> crate::Result<()> {
259        self.backend.runtime().synchronize()
260    }
261
262    /// Borrow the type-safe raw CUDA extension session for one operation.
263    ///
264    /// The enter/exit protocol is fully contained in this call: a definite
265    /// CubeCL stream is captured on the current thread, pending CubeCL work is
266    /// flushed, the calling thread's previous device/context is saved, the
267    /// tenferro primary context is activated, the callback runs, and the
268    /// previous device/context is best-effort restored on return, `Err`, or
269    /// unwind (restoration failures are logged to stderr). The success path
270    /// does not synchronize.
271    ///
272    /// # Errors
273    ///
274    /// Returns [`crate::Error::BackendSource`] when CubeCL cannot expose or
275    /// flush the stream, or when the CUDA context cannot be entered. Context
276    /// restoration on exit is best-effort: a failure to restore the caller's
277    /// previous device/context is logged to stderr rather than propagated, so
278    /// a callback result is never replaced by a restore error.
279    ///
280    /// # Examples
281    ///
282    /// ```
283    /// use tenferro_gpu::cuda::CudaExecSession;
284    ///
285    /// // Method-call check only: `CudaExecSession` is not user-constructible.
286    /// fn check(session: &mut CudaExecSession<'_>) -> tenferro_tensor::Result<()> {
287    ///     session.with_raw("test.raw", |raw| {
288    ///         let _ = raw.stream();
289    ///         Ok(SessionOutcome::Done)
290    ///     })?;
291    ///     Ok(())
292    /// }
293    /// enum SessionOutcome { Done }
294    /// let _ = check;
295    /// ```
296    pub fn with_raw<R>(
297        &mut self,
298        op: &'static str,
299        f: impl for<'s> FnOnce(&mut raw::Session<'s>) -> crate::Result<R>,
300    ) -> crate::Result<R> {
301        let runtime = self.backend.runtime().clone();
302        let cache = self.backend.cuda_extension_cache();
303        // 1. Capture the definite CubeCL stream on this thread.
304        let stream = runtime.raw_cuda_stream()?;
305        // 2. Flush pending CubeCL work so raw library calls observe it.
306        runtime.flush_cubecl(op)?;
307        // 3-4. Save previous context, activate the tenferro primary context.
308        let device_ordinal = i32::try_from(runtime.device_ordinal())
309            .map_err(|source| crate::Error::backend_source(op, source))?;
310        let _guard = RawContextRestore::enter(op, device_ordinal, runtime.primary_context())?;
311        // 5. Build the unique raw session and run the callback.
312        // SAFETY: `_guard` keeps the primary context current for the whole
313        // `Session<'s>` borrow; `stream` is the captured CubeCL stream bound to
314        // the current thread.
315        let mut session = unsafe { raw::Session::new(runtime, cache, stream) };
316        f(&mut session)
317    }
318
319    /// Borrow the public tenferro-wide CubeCL session for one operation.
320    ///
321    /// The session exposes the exact tenferro CubeCL client bound to this
322    /// runtime. Pending CubeCL work is flushed before entering and again on
323    /// exit (including `Err` and unwind) so a later raw-session or host read
324    /// observes the enqueued work. The success path does not synchronize.
325    ///
326    /// # Examples
327    ///
328    /// ```
329    /// use tenferro_gpu::cuda::CudaExecSession;
330    ///
331    /// fn check(session: &mut CudaExecSession<'_>) {
332    ///     let _ = session.with_cubecl("test.cubecl", |_cubecl| Ok(()));
333    /// }
334    /// let _ = check;
335    /// ```
336    ///
337    /// # Errors
338    ///
339    /// Returns the callback's error, or [`crate::Error::BackendSource`] when
340    /// pending CubeCL work cannot be flushed on entry or exit.
341    pub fn with_cubecl<R>(
342        &mut self,
343        op: &'static str,
344        f: impl for<'s> FnOnce(&session_cubecl::Session<'s>) -> crate::Result<R>,
345    ) -> crate::Result<R> {
346        let runtime = self.backend.runtime().clone();
347        runtime.flush_cubecl(op)?;
348        let session = unsafe { session_cubecl::Session::new(runtime) };
349        // Best-effort exit flush on every path via Drop.
350        let mut _flush_guard = CubeclExitFlush::new(op, session.client());
351        let result = f(&session);
352        let flush_result = _flush_guard.flush_now();
353        match result {
354            Ok(value) => {
355                flush_result?;
356                Ok(value)
357            }
358            Err(err) => {
359                let _ = flush_result;
360                Err(err)
361            }
362        }
363    }
364
365    #[doc(hidden)]
366    pub fn tril_typed<T>(&self, input: &TypedTensor<T>, k: i64) -> crate::Result<TypedTensor<T>>
367    where
368        T: CubeElement + TensorScalar + CubePrimitive + Clone,
369    {
370        self.backend.tril_typed(input, k)
371    }
372
373    #[doc(hidden)]
374    pub fn slice_typed<T>(
375        &self,
376        input: &TypedTensor<T>,
377        config: &SliceConfig,
378    ) -> crate::Result<TypedTensor<T>>
379    where
380        T: CubeElement + TensorScalar + CubePrimitive + Clone,
381    {
382        self.backend.slice_typed(input, config)
383    }
384
385    /// Borrow the CUDA extension cache owned by the provider runtime.
386    #[doc(hidden)]
387    pub fn cuda_extension_cache(&self) -> &CudaExtensionCache {
388        self.backend.cuda_extension_cache()
389    }
390
391    #[doc(hidden)]
392    pub fn triu_typed<T>(&self, input: &TypedTensor<T>, k: i64) -> crate::Result<TypedTensor<T>>
393    where
394        T: CubeElement + TensorScalar + CubePrimitive + Clone,
395    {
396        self.backend.triu_typed(input, k)
397    }
398}
399
400// Typed view canonicalization runs on the session, never on the backend
401// owner: the owner is not an execution surface (#1946 F6).
402macro_rules! impl_session_view_canonicalization {
403    ($to_contiguous:ident; $($ty:ty),* $(,)?) => {
404        $(
405            impl<R> TensorViewCanonicalization<$ty, R> for CudaExecSession<'_>
406            where
407                R: TensorRank,
408            {
409                fn to_contiguous(
410                    &mut self,
411                    view: &TypedTensorView<'_, $ty, R>,
412                ) -> crate::Result<TypedTensor<$ty, R>> {
413                    self.backend
414                        .$to_contiguous(view, "CudaExecSession::to_contiguous")
415                }
416
417                fn copy_into(
418                    &mut self,
419                    src: &TypedTensorView<'_, $ty, R>,
420                    dst: &mut TypedTensorViewMut<'_, $ty, R>,
421                ) -> crate::Result<()> {
422                    self.backend
423                        .copy_view_to_view_typed(src, dst, "CudaExecSession::copy_into")
424                }
425            }
426        )*
427    };
428}
429
430impl_session_view_canonicalization!(
431    to_contiguous_view_cutensor_or_cubecl; f32, f64, Complex32, Complex64
432);
433impl_session_view_canonicalization!(to_contiguous_view_typed; i32, i64);
434
435impl<R> TensorViewCanonicalization<bool, R> for CudaExecSession<'_>
436where
437    R: TensorRank,
438{
439    fn to_contiguous(
440        &mut self,
441        _view: &TypedTensorView<'_, bool, R>,
442    ) -> crate::Result<TypedTensor<bool, R>> {
443        Err(super::error::unsupported_dtype(
444            "CudaExecSession::to_contiguous",
445            crate::DType::Bool,
446        ))
447    }
448
449    fn copy_into(
450        &mut self,
451        _src: &TypedTensorView<'_, bool, R>,
452        _dst: &mut TypedTensorViewMut<'_, bool, R>,
453    ) -> crate::Result<()> {
454        Err(super::error::unsupported_dtype(
455            "CudaExecSession::copy_into",
456            crate::DType::Bool,
457        ))
458    }
459}
460
461/// Visit a CUDA execution session through the erased backend-session surface.
462///
463/// This is the public entry point that borrows CUDA execution authority for
464/// the duration of the callback (issue #1597). The callback cannot return a
465/// borrow of the reconstructed session, so the authority cannot escape the
466/// scope.
467///
468/// Returns `None` when `session` is not a CUDA execution session.
469///
470/// # Examples
471///
472/// ```
473/// use tenferro_gpu::cuda::{with_cuda_exec_session, CudaExecSession};
474///
475/// // Call-check only: the visitor borrows CUDA execution authority for the
476/// // duration of the callback.
477/// fn check(session: &mut dyn tenferro_tensor::backend::BackendSession) {
478///     let _ = with_cuda_exec_session(session, |_session| 0usize);
479/// }
480/// let _ = check;
481/// ```
482pub fn with_cuda_exec_session<B, R>(
483    session: &mut B,
484    f: impl for<'a> FnOnce(&'a mut CudaExecSession<'a>) -> R,
485) -> Option<R>
486where
487    B: BackendSession + ?Sized,
488{
489    let data = session
490        .native_session()?
491        .into_marked_ptr::<CudaExecSessionMarker>()?;
492    // SAFETY: only `CudaExecSession::native_session` creates a token with the
493    // crate-private `CudaExecSessionMarker`, and it points that token at a live
494    // `CudaExecSession`. The token borrowed `*session` exclusively, and this function
495    // keeps holding `session: &mut B` for the whole scoped visit.
496    Some(unsafe { f(data.cast::<CudaExecSession<'static>>().as_mut()) })
497}
498
499macro_rules! delegate {
500    ($trait:path {
501        $(fn $method:ident($($arg:ident: $arg_ty:ty),* $(,)?) -> $ret:ty;)*
502    }) => {
503        impl $trait for CudaExecSession<'_> {
504            $(
505                fn $method(&mut self, $($arg: $arg_ty),*) -> $ret {
506                    self.backend.$method($($arg),*)
507                }
508            )*
509        }
510    };
511}
512
513macro_rules! delegate_ops {
514    ($trait:path {
515        $(fn $method:ident($($arg:ident: $arg_ty:ty),* $(,)?) -> $ret:ty;)*
516    } $(override { $($custom:item)* })?) => {
517        impl $trait for CudaExecSession<'_> {
518            $(
519                fn $method(&mut self, $($arg: $arg_ty),*) -> $ret {
520                    ops::$method(self.backend, $($arg),*)
521                }
522            )*
523            $($($custom)*)?
524        }
525    };
526}
527
528delegate_ops!(TensorElementwise {
529    fn add_read(lhs: TensorRead<'_>, rhs: TensorRead<'_>) -> crate::Result<Tensor>;
530    fn sub_read(lhs: TensorRead<'_>, rhs: TensorRead<'_>) -> crate::Result<Tensor>;
531    fn mul_read(lhs: TensorRead<'_>, rhs: TensorRead<'_>) -> crate::Result<Tensor>;
532    fn neg_read(input: TensorRead<'_>) -> crate::Result<Tensor>;
533    fn conj_read(input: TensorRead<'_>) -> crate::Result<Tensor>;
534    fn div_read(lhs: TensorRead<'_>, rhs: TensorRead<'_>) -> crate::Result<Tensor>;
535    fn rem_read(lhs: TensorRead<'_>, rhs: TensorRead<'_>) -> crate::Result<Tensor>;
536    fn abs_read(input: TensorRead<'_>) -> crate::Result<Tensor>;
537    fn sign_read(input: TensorRead<'_>) -> crate::Result<Tensor>;
538    fn maximum_read(lhs: TensorRead<'_>, rhs: TensorRead<'_>) -> crate::Result<Tensor>;
539    fn minimum_read(lhs: TensorRead<'_>, rhs: TensorRead<'_>) -> crate::Result<Tensor>;
540    fn compare_read(lhs: TensorRead<'_>, rhs: TensorRead<'_>, dir: &CompareDir) -> crate::Result<Tensor>;
541    fn select_read(pred: TensorRead<'_>, on_true: TensorRead<'_>, on_false: TensorRead<'_>) -> crate::Result<Tensor>;
542    fn clamp_read(input: TensorRead<'_>, lower: TensorRead<'_>, upper: TensorRead<'_>) -> crate::Result<Tensor>;
543    fn rem(lhs: &Tensor, rhs: &Tensor) -> crate::Result<Tensor>;
544} override {
545    // Read-into elementwise dispatch must stay session-shaped: the allocating
546    // fallback in `tenferro-tensor` is generic over `TensorElementwise`, and the
547    // session is the only type in this crate that implements it now. The native
548    // read-into kernels still run first, so no work moves onto the allocating path.
549    fn elementwise_read_into(
550        &mut self,
551        op: ElementwiseReadOp,
552        inputs: &[TensorRead<'_>],
553        mut out: TensorWrite<'_>,
554    ) -> crate::Result<()> {
555        if inputs.len() != op.arity() {
556            return Err(crate::Error::invalid_argument(
557                op.label(),
558                "inputs",
559                format!("expected {} inputs, got {}", op.arity(), inputs.len()),
560            ));
561        }
562        tenferro_tensor::backend::validate_read_into_destination(op.label(), inputs, &out)?;
563        if let Some(result) = self.backend.elementwise_read_into_native(op, inputs, &mut out) {
564            return result;
565        }
566        tenferro_tensor::backend::elementwise_read_into_via_allocating_ops(self, op, inputs, out)
567    }
568});
569
570delegate_ops!(TensorAnalytic {
571    fn exp_read(input: TensorRead<'_>) -> crate::Result<Tensor>;
572    fn log_read(input: TensorRead<'_>) -> crate::Result<Tensor>;
573    fn sin_read(input: TensorRead<'_>) -> crate::Result<Tensor>;
574    fn cos_read(input: TensorRead<'_>) -> crate::Result<Tensor>;
575    fn tanh_read(input: TensorRead<'_>) -> crate::Result<Tensor>;
576    fn sqrt_read(input: TensorRead<'_>) -> crate::Result<Tensor>;
577    fn rsqrt_read(input: TensorRead<'_>) -> crate::Result<Tensor>;
578    fn pow_read(lhs: TensorRead<'_>, rhs: TensorRead<'_>) -> crate::Result<Tensor>;
579    fn expm1_read(input: TensorRead<'_>) -> crate::Result<Tensor>;
580    fn log1p_read(input: TensorRead<'_>) -> crate::Result<Tensor>;
581    fn erf_read(input: TensorRead<'_>) -> crate::Result<Tensor>;
582});
583
584delegate_ops!(TensorStructural {
585    fn transpose_read(input: TensorRead<'_>, perm: &[usize]) -> crate::Result<Tensor>;
586    fn reshape_read(input: TensorRead<'_>, shape: &[usize]) -> crate::Result<Tensor>;
587    fn broadcast_in_dim_read(input: TensorRead<'_>, shape: &[usize], dims: &[usize]) -> crate::Result<Tensor>;
588    fn to_contiguous_read(input: TensorRead<'_>) -> crate::Result<Tensor>;
589    fn copy_read_into(src: TensorRead<'_>, dst: TensorWrite<'_>) -> crate::Result<()>;
590    fn cast(input: &Tensor, to: tenferro_tensor::DType) -> crate::Result<Tensor>;
591    fn extract_diagonal(input: &Tensor, axis_a: usize, axis_b: usize) -> crate::Result<Tensor>;
592    fn embed_diagonal(input: &Tensor, axis_a: usize, axis_b: usize) -> crate::Result<Tensor>;
593    fn tril(input: &Tensor, k: i64) -> crate::Result<Tensor>;
594    fn triu(input: &Tensor, k: i64) -> crate::Result<Tensor>;
595});
596
597delegate_ops!(TensorReduction {
598    fn reduce_sum_read(input: TensorRead<'_>, axes: &[usize]) -> crate::Result<Tensor>;
599    fn reduce_prod_read(input: TensorRead<'_>, axes: &[usize]) -> crate::Result<Tensor>;
600    fn reduce_max_read(input: TensorRead<'_>, axes: &[usize]) -> crate::Result<Tensor>;
601    fn reduce_min_read(input: TensorRead<'_>, axes: &[usize]) -> crate::Result<Tensor>;
602    fn reduce_sum_squares_read(input: TensorRead<'_>, axes: &[usize]) -> crate::Result<Tensor>;
603});
604
605delegate_ops!(TensorDot {
606    fn dot_general_with_conj(
607        lhs: &Tensor,
608        rhs: &Tensor,
609        config: &DotGeneralConfig,
610        lhs_conj: bool,
611        rhs_conj: bool,
612    ) -> crate::Result<Tensor>;
613    fn dot_general_read(
614        lhs: TensorRead<'_>,
615        rhs: TensorRead<'_>,
616        config: &DotGeneralConfig,
617    ) -> crate::Result<Tensor>;
618    fn dot_general_read_into_accum(
619        lhs: TensorRead<'_>,
620        rhs: TensorRead<'_>,
621        config: &DotGeneralConfig,
622        accumulation: DotGeneralAccumulation,
623        out: TensorWrite<'_>,
624    ) -> crate::Result<()>;
625});
626
627delegate_ops!(TensorIndexing {
628    fn gather(
629        operand: &Tensor,
630        start_indices: &Tensor,
631        config: &GatherConfig,
632    ) -> crate::Result<Tensor>;
633    fn scatter(
634        operand: &Tensor,
635        scatter_indices: &Tensor,
636        updates: &Tensor,
637        config: &ScatterConfig,
638    ) -> crate::Result<Tensor>;
639    fn slice(input: &Tensor, config: &SliceConfig) -> crate::Result<Tensor>;
640    fn dynamic_slice(
641        input: &Tensor,
642        starts: &Tensor,
643        slice_sizes: &[usize],
644    ) -> crate::Result<Tensor>;
645    fn dynamic_update_slice(
646        operand: &Tensor,
647        update: &Tensor,
648        starts: &Tensor,
649    ) -> crate::Result<Tensor>;
650    fn pad(input: &Tensor, config: &PadConfig) -> crate::Result<Tensor>;
651    fn concatenate(inputs: &[&Tensor], axis: usize) -> crate::Result<Tensor>;
652    fn reverse(input: &Tensor, axes: &[usize]) -> crate::Result<Tensor>;
653});
654
655delegate_ops!(TensorFusion {
656    fn execute_elementwise_fusion(
657        inputs: &[&Tensor],
658        plan: &ElementwiseFusionPlan,
659    ) -> crate::Result<Option<Vec<Tensor>>>;
660    fn execute_broadcast_multiply(
661        lhs: TensorRead<'_>,
662        lhs_shape: &[usize],
663        lhs_dims: &[usize],
664        rhs: TensorRead<'_>,
665        rhs_shape: &[usize],
666        rhs_dims: &[usize],
667    ) -> crate::Result<Option<Tensor>>;
668});
669
670// CUDA device buffers return to the runtime allocator on drop, so the
671// session keeps the trait's no-op reclaim; the owner is not a buffer surface.
672impl TensorBuffer for CudaExecSession<'_> {}
673
674delegate!(TensorDeviceTransfer {
675    fn download_to_host(tensor: TensorRead<'_>) -> crate::Result<Tensor>;
676    fn upload_host_tensor(tensor: TensorRead<'_>) -> crate::Result<Tensor>;
677});
678
679impl SessionCachedDot for CudaExecSession<'_> {
680    // Read-based cached dot paths keep strided operands on device; the plan
681    // cache is per backend, so the runtime cache slot stays unused.
682    fn dot_general_read_cached(
683        &mut self,
684        _cache_slot: Option<usize>,
685        lhs: TensorRead<'_>,
686        rhs: TensorRead<'_>,
687        config: &DotGeneralConfig,
688    ) -> crate::Result<Tensor> {
689        gemm::dot_general_read_allocating(self.backend, lhs, rhs, config, false, false)
690    }
691
692    fn dot_general_with_conj_read_cached(
693        &mut self,
694        _cache_slot: Option<usize>,
695        lhs: TensorRead<'_>,
696        rhs: TensorRead<'_>,
697        config: &DotGeneralConfig,
698        lhs_conj: bool,
699        rhs_conj: bool,
700    ) -> crate::Result<Tensor> {
701        gemm::dot_general_read_allocating(self.backend, lhs, rhs, config, lhs_conj, rhs_conj)
702    }
703}
704
705impl BackendSession for CudaExecSession<'_> {
706    fn vdot_read(&mut self, lhs: TensorRead<'_>, rhs: TensorRead<'_>) -> crate::Result<Tensor> {
707        ops::vdot_read(self.backend, lhs, rhs)
708    }
709
710    fn norm_squared_read(&mut self, input: TensorRead<'_>) -> crate::Result<Tensor> {
711        ops::norm_squared_read(self.backend, input)
712    }
713
714    fn axpby_read_into_accum(
715        &mut self,
716        alpha: tenferro_tensor::ContractionScalar,
717        x: TensorRead<'_>,
718        beta: tenferro_tensor::ContractionScalar,
719        y: TensorWrite<'_>,
720    ) -> crate::Result<()> {
721        ops::axpby_read_into_accum(self.backend, alpha, x, beta, y)
722    }
723
724    fn native_session(&mut self) -> Option<tenferro_tensor::NativeSessionRef<'_>> {
725        // SAFETY: `CudaExecSessionMarker` is private to this crate, and this is the only
726        // place a token carrying it is created; it always points to a
727        // `CudaExecSession`, exclusively borrowed for the token lifetime.
728        Some(unsafe { tenferro_tensor::NativeSessionRef::new::<CudaExecSessionMarker, _>(self) })
729    }
730}
731
732impl BackendSessionHost for CudaBackend {
733    fn with_backend_session<R: Send>(
734        &mut self,
735        f: impl FnOnce(&mut dyn BackendSession) -> R + Send,
736    ) -> Result<R, tenferro_tensor::SessionEntryError> {
737        let mut session = CudaExecSession {
738            backend: self,
739            _not_send_sync: PhantomData,
740        };
741        // The portable in-session guard rejects nested entry before `f` runs;
742        // the CUDA runtime must never re-enter a session closure.
743        with_session_entry_guard("CudaBackend", || f(&mut session))
744    }
745}