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
27struct 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 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
68pub(super) struct CudaExecSessionMarker;
71
72#[derive(Debug)]
93pub struct CudaExecSession<'a> {
94 backend: &'a mut CudaBackend,
95 _not_send_sync: PhantomData<Rc<()>>,
96}
97
98fn 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 pub fn runtime(&self) -> &CudaRuntime {
117 self.backend.runtime()
118 }
119
120 pub fn runtime_identity(&self) -> CudaRuntimeIdentity {
122 self.backend.runtime_identity()
123 }
124
125 pub fn supports(&self, capability: GpuExtensionCapability) -> bool {
140 self.backend.runtime().supports_extension(capability)
141 }
142
143 pub fn device_info(&self) -> &CudaDeviceInfo {
158 self.backend.runtime().device_info()
159 }
160
161 pub fn allocation_domain(&self) -> tenferro_tensor::AllocationDomainId {
176 self.backend.runtime().allocation_domain()
177 }
178
179 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 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 pub fn synchronize(&mut self) -> crate::Result<()> {
259 self.backend.runtime().synchronize()
260 }
261
262 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 let stream = runtime.raw_cuda_stream()?;
305 runtime.flush_cubecl(op)?;
307 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 let mut session = unsafe { raw::Session::new(runtime, cache, stream) };
316 f(&mut session)
317 }
318
319 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 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 #[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
400macro_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
461pub 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 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 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
670impl 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 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 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 with_session_entry_guard("CudaBackend", || f(&mut session))
744 }
745}