1use std::any::{Any, TypeId};
64use std::collections::{HashMap, VecDeque};
65use std::fmt;
66use std::num::NonZeroUsize;
67use std::ops::Deref;
68use std::ptr::NonNull;
69use std::sync::atomic::{AtomicU64, Ordering};
70use std::sync::{Arc, Mutex, MutexGuard, OnceLock};
71
72use cubecl::client::ComputeClient;
73use cubecl::features::AtomicUsage;
74use cubecl::prelude::{
75 ArrayArg, ComplexCore as CubeComplex, CubeDim, CubeElement, CubePrimitive, Float as CubeFloat,
76 Numeric as CubeNumeric,
77};
78use cubecl::prelude::{CubeCount, Int as CubeInt, StorageType, TensorBinding, Type};
79use cubecl_cuda::CudaRuntime as CubeclCudaRuntime;
80use num_complex::{Complex32, Complex64};
81use tenferro_core_ops::PrimitiveOpKind;
82
83use tenferro_tensor::CacheStats;
84use tenferro_tensor::{
85 ContractionScalar, DType, DotGeneralAccumulation, ElementwiseReadOp, TensorRead, TensorWrite,
86};
87
88use crate::backend::{BackendRuntimeCache, TensorBackend, TensorDeviceTransfer};
89use crate::config::{
90 CompareDir, DotGeneralConfig, GatherConfig, PadConfig, ScatterConfig, SliceConfig,
91};
92use crate::kernels::reduce::{self as cubecl_reduce, ReduceStrategy};
93use crate::kernels::{diagonal, elementwise, indexing, structural};
94use crate::native_permutation::{
95 NativePermutationKind, NativePermutationPlan, NativeStridedCopyPlan, NativeTransposeTile,
96};
97use crate::{
98 DeviceId, DeviceKind, GpuBackendKind, MemoryKind, Placement, StorageBuffer, Tensor, TensorRank,
99 TensorScalar, TensorView, TensorViewMut, TypedTensor, TypedTensorView, TypedTensorViewMut,
100};
101
102macro_rules! preset_scalar {
104 (F32) => {
105 f32
106 };
107 (F64) => {
108 f64
109 };
110 (I32) => {
111 i32
112 };
113 (I64) => {
114 i64
115 };
116 (Bool) => {
117 bool
118 };
119 (C32) => {
120 num_complex::Complex32
121 };
122 (C64) => {
123 num_complex::Complex64
124 };
125}
126mod blas1;
127mod capability;
128mod device;
129pub(crate) mod dispatch;
130mod error;
131mod event_domain;
132mod exec_session;
133mod ffi;
134mod fusion;
135mod gemm;
136mod identity;
137pub(crate) mod interop;
138mod memory;
139pub(crate) mod op_descriptor;
140mod permutation;
141mod plan_cache;
142pub(crate) mod raw;
143mod runtime;
144mod runtime_adapter;
145pub(crate) mod session_cubecl;
146mod workspace_retirement;
147
148pub use workspace_retirement::WorkspaceRetirementStats;
149
150pub use gemm::CutensorWorkspaceStats;
151
152use dispatch::{
153 alloc_bool_output, alloc_output, bool_tensor_array_arg, bool_view_array_arg, comptime_sequence,
154 cube_count_for_len, cube_dim_1d, dtype_mismatch, ensure_axes_unique, ensure_axis, ensure_rank,
155 ensure_resident_on_runtime, ensure_view_mut_resident_on_runtime,
156 ensure_view_resident_on_runtime, launch_binary, launch_binary_bool_tensor, launch_binary_parts,
157 launch_binary_tensor, launch_bool_tensor_into, launch_compare_bool, launch_nullary_bool_into,
158 launch_nullary_into, launch_select_bool, launch_ternary, launch_unary,
159 launch_unary_bool_tensor, launch_unary_tensor, launch_unary_tensor_into, runtime_sequence,
160 ternary_dtype_mismatch, typed_tensor_array_arg, typed_tensor_array_arg_as,
161 typed_tensor_binding, typed_tensor_mut_array_arg, typed_view_array_arg, typed_view_binding,
162 typed_view_mut_array_arg,
163};
164use error::{unsupported_dtype, unsupported_operation};
165
166pub use capability::cuda_capabilities;
167pub use device::{cuda_devices, CudaDeviceError, CudaDeviceId, CudaDeviceInfo};
168#[doc(hidden)]
169pub use exec_session::{with_cuda_exec_session, CudaExecSession};
170pub use identity::{CudaComputeCapability, CudaDeviceUuid, GpuExtensionCapability};
171pub use memory::{download_tensor, upload_tensor};
172pub use runtime::{gpu_available, CudaRuntime, CudaRuntimeIdentity};
173pub use runtime_adapter::{cuda_runtime_engine_registration, cuda_runtime_hardware_class};
174
175fn op_name(
176 kind: PrimitiveOpKind,
177 launch: op_descriptor::GpuLaunchKind,
178) -> crate::Result<&'static str> {
179 op_descriptor::require_gpu_descriptor(kind, launch).map(|descriptor| descriptor.name)
180}
181
182fn ensure_atomic_add_supported<T: CubePrimitive>(
183 client: &ComputeClient<CubeclCudaRuntime>,
184 op: &'static str,
185) -> crate::Result<()> {
186 let elem = T::as_type_native_unchecked().elem_type();
187 let atomic_ty = Type::new(StorageType::Atomic(elem));
188 if client
189 .properties()
190 .atomic_type_usage(atomic_ty)
191 .contains(AtomicUsage::Add)
192 {
193 Ok(())
194 } else {
195 Err(unsupported_operation(
196 op,
197 "CubeCL runtime does not support atomic add",
198 ))
199 }
200}
201
202fn checked_dim_product(
203 op: &'static str,
204 role: &'static str,
205 shape: &[usize],
206) -> crate::Result<usize> {
207 shape.iter().try_fold(1usize, |acc, &dim| {
208 acc.checked_mul(dim).ok_or_else(|| {
209 crate::Error::invalid_argument(
210 op,
211 role,
212 format!("{role} product overflow for shape {shape:?}"),
213 )
214 })
215 })
216}
217
218fn view_strides_i64(strides: &[isize], op: &'static str) -> crate::Result<Vec<i64>> {
219 strides
220 .iter()
221 .map(|&stride| {
222 i64::try_from(stride).map_err(|_| {
223 crate::Error::invalid_argument(
224 op,
225 "layout",
226 format!("view stride {stride} exceeds CubeCL i64 metadata limit"),
227 )
228 })
229 })
230 .collect()
231}
232
233fn view_offset_i64(offset: isize, op: &'static str) -> crate::Result<i64> {
234 i64::try_from(offset).map_err(|_| {
235 crate::Error::invalid_argument(
236 op,
237 "layout",
238 format!("view offset {offset} exceeds CubeCL i64 metadata limit"),
239 )
240 })
241}
242
243fn launch_native_materialization<E: CubePrimitive>(
244 backend: &CudaBackend,
245 output: ArrayArg<CubeclCudaRuntime>,
246 input: ArrayArg<CubeclCudaRuntime>,
247 plan: &NativePermutationPlan,
248 op: &'static str,
249) -> crate::Result<()> {
250 if plan.len == 0 {
251 return Ok(());
252 }
253 if plan.kind == NativePermutationKind::TiledTranspose {
254 if let Some(config) = NativeTransposeTile::selected(op)? {
255 let block_rows = config.block_rows as usize;
256 let padding = config.padding as usize;
257 let vector_width = config.vector_width as usize;
258 let src_offset = usize::try_from(plan.src_offset).map_err(|_| {
259 crate::Error::invalid_argument(
260 op,
261 "offset",
262 "tiled transpose requires a non-negative source offset",
263 )
264 })?;
265 if let Some((cubes_x, cubes_y, cubes_z)) = config.dispatch_grid(
266 op,
267 plan.dims[0],
268 plan.dims[1],
269 plan.dims.get(2).copied().unwrap_or(1),
270 65_535,
271 )? {
272 let batch_stride = plan.tiled_matrix_len(op)?;
273 unsafe {
274 structural::tiled_transpose_kernel::launch_unchecked::<E, CubeclCudaRuntime>(
278 backend.runtime().client(),
279 CubeCount::Static(cubes_x, cubes_y, cubes_z),
280 CubeDim::new_2d(config.tile / config.vector_width, config.block_rows),
281 output,
282 input,
283 src_offset,
284 batch_stride,
285 plan.dims[0],
286 plan.dims[1],
287 config.tile as usize,
288 block_rows,
289 padding,
290 vector_width,
291 );
292 }
293 return Ok(());
294 }
295 }
296 }
297 let src_strides = view_strides_i64(&plan.src_strides, op)?;
298 let src_offset = view_offset_i64(plan.src_offset, op)?;
299 unsafe {
300 structural::materialize_strided_kernel::launch_unchecked::<E, CubeclCudaRuntime>(
303 backend.runtime().client(),
304 cube_count_for_len(plan.len)?,
305 cube_dim_1d(),
306 output,
307 input,
308 runtime_sequence(&plan.dims),
309 runtime_sequence(&src_strides),
310 src_offset,
311 plan.len,
312 plan.dims.len(),
313 );
314 }
315 Ok(())
316}
317
318fn scatter_update_len(meta: &ScatterLaunchMeta) -> crate::Result<usize> {
319 let batch_len = checked_dim_product("scatter", "batch shape", &meta.batch_shape)?;
320 let window_len =
321 checked_dim_product("scatter", "window update shape", &meta.window_shape_updates)?;
322 batch_len.checked_mul(window_len).ok_or_else(|| {
323 crate::Error::invalid_argument(
324 "scatter",
325 "shape",
326 format!(
327 "scatter update domain product overflow for batch {:?} and window {:?}",
328 meta.batch_shape, meta.window_shape_updates
329 ),
330 )
331 })
332}
333
334mod ops;
335
336#[derive(Clone)]
337pub struct CudaBackend {
338 inner: Arc<CudaBackendState>,
339}
340
341struct CudaBackendState {
342 cutensor: OnceLock<ffi::cutensor::CutensorHandle>,
346 extension_cache: CudaExtensionCache,
347 cutensor_workspace_max_retained_bytes: AtomicU64,
350 cutensor_workspace_temporary_uses: AtomicU64,
354 rt: CudaRuntime,
355}
356
357impl fmt::Debug for CudaBackend {
358 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
359 f.debug_struct("CudaBackend")
360 .field("runtime", &self.inner.rt)
361 .field("cuda_extension_cache", &self.inner.extension_cache)
362 .field("cutensor_initialized", &self.inner.cutensor.get().is_some())
363 .finish_non_exhaustive()
364 }
365}
366
367#[doc(hidden)]
369pub struct CudaExtensionCache {
370 inner: Mutex<CudaExtensionCacheInner>,
371}
372
373impl fmt::Debug for CudaExtensionCache {
374 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
375 f.debug_struct("CudaExtensionCache")
376 .field("max_entries", &self.max_entries())
377 .field("stats", &self.stats())
378 .finish_non_exhaustive()
379 }
380}
381
382const DEFAULT_CUDA_EXTENSION_CACHE_MAX_ENTRIES: usize = 16;
383const DEFAULT_CUDA_EXTENSION_CACHE_RETAINED_BYTES: usize = 64 * 1024 * 1024;
384
385const DEFAULT_CUTENSOR_WORKSPACE_MAX_RETAINED_BYTES: u64 = 10 << 30;
394
395struct CudaExtensionCacheEntry {
396 value: Box<dyn Any + Send>,
397 retained_bytes: usize,
398}
399
400struct CudaExtensionCacheInner {
401 max_entries: NonZeroUsize,
402 max_retained_bytes: NonZeroUsize,
403 entries: HashMap<TypeId, CudaExtensionCacheEntry>,
404 order: VecDeque<TypeId>,
405 retained_bytes: usize,
406 stats: CacheStats,
407}
408
409impl CudaExtensionCacheInner {
410 fn new(max_entries: NonZeroUsize) -> Self {
411 Self {
412 max_entries,
413 max_retained_bytes: NonZeroUsize::new(DEFAULT_CUDA_EXTENSION_CACHE_RETAINED_BYTES)
414 .unwrap_or(NonZeroUsize::MIN),
415 entries: HashMap::new(),
416 order: VecDeque::new(),
417 retained_bytes: 0,
418 stats: CacheStats::empty(),
419 }
420 }
421
422 fn evict_to_limit(&mut self) {
423 while self.entries.len() > self.max_entries.get()
424 || self.retained_bytes > self.max_retained_bytes.get()
425 {
426 let Some(type_id) = self.order.pop_front() else {
427 break;
428 };
429 if let Some(entry) = self.entries.remove(&type_id) {
430 self.retained_bytes = self.retained_bytes.saturating_sub(entry.retained_bytes);
431 self.stats.evictions = self.stats.evictions.saturating_add(1);
432 }
433 }
434 }
435
436 fn insert<T: Send + 'static>(&mut self, type_id: TypeId, value: T, retained_bytes: usize) {
437 self.entries.insert(
438 type_id,
439 CudaExtensionCacheEntry {
440 value: Box::new(value),
441 retained_bytes,
442 },
443 );
444 self.order.retain(|&existing| existing != type_id);
445 self.order.push_back(type_id);
446 self.retained_bytes = self
447 .entries
448 .values()
449 .map(|entry| entry.retained_bytes)
450 .sum();
451 self.evict_to_limit();
452 }
453
454 fn snapshot_stats(&self) -> CacheStats {
455 CacheStats {
456 entries: self.entries.len(),
457 retained_bytes: self.retained_bytes,
458 ..self.stats
459 }
460 }
461
462 fn refresh_retained_bytes(&mut self) {
463 self.retained_bytes = self
464 .entries
465 .values()
466 .map(|entry| entry.retained_bytes)
467 .sum();
468 }
469}
470
471impl CudaExtensionCache {
472 fn poisoned_lock_error() -> crate::Error {
473 crate::Error::runtime_state("cuda_extension_cache", "extension cache lock poisoned")
474 }
475
476 fn lock_inner(&self) -> crate::Result<MutexGuard<'_, CudaExtensionCacheInner>> {
477 self.inner.lock().map_err(|_| Self::poisoned_lock_error())
478 }
479
480 pub fn new() -> Self {
497 let max_entries = NonZeroUsize::new(DEFAULT_CUDA_EXTENSION_CACHE_MAX_ENTRIES)
498 .unwrap_or(NonZeroUsize::MIN);
499 Self::with_max_entries(max_entries)
500 }
501
502 pub fn with_max_entries(max_entries: NonZeroUsize) -> Self {
504 Self {
505 inner: Mutex::new(CudaExtensionCacheInner::new(max_entries)),
506 }
507 }
508
509 pub fn is_empty(&self) -> crate::Result<bool> {
523 Ok(self.lock_inner()?.entries.is_empty())
524 }
525
526 pub fn clear(&self) -> crate::Result<()> {
535 let mut inner = self.lock_inner()?;
536 inner.entries.clear();
537 inner.order.clear();
538 inner.retained_bytes = 0;
539 let clears = inner.stats.clears.saturating_add(1);
540 inner.stats = CacheStats {
541 clears,
542 ..CacheStats::empty()
543 };
544 Ok(())
545 }
546
547 pub fn stats(&self) -> crate::Result<CacheStats> {
552 let inner = self.lock_inner()?;
553 Ok(inner.snapshot_stats())
554 }
555
556 pub fn max_entries(&self) -> crate::Result<NonZeroUsize> {
561 Ok(self.lock_inner()?.max_entries)
562 }
563
564 pub fn max_retained_bytes(&self) -> crate::Result<NonZeroUsize> {
569 Ok(self.lock_inner()?.max_retained_bytes)
570 }
571
572 pub fn set_max_entries(&self, max_entries: NonZeroUsize) -> crate::Result<()> {
578 let mut inner = self.lock_inner()?;
579 inner.max_entries = max_entries;
580 inner.evict_to_limit();
581 Ok(())
582 }
583
584 pub fn set_max_retained_bytes(&self, max_retained_bytes: NonZeroUsize) -> crate::Result<()> {
591 let mut inner = self.lock_inner()?;
592 inner.max_retained_bytes = max_retained_bytes;
593 inner.evict_to_limit();
594 Ok(())
595 }
596
597 pub fn get_or_try_init<T>(
614 &self,
615 init: impl FnOnce() -> crate::Result<T>,
616 ) -> crate::Result<CudaExtensionCacheGuard<'_, T>>
617 where
618 T: Send + 'static,
619 {
620 let type_id = TypeId::of::<T>();
621 let mut inner = self.lock_inner()?;
622 if !inner.entries.contains_key(&type_id) {
623 inner.stats.misses = inner.stats.misses.saturating_add(1);
624 inner.insert(type_id, init()?, std::mem::size_of::<T>());
625 } else {
626 inner.stats.hits = inner.stats.hits.saturating_add(1);
627 }
628 let value = inner
629 .entries
630 .get(&type_id)
631 .and_then(|entry| entry.value.downcast_ref::<T>())
632 .map(NonNull::from)
633 .ok_or_else(|| {
634 crate::Error::runtime_state(
635 "cuda_extension_cache",
636 format!(
637 "stored entry for {} is missing or has the wrong type",
638 std::any::type_name::<T>()
639 ),
640 )
641 })?;
642 Ok(CudaExtensionCacheGuard {
643 inner,
644 type_id,
645 value,
646 _marker: std::marker::PhantomData,
647 })
648 }
649
650 pub(crate) fn get_cloned<T>(&self) -> crate::Result<Option<T>>
651 where
652 T: Clone + 'static,
653 {
654 let inner = self.lock_inner()?;
655 inner
656 .entries
657 .get(&TypeId::of::<T>())
658 .map(|entry| {
659 entry.value.downcast_ref::<T>().cloned().ok_or_else(|| {
660 crate::Error::runtime_state(
661 "cuda_extension_cache",
662 format!(
663 "stored entry for {} is missing or has the wrong type",
664 std::any::type_name::<T>()
665 ),
666 )
667 })
668 })
669 .transpose()
670 }
671
672 pub(crate) fn update_retained_bytes<T: 'static>(
681 &self,
682 retained_bytes: usize,
683 ) -> crate::Result<()> {
684 let type_id = TypeId::of::<T>();
685 let mut inner = self.lock_inner()?;
686 if let Some(entry) = inner.entries.get_mut(&type_id) {
687 entry.retained_bytes = retained_bytes;
688 inner.refresh_retained_bytes();
689 if inner.retained_bytes > inner.max_retained_bytes.get() {
690 inner.entries.remove(&type_id);
691 inner.order.retain(|candidate| *candidate != type_id);
692 inner.stats.evictions = inner.stats.evictions.saturating_add(1);
693 inner.refresh_retained_bytes();
694 }
695 inner.evict_to_limit();
696 }
697 Ok(())
698 }
699}
700
701impl Default for CudaExtensionCache {
702 fn default() -> Self {
703 Self::new()
704 }
705}
706
707#[doc(hidden)]
709pub struct CudaExtensionCacheGuard<'a, T> {
710 inner: MutexGuard<'a, CudaExtensionCacheInner>,
711 type_id: TypeId,
712 value: NonNull<T>,
713 _marker: std::marker::PhantomData<&'a T>,
714}
715
716impl<T: 'static> fmt::Debug for CudaExtensionCacheGuard<'_, T> {
717 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
718 let retained_bytes = self
719 .inner
720 .entries
721 .get(&self.type_id)
722 .map(|entry| entry.retained_bytes)
723 .unwrap_or(0);
724 f.debug_struct("CudaExtensionCacheGuard")
725 .field("value_type", &std::any::type_name::<T>())
726 .field("retained_bytes", &retained_bytes)
727 .finish_non_exhaustive()
728 }
729}
730
731impl<T: 'static> Deref for CudaExtensionCacheGuard<'_, T> {
732 type Target = T;
733
734 fn deref(&self) -> &Self::Target {
735 unsafe { self.value.as_ref() }
739 }
740}
741
742impl CudaBackend {
743 fn duplicate_typed<T>(&self, input: &TypedTensor<T>) -> crate::Result<TypedTensor<T>>
744 where
745 T: CubeElement + TensorScalar + CubePrimitive + Clone + Send + Sync + 'static,
746 {
747 self.to_contiguous_view_typed(&input.as_view(), "cast")
751 }
752
753 fn duplicate_bool(
754 &self,
755 input: &TypedTensor<bool>,
756 op: &'static str,
757 ) -> crate::Result<TypedTensor<bool>> {
758 launch_unary_bool_tensor(
759 self.runtime(),
760 input,
761 input.shape(),
762 op,
763 |client, count, dim, out, input_arg| unsafe {
764 structural::copy_bool_kernel::launch_unchecked::<CubeclCudaRuntime>(
765 client,
766 count,
767 dim,
768 out.into_array_arg(),
769 input_arg.into_array_arg(),
770 );
771 },
772 )
773 }
774
775 pub fn new(device_id: CudaDeviceId) -> Result<Self, CudaDeviceError> {
791 Ok(Self {
792 inner: Arc::new(CudaBackendState {
793 cutensor: OnceLock::new(),
794 extension_cache: CudaExtensionCache::new(),
795 cutensor_workspace_max_retained_bytes: AtomicU64::new(
796 DEFAULT_CUTENSOR_WORKSPACE_MAX_RETAINED_BYTES,
797 ),
798 cutensor_workspace_temporary_uses: AtomicU64::new(0),
799 rt: CudaRuntime::new(device_id)?,
800 }),
801 })
802 }
803
804 pub fn runtime(&self) -> &CudaRuntime {
814 &self.inner.rt
815 }
816
817 pub fn device_id(&self) -> CudaDeviceId {
827 self.inner.rt.device_id()
828 }
829
830 pub fn runtime_identity(&self) -> CudaRuntimeIdentity {
844 self.inner.rt.runtime_identity()
845 }
846
847 fn cutensor_handle(&self) -> crate::Result<&ffi::cutensor::CutensorHandle> {
848 if let Some(handle) = self.inner.cutensor.get() {
849 return Ok(handle);
850 }
851 let handle =
852 ffi::cutensor::CutensorHandle::load(self.inner.rt.device_info().compute_capability())?;
853 let _ = self.inner.cutensor.set(handle);
854 self.inner.cutensor.get().ok_or_else(|| {
855 crate::Error::runtime_state(
856 "cuda_cutensor",
857 "cuTENSOR handle initialization completed without a stored handle",
858 )
859 })
860 }
861
862 #[doc(hidden)]
863 pub fn cuda_extension_cache(&self) -> &CudaExtensionCache {
864 &self.inner.extension_cache
865 }
866
867 pub fn clear_cuda_extension_cache(&self) -> crate::Result<()> {
874 self.inner.extension_cache.clear()
875 }
876
877 pub fn cuda_extension_cache_stats(&self) -> crate::Result<CacheStats> {
884 self.inner.extension_cache.stats()
885 }
886
887 pub fn cuda_extension_cache_max_entries(&self) -> crate::Result<NonZeroUsize> {
894 self.inner.extension_cache.max_entries()
895 }
896
897 pub fn cuda_extension_cache_max_retained_bytes(&self) -> crate::Result<NonZeroUsize> {
904 self.inner.extension_cache.max_retained_bytes()
905 }
906
907 pub fn set_cuda_extension_cache_max_entries(
914 &self,
915 max_entries: NonZeroUsize,
916 ) -> crate::Result<()> {
917 self.inner.extension_cache.set_max_entries(max_entries)
918 }
919
920 pub fn set_cuda_extension_cache_max_retained_bytes(
927 &self,
928 max_retained_bytes: NonZeroUsize,
929 ) -> crate::Result<()> {
930 self.inner
931 .extension_cache
932 .set_max_retained_bytes(max_retained_bytes)
933 }
934
935 pub fn cutensor_plan_cache_stats(&self) -> crate::Result<CacheStats> {
946 gemm::cutensor_plan_cache_stats(self)
947 }
948
949 pub fn cutensor_workspace_stats(&self) -> crate::Result<CutensorWorkspaceStats> {
985 gemm::cutensor_workspace_stats(self)
986 }
987
988 pub fn cutensor_workspace_bytes(&self) -> crate::Result<u64> {
1012 Ok(self.cutensor_workspace_stats()?.retained_bytes)
1013 }
1014
1015 pub fn cutensor_workspace_max_retained_bytes(&self) -> u64 {
1041 self.cutensor_workspace_limit()
1042 }
1043
1044 pub fn set_cutensor_workspace_max_retained_bytes(&self, bytes: u64) -> crate::Result<()> {
1090 self.inner
1091 .cutensor_workspace_max_retained_bytes
1092 .store(bytes, Ordering::Relaxed);
1093 gemm::set_cutensor_workspace_max_retained_bytes(self, bytes)
1094 }
1095
1096 fn cutensor_workspace_limit(&self) -> u64 {
1098 self.inner
1099 .cutensor_workspace_max_retained_bytes
1100 .load(Ordering::Relaxed)
1101 }
1102
1103 pub fn cutensor_workspace_temporary_uses(&self) -> u64 {
1136 self.inner
1137 .cutensor_workspace_temporary_uses
1138 .load(Ordering::Relaxed)
1139 }
1140
1141 fn note_cutensor_temporary_workspace(&self) {
1143 self.inner
1144 .cutensor_workspace_temporary_uses
1145 .fetch_add(1, Ordering::Relaxed);
1146 }
1147
1148 pub fn cutensor_workspace_retirement_stats(&self) -> crate::Result<WorkspaceRetirementStats> {
1159 gemm::cutensor_workspace_retirement_stats(self)
1160 }
1161
1162 pub fn cutensor_plan_cache_max_entries(&self) -> crate::Result<NonZeroUsize> {
1167 gemm::cutensor_plan_cache_max_entries(self)
1168 }
1169
1170 pub fn set_cutensor_plan_cache_max_entries(
1178 &self,
1179 max_entries: NonZeroUsize,
1180 ) -> crate::Result<()> {
1181 gemm::set_cutensor_plan_cache_max_entries(self, max_entries)
1182 }
1183
1184 pub fn cutensor_permutation_plan_cache_stats(&self) -> crate::Result<CacheStats> {
1193 permutation::cutensor_permutation_plan_cache_stats(self)
1194 }
1195
1196 pub fn cutensor_permutation_plan_cache_max_entries(&self) -> crate::Result<NonZeroUsize> {
1201 permutation::cutensor_permutation_plan_cache_max_entries(self)
1202 }
1203
1204 pub fn set_cutensor_permutation_plan_cache_max_entries(
1212 &self,
1213 max_entries: NonZeroUsize,
1214 ) -> crate::Result<()> {
1215 permutation::set_cutensor_permutation_plan_cache_max_entries(self, max_entries)
1216 }
1217
1218 fn transpose_typed<T>(
1219 &self,
1220 input: &TypedTensor<T>,
1221 perm: &[usize],
1222 ) -> crate::Result<TypedTensor<T>>
1223 where
1224 T: CubeElement + TensorScalar + CubePrimitive + Clone + Send + Sync + 'static,
1225 {
1226 validate_permutation("transpose", perm, input.shape().len())?;
1227 let output_shape: Vec<usize> = perm.iter().map(|&axis| input.shape()[axis]).collect();
1228 ensure_resident_on_runtime(self.runtime(), input, "transpose")?;
1229 let input_strides =
1230 crate::native_permutation::compact_col_major_strides("transpose", input.shape())?;
1231 let plan = NativePermutationPlan::for_transpose(
1232 "transpose",
1233 input.shape(),
1234 &input_strides,
1235 perm,
1236 0,
1237 input.n_elements(),
1238 input.n_elements(),
1239 false,
1240 )?;
1241 let output = alloc_output::<T>(self.runtime(), &output_shape)?;
1242 let output_arg = typed_tensor_array_arg(&output, "transpose")?;
1243 let input_arg = typed_tensor_array_arg(input, "transpose")?;
1244 launch_native_materialization::<T>(self, output_arg, input_arg, &plan, "transpose")?;
1245 Ok(output)
1246 }
1247
1248 fn transpose_bool(
1249 &self,
1250 input: &TypedTensor<bool>,
1251 perm: &[usize],
1252 ) -> crate::Result<TypedTensor<bool>> {
1253 validate_permutation("transpose", perm, input.shape().len())?;
1254 let output_shape: Vec<usize> = perm.iter().map(|&axis| input.shape()[axis]).collect();
1255 ensure_resident_on_runtime(self.runtime(), input, "transpose")?;
1256 let input_strides =
1257 crate::native_permutation::compact_col_major_strides("transpose", input.shape())?;
1258 let plan = NativePermutationPlan::for_transpose(
1259 "transpose",
1260 input.shape(),
1261 &input_strides,
1262 perm,
1263 0,
1264 input.n_elements(),
1265 input.n_elements(),
1266 false,
1267 )?;
1268 let output = alloc_bool_output(self.runtime(), &output_shape)?;
1269 let output_arg = bool_tensor_array_arg(&output, "transpose")?;
1270 let input_arg = bool_tensor_array_arg(input, "transpose")?;
1271 launch_native_materialization::<u8>(self, output_arg, input_arg, &plan, "transpose")?;
1272 Ok(output)
1273 }
1274
1275 fn broadcast_typed<T>(
1276 &self,
1277 input: &TypedTensor<T>,
1278 shape: &[usize],
1279 dims: &[usize],
1280 ) -> crate::Result<TypedTensor<T>>
1281 where
1282 T: CubeElement + TensorScalar + CubePrimitive + Clone,
1283 {
1284 validate_broadcast_in_dim(input.shape(), shape, dims)?;
1285 launch_unary_tensor(
1286 self.runtime(),
1287 input,
1288 shape,
1289 "broadcast_in_dim",
1290 |client, count, dim, out, input_arg| unsafe {
1291 structural::broadcast_in_dim_kernel::launch_unchecked::<T, CubeclCudaRuntime>(
1292 client,
1293 count,
1294 dim,
1295 out.into_tensor_arg(),
1296 input_arg.into_tensor_arg(),
1297 comptime_sequence(dims),
1298 shape.len(),
1299 );
1300 },
1301 )
1302 }
1303
1304 fn broadcast_bool(
1305 &self,
1306 input: &TypedTensor<bool>,
1307 shape: &[usize],
1308 dims: &[usize],
1309 ) -> crate::Result<TypedTensor<bool>> {
1310 validate_broadcast_in_dim(input.shape(), shape, dims)?;
1311 launch_unary_bool_tensor(
1312 self.runtime(),
1313 input,
1314 shape,
1315 "broadcast_in_dim",
1316 |client, count, dim, out, input_arg| unsafe {
1317 structural::broadcast_in_dim_kernel::launch_unchecked::<u8, CubeclCudaRuntime>(
1318 client,
1319 count,
1320 dim,
1321 out.into_tensor_arg(),
1322 input_arg.into_tensor_arg(),
1323 comptime_sequence(dims),
1324 shape.len(),
1325 );
1326 },
1327 )
1328 }
1329
1330 fn reverse_typed<T>(
1331 &self,
1332 input: &TypedTensor<T>,
1333 axes: &[usize],
1334 ) -> crate::Result<TypedTensor<T>>
1335 where
1336 T: CubeElement + TensorScalar + CubePrimitive + Clone,
1337 {
1338 ensure_axes_unique("reverse", "axes", axes, input.shape().len())?;
1339 launch_unary_tensor(
1340 self.runtime(),
1341 input,
1342 input.shape(),
1343 "reverse",
1344 |client, count, dim, out, input_arg| unsafe {
1345 structural::reverse_kernel::launch_unchecked::<T, CubeclCudaRuntime>(
1346 client,
1347 count,
1348 dim,
1349 out.into_tensor_arg(),
1350 input_arg.into_tensor_arg(),
1351 comptime_sequence(axes),
1352 input.shape().len(),
1353 );
1354 },
1355 )
1356 }
1357
1358 fn reverse_bool(
1359 &self,
1360 input: &TypedTensor<bool>,
1361 axes: &[usize],
1362 ) -> crate::Result<TypedTensor<bool>> {
1363 ensure_axes_unique("reverse", "axes", axes, input.shape().len())?;
1364 launch_unary_bool_tensor(
1365 self.runtime(),
1366 input,
1367 input.shape(),
1368 "reverse",
1369 |client, count, dim, out, input_arg| unsafe {
1370 structural::reverse_kernel::launch_unchecked::<u8, CubeclCudaRuntime>(
1371 client,
1372 count,
1373 dim,
1374 out.into_tensor_arg(),
1375 input_arg.into_tensor_arg(),
1376 comptime_sequence(axes),
1377 input.shape().len(),
1378 );
1379 },
1380 )
1381 }
1382
1383 fn alloc_ranked_output<T, R>(
1384 &self,
1385 shape: &[usize],
1386 op: &'static str,
1387 ) -> crate::Result<TypedTensor<T, R>>
1388 where
1389 T: CubeElement + TensorScalar + Clone + Send + Sync + 'static,
1390 R: TensorRank,
1391 {
1392 let len = checked_dim_product(op, "output shape", shape)?;
1393 let bytes = len.checked_mul(core::mem::size_of::<T>()).ok_or_else(|| {
1394 crate::Error::invalid_argument(
1395 op,
1396 "shape",
1397 format!("CubeCL output byte length overflow for shape {shape:?}"),
1398 )
1399 })?;
1400 let handle = self.runtime().client().empty(bytes);
1401 let shape = R::shape_from_vec(shape.to_vec().into())
1402 .map_err(|err| crate::Error::validation(op, err))?;
1403 TypedTensor::from_buffer_col_major(
1404 shape,
1405 StorageBuffer::Backend(Box::new(crate::CubeclBuffer::new(
1406 handle,
1407 bytes,
1408 self.runtime().device_ordinal(),
1409 self.runtime().allocation_domain_id(),
1410 ))),
1411 Placement {
1412 memory_kind: MemoryKind::Device,
1413 device: Some(DeviceId {
1414 kind: DeviceKind::Gpu(GpuBackendKind::Cuda),
1415 ordinal: self.runtime().device_ordinal(),
1416 }),
1417 cpu_affinity: None,
1418 },
1419 )
1420 }
1421
1422 fn to_contiguous_view_typed<T, R>(
1423 &self,
1424 view: &TypedTensorView<'_, T, R>,
1425 op: &'static str,
1426 ) -> crate::Result<TypedTensor<T, R>>
1427 where
1428 T: CubeElement + TensorScalar + CubePrimitive + Clone + Send + Sync + 'static,
1429 R: TensorRank,
1430 {
1431 ensure_view_resident_on_runtime(self.runtime(), view, op)?;
1432 let len = checked_dim_product(op, "output shape", view.shape())?;
1433 let source_allocation_len = view
1434 .backend_buffer()
1435 .map(|buffer| buffer.len())
1436 .ok_or_else(|| {
1437 crate::Error::runtime_state(op, "expected CUDA backend view, got host view")
1438 })?;
1439 let plan = NativePermutationPlan::for_contiguous_output(
1440 op,
1441 view.shape(),
1442 view.strides(),
1443 view.offset(),
1444 source_allocation_len,
1445 len,
1446 false,
1447 )?;
1448 let output = self.alloc_ranked_output::<T, R>(view.shape(), op)?;
1449 let output_arg = typed_tensor_array_arg(&output, op)?;
1450 let input_arg = typed_view_array_arg(view, op)?;
1451 launch_native_materialization::<T>(self, output_arg, input_arg, &plan, op)?;
1452 Ok(output)
1453 }
1454
1455 fn to_contiguous_view_bool<R: TensorRank>(
1458 &self,
1459 view: &TypedTensorView<'_, bool, R>,
1460 op: &'static str,
1461 ) -> crate::Result<TypedTensor<bool>> {
1462 ensure_view_resident_on_runtime(self.runtime(), view, op)?;
1463 let len = checked_dim_product(op, "output shape", view.shape())?;
1464 let source_allocation_len = view
1465 .backend_buffer()
1466 .map(|buffer| buffer.len())
1467 .ok_or_else(|| {
1468 crate::Error::runtime_state(op, "expected CUDA backend view, got host view")
1469 })?;
1470 let plan = NativePermutationPlan::for_contiguous_output(
1471 op,
1472 view.shape(),
1473 view.strides(),
1474 view.offset(),
1475 source_allocation_len,
1476 len,
1477 false,
1478 )?;
1479 let output = alloc_bool_output(self.runtime(), view.shape())?;
1480 let output_arg = bool_tensor_array_arg(&output, op)?;
1481 let input_arg = bool_view_array_arg(view, op)?;
1482 launch_native_materialization::<u8>(self, output_arg, input_arg, &plan, op)?;
1483 Ok(output)
1484 }
1485
1486 fn to_contiguous_view_cutensor_or_cubecl<T, R>(
1487 &self,
1488 view: &TypedTensorView<'_, T, R>,
1489 op: &'static str,
1490 ) -> crate::Result<TypedTensor<T, R>>
1491 where
1492 T: permutation::CutensorPermutationScalar,
1493 R: TensorRank,
1494 {
1495 if view.strides().iter().any(|&stride| stride <= 0) {
1496 return self.to_contiguous_view_typed(view, op);
1501 }
1502 permutation::to_contiguous_view(self, view, op)
1503 }
1504
1505 fn copy_view_to_view_typed<T, R>(
1506 &self,
1507 src: &TypedTensorView<'_, T, R>,
1508 dst: &mut TypedTensorViewMut<'_, T, R>,
1509 op: &'static str,
1510 ) -> crate::Result<()>
1511 where
1512 T: CubeElement + TensorScalar + CubePrimitive + Clone + Send + Sync + 'static,
1513 R: TensorRank,
1514 {
1515 ensure_view_resident_on_runtime(self.runtime(), src, op)?;
1516 ensure_view_mut_resident_on_runtime(self.runtime(), dst, op)?;
1517 if src.shape() != dst.shape() {
1518 return Err(crate::Error::shape_mismatch(
1519 op,
1520 src.shape().to_vec(),
1521 dst.shape().to_vec(),
1522 ));
1523 }
1524 let source_buffer = src.backend_buffer().ok_or_else(|| {
1525 crate::Error::runtime_state(
1526 op,
1527 "CUDA backend expected a GPU source view; call upload_tensor() first",
1528 )
1529 })?;
1530 let destination_buffer = dst.backend_buffer().ok_or_else(|| {
1531 crate::Error::runtime_state(
1532 op,
1533 "CUDA backend expected a GPU destination view; call upload_tensor() first",
1534 )
1535 })?;
1536 if std::ptr::eq(source_buffer, destination_buffer) {
1537 return Err(crate::Error::invalid_argument(
1538 op,
1539 "source/destination",
1540 "CUDA copy_into source and destination allocations must not alias",
1541 ));
1542 }
1543 let len = src.n_elements();
1544 if len == 0 {
1545 return Ok(());
1546 }
1547 let source_allocation_len = source_buffer.len();
1548 let destination_allocation_len = destination_buffer.len();
1549 if let Some(plan) = self.transpose_copy_plan(src, dst, op)? {
1550 let dst_arg = typed_view_mut_array_arg(dst, op)?;
1551 let src_arg = typed_view_array_arg(src, op)?;
1552 return launch_native_materialization::<T>(self, dst_arg, src_arg, &plan, op);
1553 }
1554 let plan = NativeStridedCopyPlan::new(
1559 op,
1560 src.shape(),
1561 src.strides(),
1562 src.offset(),
1563 source_allocation_len,
1564 dst.strides(),
1565 dst.offset(),
1566 destination_allocation_len,
1567 false,
1568 )?;
1569 if dst.offset() == 0 && self.launch_tiled_transpose(&plan, src, dst, op)? {
1576 return Ok(());
1577 }
1578 if src.offset() == 0 && src.is_col_major_contiguous()? {
1579 let strides = view_strides_i64(dst.strides(), op)?;
1580 let base_offset = view_offset_i64(dst.offset(), op)?;
1581 let src_arg = typed_view_binding(src, op)?;
1582 let dst_arg = typed_view_mut_array_arg(dst, op)?;
1583 let rank = dst.shape().len();
1584 unsafe {
1585 structural::contiguous_to_view_kernel::launch_unchecked::<T, CubeclCudaRuntime>(
1592 self.runtime().client(),
1593 cube_count_for_len(len)?,
1594 cube_dim_1d(),
1595 dst_arg,
1596 src_arg.into_tensor_arg(),
1597 runtime_sequence(&strides),
1598 base_offset,
1599 rank,
1600 );
1601 }
1602 return Ok(());
1603 }
1604 let src_strides = view_strides_i64(&plan.src_strides, op)?;
1605 let dst_strides = view_strides_i64(&plan.dst_strides, op)?;
1606 let src_offset = view_offset_i64(plan.src_offset, op)?;
1607 let dst_offset = view_offset_i64(plan.dst_offset, op)?;
1608 let rank = plan.dims.len();
1609 let src_arg = typed_view_array_arg(src, op)?;
1610 let dst_arg = typed_view_mut_array_arg(dst, op)?;
1611 unsafe {
1612 structural::strided_to_strided_kernel::launch_unchecked::<T, CubeclCudaRuntime>(
1619 self.runtime().client(),
1620 cube_count_for_len(plan.len)?,
1621 cube_dim_1d(),
1622 dst_arg,
1623 src_arg,
1624 runtime_sequence(&plan.dims),
1625 runtime_sequence(&src_strides),
1626 runtime_sequence(&dst_strides),
1627 src_offset,
1628 dst_offset,
1629 plan.len,
1630 rank,
1631 );
1632 }
1633 Ok(())
1634 }
1635
1636 fn launch_tiled_transpose<T, R>(
1647 &self,
1648 plan: &NativeStridedCopyPlan,
1649 src: &TypedTensorView<'_, T, R>,
1650 dst: &mut TypedTensorViewMut<'_, T, R>,
1651 op: &'static str,
1652 ) -> crate::Result<bool>
1653 where
1654 T: CubeElement + TensorScalar + CubePrimitive + Clone + Send + Sync + 'static,
1655 R: TensorRank,
1656 {
1657 let Some((dst_fast_extent, src_fast_extent)) = plan.tiled_transpose_matrix() else {
1658 return Ok(false);
1659 };
1660 let Some(config) = NativeTransposeTile::selected(op)? else {
1661 return Ok(false);
1662 };
1663 const MAX_SHARED_BYTES: usize = 48 * 1024;
1669 let element_bytes = std::mem::size_of::<T>();
1670 let mut launched = None;
1671 for tile in [config.tile, 32] {
1672 let candidate = config.with_tile(tile);
1673 if candidate.shared_bytes(element_bytes) > MAX_SHARED_BYTES {
1674 continue;
1675 }
1676 if let Some(grid) =
1677 candidate.dispatch_grid(op, dst_fast_extent, src_fast_extent, 1, 65_535)?
1678 {
1679 launched = Some((candidate, grid));
1680 break;
1681 }
1682 }
1683 let Some((config, (cubes_x, cubes_y, cubes_z))) = launched else {
1684 return Ok(false);
1685 };
1686 let batch_stride = dst_fast_extent
1687 .checked_mul(src_fast_extent)
1688 .ok_or_else(|| {
1689 crate::Error::invalid_argument(
1690 op,
1691 "shape",
1692 "tiled transpose matrix extent overflows usize",
1693 )
1694 })?;
1695 let src_offset = usize::try_from(plan.src_offset).map_err(|_| {
1696 crate::Error::invalid_argument(
1697 op,
1698 "offset",
1699 "tiled transpose requires a non-negative source offset",
1700 )
1701 })?;
1702 let dst_arg = typed_view_mut_array_arg(dst, op)?;
1703 let src_arg = typed_view_array_arg(src, op)?;
1704 unsafe {
1705 structural::tiled_transpose_kernel::launch_unchecked::<T, CubeclCudaRuntime>(
1713 self.runtime().client(),
1714 CubeCount::Static(cubes_x, cubes_y, cubes_z),
1715 CubeDim::new_2d(config.tile / config.vector_width, config.block_rows),
1716 dst_arg,
1717 src_arg,
1718 src_offset,
1719 batch_stride,
1720 dst_fast_extent,
1721 src_fast_extent,
1722 config.tile as usize,
1723 config.block_rows as usize,
1724 config.padding as usize,
1725 config.vector_width as usize,
1726 );
1727 }
1728 Ok(true)
1729 }
1730
1731 fn transpose_copy_plan<T, R>(
1750 &self,
1751 src: &TypedTensorView<'_, T, R>,
1752 dst: &TypedTensorViewMut<'_, T, R>,
1753 op: &'static str,
1754 ) -> crate::Result<Option<NativePermutationPlan>>
1755 where
1756 T: CubeElement + TensorScalar + CubePrimitive + Clone + Send + Sync + 'static,
1757 R: TensorRank,
1758 {
1759 let shape = dst.shape();
1760 if !(2..=3).contains(&shape.len()) || dst.offset() != 0 {
1761 return Ok(None);
1762 }
1763 let mut row_major_compact = vec![1isize; shape.len()];
1767 for axis in (0..shape.len() - 1).rev() {
1768 let extent = isize::try_from(shape[axis + 1]).map_err(|_| {
1769 crate::Error::invalid_argument(
1770 op,
1771 "shape",
1772 "row-major stride extent exceeds the isize metadata limit",
1773 )
1774 })?;
1775 row_major_compact[axis] =
1776 row_major_compact[axis + 1]
1777 .checked_mul(extent)
1778 .ok_or_else(|| {
1779 crate::Error::invalid_argument(
1780 op,
1781 "shape",
1782 "row-major stride product overflow in the copy destination",
1783 )
1784 })?;
1785 }
1786 if dst.strides() != row_major_compact {
1787 return Ok(None);
1788 }
1789 let source_allocation_len = src
1790 .backend_buffer()
1791 .map(|buffer| buffer.len())
1792 .ok_or_else(|| crate::Error::runtime_state(op, "expected a CUDA source view"))?;
1793 let destination_allocation_len = dst
1794 .backend_buffer()
1795 .map(|buffer| buffer.len())
1796 .ok_or_else(|| crate::Error::runtime_state(op, "expected a CUDA destination view"))?;
1797 let permutation: Vec<usize> = (0..shape.len()).rev().collect();
1798 let plan = NativePermutationPlan::for_transpose(
1799 op,
1800 src.shape(),
1801 src.strides(),
1802 &permutation,
1803 src.offset(),
1804 source_allocation_len,
1805 destination_allocation_len,
1806 false,
1807 )?;
1808 if plan.kind != NativePermutationKind::TiledTranspose {
1809 return Ok(None);
1810 }
1811 Ok(Some(plan))
1812 }
1813
1814 fn copy_view_to_view_cutensor_or_cubecl<T, R>(
1827 &self,
1828 src: &TypedTensorView<'_, T, R>,
1829 dst: &mut TypedTensorViewMut<'_, T, R>,
1830 op: &'static str,
1831 ) -> crate::Result<()>
1832 where
1833 T: permutation::CutensorPermutationScalar,
1834 R: TensorRank,
1835 {
1836 if dst.strides().iter().any(|&stride| stride < 0)
1841 || src.strides().iter().any(|&stride| stride < 1)
1842 {
1843 return self.copy_view_to_view_typed(src, dst, op);
1844 }
1845 if T::REAL_VIEW
1852 && dst.offset() == 0
1853 && NativeStridedCopyPlan::new(
1854 op,
1855 src.shape(),
1856 src.strides(),
1857 src.offset(),
1858 src.backend_buffer().map_or(0, |buffer| buffer.len()),
1859 dst.strides(),
1860 dst.offset(),
1861 dst.backend_buffer().map_or(0, |buffer| buffer.len()),
1862 false,
1863 )?
1864 .tiled_transpose_matrix()
1865 .is_some()
1866 {
1867 return self.copy_view_to_view_typed(src, dst, op);
1868 }
1869 permutation::copy_view_into(self, src, dst, op)
1870 }
1871
1872 fn convert_float_to_float<In, Out>(
1873 &self,
1874 input: &TypedTensor<In>,
1875 ) -> crate::Result<TypedTensor<Out>>
1876 where
1877 In: CubeElement + TensorScalar + CubeFloat + Clone,
1878 Out: CubeElement + TensorScalar + CubeFloat + Clone,
1879 {
1880 launch_unary(
1881 self.runtime(),
1882 input,
1883 input.shape(),
1884 "convert",
1885 |client, count, dim, out, input_arg| unsafe {
1886 structural::convert_float_to_float::launch_unchecked::<Out, In, CubeclCudaRuntime>(
1887 client, count, dim, out, input_arg,
1888 );
1889 },
1890 )
1891 }
1892
1893 fn convert_numeric<In, Out>(&self, input: &TypedTensor<In>) -> crate::Result<TypedTensor<Out>>
1894 where
1895 In: CubeElement + TensorScalar + CubeNumeric + Clone,
1896 Out: CubeElement + TensorScalar + CubeNumeric + Clone,
1897 {
1898 self.launch_cast_unary(input, |client, count, dim, out, input| unsafe {
1899 structural::convert_numeric::launch_unchecked::<Out, In, CubeclCudaRuntime>(
1900 client, count, dim, out, input,
1901 );
1902 })
1903 }
1904
1905 fn launch_cast_unary<In, Out>(
1906 &self,
1907 input: &TypedTensor<In>,
1908 launch: impl FnOnce(
1909 &ComputeClient<CubeclCudaRuntime>,
1910 CubeCount,
1911 CubeDim,
1912 ArrayArg<CubeclCudaRuntime>,
1913 ArrayArg<CubeclCudaRuntime>,
1914 ),
1915 ) -> crate::Result<TypedTensor<Out>>
1916 where
1917 In: CubeElement + TensorScalar + Clone,
1918 Out: CubeElement + TensorScalar + Clone,
1919 {
1920 ensure_resident_on_runtime(self.runtime(), input, "cast")?;
1921 let input_arg = typed_tensor_array_arg(input, "cast")?;
1922 let n = input.n_elements();
1923 let count = if n == 0 {
1924 None
1925 } else {
1926 Some(cube_count_for_len(n)?)
1927 };
1928 let output = alloc_output::<Out>(self.runtime(), input.shape())?;
1929 let Some(count) = count else {
1930 return Ok(output);
1931 };
1932 let output_arg = typed_tensor_array_arg(&output, "cast")?;
1933 launch(
1934 self.runtime().client(),
1935 count,
1936 cube_dim_1d(),
1937 output_arg,
1938 input_arg,
1939 );
1940 Ok(output)
1941 }
1942
1943 fn convert_numeric_to_bool<In>(
1944 &self,
1945 input: &TypedTensor<In>,
1946 ) -> crate::Result<TypedTensor<bool>>
1947 where
1948 In: CubeElement + TensorScalar + CubeNumeric + Clone,
1949 {
1950 ensure_resident_on_runtime(self.runtime(), input, "cast")?;
1951 let input_arg = typed_tensor_array_arg(input, "cast")?;
1952 let n = input.n_elements();
1953 let count = if n == 0 {
1954 None
1955 } else {
1956 Some(cube_count_for_len(n)?)
1957 };
1958 let output = alloc_bool_output(self.runtime(), input.shape())?;
1959 let Some(count) = count else {
1960 return Ok(output);
1961 };
1962 let output_arg = bool_tensor_array_arg(&output, "cast")?;
1963 unsafe {
1964 structural::convert_numeric_to_bool::launch_unchecked::<In, CubeclCudaRuntime>(
1965 self.runtime().client(),
1966 count,
1967 cube_dim_1d(),
1968 output_arg,
1969 input_arg,
1970 );
1971 }
1972 Ok(output)
1973 }
1974
1975 fn convert_bool_to_numeric<Out>(
1976 &self,
1977 input: &TypedTensor<bool>,
1978 ) -> crate::Result<TypedTensor<Out>>
1979 where
1980 Out: CubeElement + TensorScalar + CubeNumeric + Clone,
1981 {
1982 ensure_resident_on_runtime(self.runtime(), input, "cast")?;
1983 let input_arg = bool_tensor_array_arg(input, "cast")?;
1984 let n = input.n_elements();
1985 let count = if n == 0 {
1986 None
1987 } else {
1988 Some(cube_count_for_len(n)?)
1989 };
1990 let output = alloc_output::<Out>(self.runtime(), input.shape())?;
1991 let Some(count) = count else {
1992 return Ok(output);
1993 };
1994 let output_arg = typed_tensor_array_arg(&output, "cast")?;
1995 unsafe {
1996 structural::convert_bool_to_numeric::launch_unchecked::<Out, CubeclCudaRuntime>(
1997 self.runtime().client(),
1998 count,
1999 cube_dim_1d(),
2000 output_arg,
2001 input_arg,
2002 );
2003 }
2004 Ok(output)
2005 }
2006
2007 fn convert_numeric_to_complex<In, OutComplex, OutFloat>(
2008 &self,
2009 input: &TypedTensor<In>,
2010 ) -> crate::Result<TypedTensor<OutComplex>>
2011 where
2012 In: CubeElement + TensorScalar + CubeNumeric + Clone,
2013 OutComplex: CubeElement + TensorScalar + Clone,
2014 OutFloat: CubeElement + CubeFloat + Clone,
2015 {
2016 self.convert_float_to_complex_raw::<In, OutComplex, OutFloat>(
2017 input,
2018 |client, out, input, count| {
2019 unsafe {
2020 structural::convert_numeric_to_complex_raw::launch_unchecked::<
2021 OutFloat,
2022 In,
2023 CubeclCudaRuntime,
2024 >(client, count, cube_dim_1d(), out, input);
2025 }
2026 Ok(())
2027 },
2028 )
2029 }
2030
2031 fn convert_bool_to_complex<OutComplex, OutFloat>(
2032 &self,
2033 input: &TypedTensor<bool>,
2034 ) -> crate::Result<TypedTensor<OutComplex>>
2035 where
2036 OutComplex: CubeElement + TensorScalar + Clone,
2037 OutFloat: CubeElement + CubeFloat + Clone,
2038 {
2039 ensure_resident_on_runtime(self.runtime(), input, "cast")?;
2040 let n = input.n_elements();
2041 let part_len = n.checked_mul(2).ok_or_else(|| {
2042 crate::Error::invalid_argument("cast", "shape", "complex output part length overflow")
2043 })?;
2044 let input_arg = bool_tensor_array_arg(input, "cast")?;
2045 let count = if n == 0 {
2046 None
2047 } else {
2048 Some(cube_count_for_len(n)?)
2049 };
2050 let output = alloc_output::<OutComplex>(self.runtime(), input.shape())?;
2051 let Some(count) = count else {
2052 return Ok(output);
2053 };
2054 let out = typed_tensor_array_arg_as::<OutComplex, OutFloat>(&output, part_len, "cast")?;
2055 unsafe {
2056 structural::convert_bool_to_complex_raw::launch_unchecked::<OutFloat, CubeclCudaRuntime>(
2057 self.runtime().client(),
2058 count,
2059 cube_dim_1d(),
2060 out,
2061 input_arg,
2062 );
2063 }
2064 Ok(output)
2065 }
2066
2067 fn convert_complex_to_numeric<In, Out>(
2068 &self,
2069 input: &TypedTensor<In>,
2070 ) -> crate::Result<TypedTensor<Out>>
2071 where
2072 In: CubeElement + TensorScalar + CubeComplex + Clone,
2073 Out: CubeElement + TensorScalar + CubeNumeric + Clone,
2074 {
2075 self.launch_cast_unary(input, |client, count, dim, out, input| unsafe {
2076 structural::convert_complex_to_numeric::launch_unchecked::<Out, In, CubeclCudaRuntime>(
2077 client, count, dim, out, input,
2078 );
2079 })
2080 }
2081
2082 fn convert_complex_to_bool<In, F>(
2083 &self,
2084 input: &TypedTensor<In>,
2085 ) -> crate::Result<TypedTensor<bool>>
2086 where
2087 In: CubeElement + TensorScalar + CubeComplex<FloatElem = F> + Clone,
2088 F: CubeElement + TensorScalar + CubeFloat,
2089 {
2090 ensure_resident_on_runtime(self.runtime(), input, "cast")?;
2091 let part_len = input.n_elements().checked_mul(2).ok_or_else(|| {
2092 crate::Error::invalid_argument("cast", "shape", "complex input part length overflow")
2093 })?;
2094 let input_arg = typed_tensor_array_arg_as::<In, F>(input, part_len, "cast")?;
2095 let n = input.n_elements();
2096 let count = if n == 0 {
2097 None
2098 } else {
2099 Some(cube_count_for_len(n)?)
2100 };
2101 let output = alloc_bool_output(self.runtime(), input.shape())?;
2102 let Some(count) = count else {
2103 return Ok(output);
2104 };
2105 let output_arg = bool_tensor_array_arg(&output, "cast")?;
2106 unsafe {
2107 structural::convert_complex_raw_to_bool::launch_unchecked::<F, CubeclCudaRuntime>(
2108 self.runtime().client(),
2109 count,
2110 cube_dim_1d(),
2111 output_arg,
2112 input_arg,
2113 );
2114 }
2115 Ok(output)
2116 }
2117
2118 fn convert_f32_to_c32(
2119 &self,
2120 input: &TypedTensor<f32>,
2121 ) -> crate::Result<TypedTensor<Complex32>> {
2122 self.convert_float_to_complex_raw::<f32, Complex32, f32>(
2123 input,
2124 |client, out, input, count| {
2125 unsafe {
2126 structural::convert_f32_to_c32_raw::launch_unchecked::<CubeclCudaRuntime>(
2131 client,
2132 count,
2133 cube_dim_1d(),
2134 out,
2135 input,
2136 );
2137 }
2138 Ok(())
2139 },
2140 )
2141 }
2142
2143 fn convert_f32_to_c64(
2144 &self,
2145 input: &TypedTensor<f32>,
2146 ) -> crate::Result<TypedTensor<Complex64>> {
2147 self.convert_float_to_complex_raw::<f32, Complex64, f64>(
2148 input,
2149 |client, out, input, count| {
2150 unsafe {
2151 structural::convert_f32_to_c64_raw::launch_unchecked::<CubeclCudaRuntime>(
2156 client,
2157 count,
2158 cube_dim_1d(),
2159 out,
2160 input,
2161 );
2162 }
2163 Ok(())
2164 },
2165 )
2166 }
2167
2168 fn convert_f64_to_c32(
2169 &self,
2170 input: &TypedTensor<f64>,
2171 ) -> crate::Result<TypedTensor<Complex32>> {
2172 self.convert_float_to_complex_raw::<f64, Complex32, f32>(
2173 input,
2174 |client, out, input, count| {
2175 unsafe {
2176 structural::convert_f64_to_c32_raw::launch_unchecked::<CubeclCudaRuntime>(
2181 client,
2182 count,
2183 cube_dim_1d(),
2184 out,
2185 input,
2186 );
2187 }
2188 Ok(())
2189 },
2190 )
2191 }
2192
2193 fn convert_f64_to_c64(
2194 &self,
2195 input: &TypedTensor<f64>,
2196 ) -> crate::Result<TypedTensor<Complex64>> {
2197 self.convert_float_to_complex_raw::<f64, Complex64, f64>(
2198 input,
2199 |client, out, input, count| {
2200 unsafe {
2201 structural::convert_f64_to_c64_raw::launch_unchecked::<CubeclCudaRuntime>(
2206 client,
2207 count,
2208 cube_dim_1d(),
2209 out,
2210 input,
2211 );
2212 }
2213 Ok(())
2214 },
2215 )
2216 }
2217
2218 fn convert_float_to_complex_raw<InFloat, OutComplex, OutFloat>(
2223 &self,
2224 input: &TypedTensor<InFloat>,
2225 launch: impl FnOnce(
2226 &cubecl::client::ComputeClient<CubeclCudaRuntime>,
2227 ArrayArg<CubeclCudaRuntime>,
2228 ArrayArg<CubeclCudaRuntime>,
2229 CubeCount,
2230 ) -> crate::Result<()>,
2231 ) -> crate::Result<TypedTensor<OutComplex>>
2232 where
2233 InFloat: CubeElement + TensorScalar + Clone,
2234 OutComplex: CubeElement + TensorScalar + Clone,
2235 OutFloat: CubeElement + Clone,
2236 {
2237 ensure_resident_on_runtime(self.runtime(), input, "convert")?;
2238 let input_arg = typed_tensor_array_arg(input, "convert")?;
2239 let n = input.n_elements();
2240 let output_part_len = n.checked_mul(2).ok_or_else(|| {
2241 crate::Error::invalid_argument(
2242 "convert",
2243 "shape",
2244 "complex output part length overflow",
2245 )
2246 })?;
2247 let count = if n == 0 {
2248 None
2249 } else {
2250 Some(cube_count_for_len(n)?)
2251 };
2252 let output = alloc_output::<OutComplex>(self.runtime(), input.shape())?;
2253 let Some(count) = count else {
2254 return Ok(output);
2255 };
2256 let output_parts =
2257 typed_tensor_array_arg_as::<OutComplex, OutFloat>(&output, output_part_len, "convert")?;
2258 launch(self.runtime().client(), output_parts, input_arg, count)?;
2262 Ok(output)
2263 }
2264
2265 fn convert_c32_to_f32(
2266 &self,
2267 input: &TypedTensor<Complex32>,
2268 ) -> crate::Result<TypedTensor<f32>> {
2269 launch_unary(
2270 self.runtime(),
2271 input,
2272 input.shape(),
2273 "convert",
2274 |client, count, dim, out, input_arg| unsafe {
2275 structural::convert_c32_to_f32::launch_unchecked::<CubeclCudaRuntime>(
2276 client, count, dim, out, input_arg,
2277 );
2278 },
2279 )
2280 }
2281
2282 fn convert_c32_to_f64(
2283 &self,
2284 input: &TypedTensor<Complex32>,
2285 ) -> crate::Result<TypedTensor<f64>> {
2286 launch_unary(
2287 self.runtime(),
2288 input,
2289 input.shape(),
2290 "convert",
2291 |client, count, dim, out, input_arg| unsafe {
2292 structural::convert_c32_to_f64::launch_unchecked::<CubeclCudaRuntime>(
2293 client, count, dim, out, input_arg,
2294 );
2295 },
2296 )
2297 }
2298
2299 fn convert_c64_to_f32(
2300 &self,
2301 input: &TypedTensor<Complex64>,
2302 ) -> crate::Result<TypedTensor<f32>> {
2303 launch_unary(
2304 self.runtime(),
2305 input,
2306 input.shape(),
2307 "convert",
2308 |client, count, dim, out, input_arg| unsafe {
2309 structural::convert_c64_to_f32::launch_unchecked::<CubeclCudaRuntime>(
2310 client, count, dim, out, input_arg,
2311 );
2312 },
2313 )
2314 }
2315
2316 fn convert_c64_to_f64(
2317 &self,
2318 input: &TypedTensor<Complex64>,
2319 ) -> crate::Result<TypedTensor<f64>> {
2320 launch_unary(
2321 self.runtime(),
2322 input,
2323 input.shape(),
2324 "convert",
2325 |client, count, dim, out, input_arg| unsafe {
2326 structural::convert_c64_to_f64::launch_unchecked::<CubeclCudaRuntime>(
2327 client, count, dim, out, input_arg,
2328 );
2329 },
2330 )
2331 }
2332
2333 fn convert_complex_to_complex<In, Out, InFloat, OutFloat>(
2334 &self,
2335 input: &TypedTensor<In>,
2336 ) -> crate::Result<TypedTensor<Out>>
2337 where
2338 In: CubeElement + TensorScalar + CubeComplex + Clone,
2339 Out: CubeElement + TensorScalar + CubeComplex + Clone,
2340 InFloat: CubeElement + CubeFloat + Clone,
2341 OutFloat: CubeElement + CubeFloat + Clone,
2342 {
2343 ensure_resident_on_runtime(self.runtime(), input, "cast")?;
2344 let parts = input.n_elements().checked_mul(2).ok_or_else(|| {
2345 crate::Error::invalid_argument("cast", "shape", "complex component length overflow")
2346 })?;
2347 let input_arg = typed_tensor_array_arg_as::<In, InFloat>(input, parts, "cast")?;
2348 let count = if parts == 0 {
2349 None
2350 } else {
2351 Some(cube_count_for_len(parts)?)
2352 };
2353 let output = alloc_output::<Out>(self.runtime(), input.shape())?;
2354 let Some(count) = count else {
2355 return Ok(output);
2356 };
2357 let output_arg = typed_tensor_array_arg_as::<Out, OutFloat>(&output, parts, "cast")?;
2358 unsafe {
2359 structural::convert_complex_raw::launch_unchecked::<OutFloat, InFloat, CubeclCudaRuntime>(
2360 self.runtime().client(),
2361 count,
2362 cube_dim_1d(),
2363 output_arg,
2364 input_arg,
2365 );
2366 }
2367 Ok(output)
2368 }
2369
2370 fn extract_diagonal_typed<T>(
2371 &self,
2372 input: &TypedTensor<T>,
2373 axis_a: usize,
2374 axis_b: usize,
2375 ) -> crate::Result<TypedTensor<T>>
2376 where
2377 T: CubeElement + TensorScalar + CubePrimitive + Clone,
2378 {
2379 let (output_shape, diag_output_axis) =
2380 extract_diagonal_shape(input.shape(), axis_a, axis_b)?;
2381 launch_unary_tensor(
2382 self.runtime(),
2383 input,
2384 &output_shape,
2385 "extract_diagonal",
2386 |client, count, dim, out, input_arg| unsafe {
2387 diagonal::extract_diagonal_kernel::launch_unchecked::<T, CubeclCudaRuntime>(
2388 client,
2389 count,
2390 dim,
2391 out.into_tensor_arg(),
2392 input_arg.into_tensor_arg(),
2393 axis_a,
2394 axis_b,
2395 diag_output_axis,
2396 input.shape().len(),
2397 output_shape.len(),
2398 );
2399 },
2400 )
2401 }
2402
2403 fn extract_diagonal_bool(
2404 &self,
2405 input: &TypedTensor<bool>,
2406 axis_a: usize,
2407 axis_b: usize,
2408 ) -> crate::Result<TypedTensor<bool>> {
2409 let (output_shape, diag_output_axis) =
2410 extract_diagonal_shape(input.shape(), axis_a, axis_b)?;
2411 launch_unary_bool_tensor(
2412 self.runtime(),
2413 input,
2414 &output_shape,
2415 "extract_diagonal",
2416 |client, count, dim, out, input_arg| unsafe {
2417 diagonal::extract_diagonal_kernel::launch_unchecked::<u8, CubeclCudaRuntime>(
2418 client,
2419 count,
2420 dim,
2421 out.into_tensor_arg(),
2422 input_arg.into_tensor_arg(),
2423 axis_a,
2424 axis_b,
2425 diag_output_axis,
2426 input.shape().len(),
2427 output_shape.len(),
2428 );
2429 },
2430 )
2431 }
2432
2433 fn embed_diagonal_typed<T>(
2434 &self,
2435 input: &TypedTensor<T>,
2436 axis_a: usize,
2437 axis_b: usize,
2438 ) -> crate::Result<TypedTensor<T>>
2439 where
2440 T: CubeElement + TensorScalar + CubePrimitive + Clone,
2441 {
2442 let output_shape = embed_diagonal_shape(input.shape(), axis_a, axis_b)?;
2443 let output = alloc_output::<T>(self.runtime(), &output_shape)?;
2444 launch_nullary_into(
2445 self.runtime(),
2446 &output,
2447 "embed_diagonal",
2448 cube_count_for_len(output.n_elements())?,
2449 cube_dim_1d(),
2450 |client, count, dim, out| unsafe {
2451 structural::fill_zero_kernel::launch_unchecked::<T, CubeclCudaRuntime>(
2452 client, count, dim, out,
2453 );
2454 },
2455 )?;
2456 launch_unary_tensor_into(
2457 self.runtime(),
2458 &output,
2459 input,
2460 "embed_diagonal",
2461 cube_count_for_len(input.n_elements())?,
2462 cube_dim_1d(),
2463 |client, count, dim, out, input_arg| unsafe {
2464 diagonal::embed_diagonal_copy_kernel::launch_unchecked::<T, CubeclCudaRuntime>(
2465 client,
2466 count,
2467 dim,
2468 out.into_tensor_arg(),
2469 input_arg.into_tensor_arg(),
2470 axis_a,
2471 axis_b,
2472 input.shape().len(),
2473 output_shape.len(),
2474 );
2475 },
2476 )?;
2477 Ok(output)
2478 }
2479
2480 fn embed_diagonal_bool(
2481 &self,
2482 input: &TypedTensor<bool>,
2483 axis_a: usize,
2484 axis_b: usize,
2485 ) -> crate::Result<TypedTensor<bool>> {
2486 let output_shape = embed_diagonal_shape(input.shape(), axis_a, axis_b)?;
2487 ensure_resident_on_runtime(self.runtime(), input, "embed_diagonal")?;
2488 typed_tensor_binding(input, "embed_diagonal")?;
2489 let output_len = checked_dim_product("embed_diagonal", "output shape", &output_shape)?;
2490 let output_count = cube_count_for_len(output_len)?;
2491 let input_count = cube_count_for_len(input.n_elements())?;
2492 let output = dispatch::alloc_bool_output(self.runtime(), &output_shape)?;
2493 launch_nullary_bool_into(
2494 self.runtime(),
2495 &output,
2496 "embed_diagonal",
2497 output_count,
2498 cube_dim_1d(),
2499 |client, count, dim, out| unsafe {
2500 structural::fill_zero_kernel::launch_unchecked::<u8, CubeclCudaRuntime>(
2501 client, count, dim, out,
2502 );
2503 },
2504 )?;
2505 launch_bool_tensor_into(
2506 self.runtime(),
2507 &output,
2508 input,
2509 "embed_diagonal",
2510 input_count,
2511 cube_dim_1d(),
2512 |client, count, dim, out, input_arg| unsafe {
2513 diagonal::embed_diagonal_copy_kernel::launch_unchecked::<u8, CubeclCudaRuntime>(
2514 client,
2515 count,
2516 dim,
2517 out.into_tensor_arg(),
2518 input_arg.into_tensor_arg(),
2519 axis_a,
2520 axis_b,
2521 input.shape().len(),
2522 output_shape.len(),
2523 );
2524 },
2525 )?;
2526 Ok(output)
2527 }
2528
2529 pub(crate) fn tril_typed<T>(
2530 &self,
2531 input: &TypedTensor<T>,
2532 k: i64,
2533 ) -> crate::Result<TypedTensor<T>>
2534 where
2535 T: CubeElement + TensorScalar + CubePrimitive + Clone,
2536 {
2537 if input.shape().len() < 2 {
2538 return Err(crate::Error::rank_mismatch("tril", 2, input.shape().len()));
2539 }
2540 launch_unary_tensor(
2541 self.runtime(),
2542 input,
2543 input.shape(),
2544 "tril",
2545 |client, count, dim, out, input_arg| unsafe {
2546 diagonal::tril_kernel::launch_unchecked::<T, CubeclCudaRuntime>(
2547 client,
2548 count,
2549 dim,
2550 out.into_tensor_arg(),
2551 input_arg.into_tensor_arg(),
2552 k,
2553 );
2554 },
2555 )
2556 }
2557
2558 fn tril_bool(&self, input: &TypedTensor<bool>, k: i64) -> crate::Result<TypedTensor<bool>> {
2559 if input.shape().len() < 2 {
2560 return Err(crate::Error::rank_mismatch("tril", 2, input.shape().len()));
2561 }
2562 launch_unary_bool_tensor(
2563 self.runtime(),
2564 input,
2565 input.shape(),
2566 "tril",
2567 |client, count, dim, out, input_arg| unsafe {
2568 diagonal::tril_kernel::launch_unchecked::<u8, CubeclCudaRuntime>(
2569 client,
2570 count,
2571 dim,
2572 out.into_tensor_arg(),
2573 input_arg.into_tensor_arg(),
2574 k,
2575 );
2576 },
2577 )
2578 }
2579
2580 pub(crate) fn triu_typed<T>(
2581 &self,
2582 input: &TypedTensor<T>,
2583 k: i64,
2584 ) -> crate::Result<TypedTensor<T>>
2585 where
2586 T: CubeElement + TensorScalar + CubePrimitive + Clone,
2587 {
2588 if input.shape().len() < 2 {
2589 return Err(crate::Error::rank_mismatch("triu", 2, input.shape().len()));
2590 }
2591 launch_unary_tensor(
2592 self.runtime(),
2593 input,
2594 input.shape(),
2595 "triu",
2596 |client, count, dim, out, input_arg| unsafe {
2597 diagonal::triu_kernel::launch_unchecked::<T, CubeclCudaRuntime>(
2598 client,
2599 count,
2600 dim,
2601 out.into_tensor_arg(),
2602 input_arg.into_tensor_arg(),
2603 k,
2604 );
2605 },
2606 )
2607 }
2608
2609 fn triu_bool(&self, input: &TypedTensor<bool>, k: i64) -> crate::Result<TypedTensor<bool>> {
2610 if input.shape().len() < 2 {
2611 return Err(crate::Error::rank_mismatch("triu", 2, input.shape().len()));
2612 }
2613 launch_unary_bool_tensor(
2614 self.runtime(),
2615 input,
2616 input.shape(),
2617 "triu",
2618 |client, count, dim, out, input_arg| unsafe {
2619 diagonal::triu_kernel::launch_unchecked::<u8, CubeclCudaRuntime>(
2620 client,
2621 count,
2622 dim,
2623 out.into_tensor_arg(),
2624 input_arg.into_tensor_arg(),
2625 k,
2626 );
2627 },
2628 )
2629 }
2630
2631 fn launch_reduce_axis_typed<T>(
2632 &self,
2633 input: &TypedTensor<T>,
2634 axis: usize,
2635 op: &'static str,
2636 launch: impl FnOnce(
2637 &ComputeClient<CubeclCudaRuntime>,
2638 TensorBinding<CubeclCudaRuntime>,
2639 TensorBinding<CubeclCudaRuntime>,
2640 ) -> crate::kernels::Result<()>,
2641 ) -> crate::Result<TypedTensor<T>>
2642 where
2643 T: CubeElement + TensorScalar + Clone,
2644 {
2645 let output_shape = reduction_keepdims_shape(input.shape(), axis);
2646 let input_binding = typed_tensor_binding(input, op)?;
2647 let output = alloc_output::<T>(self.runtime(), &output_shape)?;
2648 if output.n_elements() == 0 {
2649 return Ok(output);
2650 }
2651
2652 let output_binding = typed_tensor_binding(&output, op)?;
2653 launch(self.runtime().client(), input_binding, output_binding)
2654 .map_err(|err| crate::Error::backend_source(op, err))?;
2655 Ok(output)
2656 }
2657
2658 fn reduce_axes_typed<T>(
2659 &self,
2660 input: &TypedTensor<T>,
2661 axes: &[usize],
2662 op: &'static str,
2663 mut launch_axis: impl FnMut(&Self, &TypedTensor<T>, usize) -> crate::Result<TypedTensor<T>>,
2664 ) -> crate::Result<TypedTensor<T>>
2665 where
2666 T: CubeElement
2667 + CubePrimitive
2668 + tenferro_tensor::TensorScalar
2669 + Clone
2670 + Send
2671 + Sync
2672 + 'static,
2673 {
2674 ensure_axes_unique(op, "axes", axes, input.shape().len())?;
2675 if axes.is_empty() {
2676 return self.to_contiguous_view_typed(&input.as_view(), op);
2677 }
2678
2679 let final_shape = reduction_output_shape(input.shape(), axes);
2680 let mut sorted_axes = axes.to_vec();
2681 sorted_axes.sort_unstable();
2682
2683 let (first_axis, remaining_axes) = sorted_axes
2686 .split_first()
2687 .ok_or_else(|| crate::Error::invalid_argument(op, "axes", "axes must not be empty"))?;
2688 let mut current = launch_axis(self, input, *first_axis)?;
2689 for &axis in remaining_axes {
2690 current = launch_axis(self, ¤t, axis)?;
2691 }
2692
2693 cubecl_reshape_metadata(current, final_shape, op)
2694 }
2695
2696 fn reduce_sum_float_typed<
2697 F: CubeElement
2698 + CubePrimitive
2699 + CubeFloat
2700 + tenferro_tensor::TensorScalar
2701 + Clone
2702 + Send
2703 + Sync
2704 + 'static,
2705 >(
2706 &self,
2707 input: &TypedTensor<F>,
2708 axes: &[usize],
2709 ) -> crate::Result<TypedTensor<F>> {
2710 let op = op_name(
2711 PrimitiveOpKind::ReduceSum,
2712 op_descriptor::GpuLaunchKind::Reduction,
2713 )?;
2714 self.reduce_axes_typed(input, axes, op, |backend, current, axis| {
2715 backend.launch_reduce_axis_typed(current, axis, op, |client, input, output| {
2716 cubecl_reduce::launch_sum_float::<CubeclCudaRuntime, F>(
2717 client,
2718 input,
2719 output,
2720 axis,
2721 ReduceStrategy::Auto,
2722 )
2723 })
2724 })
2725 }
2726
2727 fn reduce_sum_squares_float_typed<
2728 F: CubeElement
2729 + CubePrimitive
2730 + CubeFloat
2731 + tenferro_tensor::TensorScalar
2732 + Clone
2733 + Send
2734 + Sync
2735 + 'static,
2736 >(
2737 &self,
2738 input: &TypedTensor<F>,
2739 axes: &[usize],
2740 ) -> crate::Result<TypedTensor<F>> {
2741 let op = op_name(
2742 PrimitiveOpKind::ReduceSumSquares,
2743 op_descriptor::GpuLaunchKind::Reduction,
2744 )?;
2745 ensure_axes_unique(op, "axes", axes, input.shape().len())?;
2746 let final_shape = reduction_output_shape(input.shape(), axes);
2747 let mut sorted_axes = axes.to_vec();
2748 sorted_axes.sort_unstable();
2749 let (&first_axis, remaining_axes) = sorted_axes
2750 .split_first()
2751 .ok_or_else(|| crate::Error::invalid_argument(op, "axes", "axes must not be empty"))?;
2752
2753 let mut current =
2754 self.launch_reduce_axis_typed(input, first_axis, op, |client, input, output| {
2755 cubecl_reduce::launch_sum_squares_float::<CubeclCudaRuntime, F>(
2756 client,
2757 input,
2758 output,
2759 first_axis,
2760 ReduceStrategy::Auto,
2761 )
2762 })?;
2763 for &axis in remaining_axes {
2764 current =
2765 self.launch_reduce_axis_typed(¤t, axis, op, |client, input, output| {
2766 cubecl_reduce::launch_sum_float::<CubeclCudaRuntime, F>(
2767 client,
2768 input,
2769 output,
2770 axis,
2771 ReduceStrategy::Auto,
2772 )
2773 })?;
2774 }
2775
2776 cubecl_reshape_metadata(current, final_shape, op)
2777 }
2778
2779 fn reduce_sum_complex_typed<
2780 C: CubeElement
2781 + CubePrimitive
2782 + CubeComplex
2783 + tenferro_tensor::TensorScalar
2784 + Clone
2785 + Send
2786 + Sync
2787 + 'static,
2788 >(
2789 &self,
2790 input: &TypedTensor<C>,
2791 axes: &[usize],
2792 ) -> crate::Result<TypedTensor<C>> {
2793 let op = op_name(
2794 PrimitiveOpKind::ReduceSum,
2795 op_descriptor::GpuLaunchKind::Reduction,
2796 )?;
2797 self.reduce_axes_typed(input, axes, op, |backend, current, axis| {
2798 backend.launch_reduce_axis_typed(current, axis, op, |client, input, output| {
2799 cubecl_reduce::launch_sum_complex::<CubeclCudaRuntime, C>(
2800 client,
2801 input,
2802 output,
2803 axis,
2804 ReduceStrategy::Auto,
2805 )
2806 })
2807 })
2808 }
2809
2810 fn reduce_sum_int_typed<
2811 I: CubeElement
2812 + CubePrimitive
2813 + CubeInt
2814 + tenferro_tensor::TensorScalar
2815 + Clone
2816 + Send
2817 + Sync
2818 + 'static,
2819 >(
2820 &self,
2821 input: &TypedTensor<I>,
2822 axes: &[usize],
2823 ) -> crate::Result<TypedTensor<I>> {
2824 let op = op_name(
2825 PrimitiveOpKind::ReduceSum,
2826 op_descriptor::GpuLaunchKind::Reduction,
2827 )?;
2828 self.reduce_axes_typed(input, axes, op, |backend, current, axis| {
2829 backend.launch_reduce_axis_typed(current, axis, op, |client, input, output| {
2830 cubecl_reduce::launch_sum_int::<CubeclCudaRuntime, I>(
2831 client,
2832 input,
2833 output,
2834 axis,
2835 ReduceStrategy::Auto,
2836 )
2837 })
2838 })
2839 }
2840
2841 fn reduce_prod_float_typed<
2842 F: CubeElement
2843 + CubePrimitive
2844 + CubeFloat
2845 + tenferro_tensor::TensorScalar
2846 + Clone
2847 + Send
2848 + Sync
2849 + 'static,
2850 >(
2851 &self,
2852 input: &TypedTensor<F>,
2853 axes: &[usize],
2854 ) -> crate::Result<TypedTensor<F>> {
2855 let op = op_name(
2856 PrimitiveOpKind::ReduceProd,
2857 op_descriptor::GpuLaunchKind::Reduction,
2858 )?;
2859 self.reduce_axes_typed(input, axes, op, |backend, current, axis| {
2860 backend.launch_reduce_axis_typed(current, axis, op, |client, input, output| {
2861 cubecl_reduce::launch_prod_float::<CubeclCudaRuntime, F>(
2862 client,
2863 input,
2864 output,
2865 axis,
2866 ReduceStrategy::Auto,
2867 )
2868 })
2869 })
2870 }
2871
2872 fn reduce_prod_complex_typed<
2873 C: CubeElement
2874 + CubePrimitive
2875 + CubeComplex
2876 + tenferro_tensor::TensorScalar
2877 + Clone
2878 + Send
2879 + Sync
2880 + 'static,
2881 >(
2882 &self,
2883 input: &TypedTensor<C>,
2884 axes: &[usize],
2885 ) -> crate::Result<TypedTensor<C>> {
2886 let op = op_name(
2887 PrimitiveOpKind::ReduceProd,
2888 op_descriptor::GpuLaunchKind::Reduction,
2889 )?;
2890 self.reduce_axes_typed(input, axes, op, |backend, current, axis| {
2891 backend.launch_reduce_axis_typed(current, axis, op, |client, input, output| {
2892 cubecl_reduce::launch_prod_complex::<CubeclCudaRuntime, C>(
2893 client,
2894 input,
2895 output,
2896 axis,
2897 ReduceStrategy::Auto,
2898 )
2899 })
2900 })
2901 }
2902
2903 fn reduce_prod_int_typed<
2904 I: CubeElement
2905 + CubePrimitive
2906 + CubeInt
2907 + tenferro_tensor::TensorScalar
2908 + Clone
2909 + Send
2910 + Sync
2911 + 'static,
2912 >(
2913 &self,
2914 input: &TypedTensor<I>,
2915 axes: &[usize],
2916 ) -> crate::Result<TypedTensor<I>> {
2917 let op = op_name(
2918 PrimitiveOpKind::ReduceProd,
2919 op_descriptor::GpuLaunchKind::Reduction,
2920 )?;
2921 self.reduce_axes_typed(input, axes, op, |backend, current, axis| {
2922 backend.launch_reduce_axis_typed(current, axis, op, |client, input, output| {
2923 cubecl_reduce::launch_prod_int::<CubeclCudaRuntime, I>(
2924 client,
2925 input,
2926 output,
2927 axis,
2928 ReduceStrategy::Auto,
2929 )
2930 })
2931 })
2932 }
2933
2934 fn reduce_max_float_typed<
2935 F: CubeElement
2936 + CubePrimitive
2937 + CubeFloat
2938 + tenferro_tensor::TensorScalar
2939 + Clone
2940 + Send
2941 + Sync
2942 + 'static,
2943 >(
2944 &self,
2945 input: &TypedTensor<F>,
2946 axes: &[usize],
2947 ) -> crate::Result<TypedTensor<F>> {
2948 let op = op_name(
2949 PrimitiveOpKind::ReduceMax,
2950 op_descriptor::GpuLaunchKind::Reduction,
2951 )?;
2952 self.reduce_axes_typed(input, axes, op, |backend, current, axis| {
2953 backend.launch_reduce_axis_typed(current, axis, op, |client, input, output| {
2954 cubecl_reduce::launch_max_float::<CubeclCudaRuntime, F>(
2955 client,
2956 input,
2957 output,
2958 axis,
2959 ReduceStrategy::Auto,
2960 )
2961 })
2962 })
2963 }
2964
2965 fn reduce_max_int_typed<
2966 I: CubeElement
2967 + CubePrimitive
2968 + CubeInt
2969 + tenferro_tensor::TensorScalar
2970 + Clone
2971 + Send
2972 + Sync
2973 + 'static,
2974 >(
2975 &self,
2976 input: &TypedTensor<I>,
2977 axes: &[usize],
2978 ) -> crate::Result<TypedTensor<I>> {
2979 let op = op_name(
2980 PrimitiveOpKind::ReduceMax,
2981 op_descriptor::GpuLaunchKind::Reduction,
2982 )?;
2983 self.reduce_axes_typed(input, axes, op, |backend, current, axis| {
2984 backend.launch_reduce_axis_typed(current, axis, op, |client, input, output| {
2985 cubecl_reduce::launch_max_int::<CubeclCudaRuntime, I>(
2986 client,
2987 input,
2988 output,
2989 axis,
2990 ReduceStrategy::Auto,
2991 )
2992 })
2993 })
2994 }
2995
2996 fn reduce_min_float_typed<
2997 F: CubeElement
2998 + CubePrimitive
2999 + CubeFloat
3000 + tenferro_tensor::TensorScalar
3001 + Clone
3002 + Send
3003 + Sync
3004 + 'static,
3005 >(
3006 &self,
3007 input: &TypedTensor<F>,
3008 axes: &[usize],
3009 ) -> crate::Result<TypedTensor<F>> {
3010 let op = op_name(
3011 PrimitiveOpKind::ReduceMin,
3012 op_descriptor::GpuLaunchKind::Reduction,
3013 )?;
3014 self.reduce_axes_typed(input, axes, op, |backend, current, axis| {
3015 backend.launch_reduce_axis_typed(current, axis, op, |client, input, output| {
3016 cubecl_reduce::launch_min_float::<CubeclCudaRuntime, F>(
3017 client,
3018 input,
3019 output,
3020 axis,
3021 ReduceStrategy::Auto,
3022 )
3023 })
3024 })
3025 }
3026
3027 fn reduce_min_int_typed<
3028 I: CubeElement
3029 + CubePrimitive
3030 + CubeInt
3031 + tenferro_tensor::TensorScalar
3032 + Clone
3033 + Send
3034 + Sync
3035 + 'static,
3036 >(
3037 &self,
3038 input: &TypedTensor<I>,
3039 axes: &[usize],
3040 ) -> crate::Result<TypedTensor<I>> {
3041 let op = op_name(
3042 PrimitiveOpKind::ReduceMin,
3043 op_descriptor::GpuLaunchKind::Reduction,
3044 )?;
3045 self.reduce_axes_typed(input, axes, op, |backend, current, axis| {
3046 backend.launch_reduce_axis_typed(current, axis, op, |client, input, output| {
3047 cubecl_reduce::launch_min_int::<CubeclCudaRuntime, I>(
3048 client,
3049 input,
3050 output,
3051 axis,
3052 ReduceStrategy::Auto,
3053 )
3054 })
3055 })
3056 }
3057
3058 pub(crate) fn slice_typed<T>(
3059 &self,
3060 input: &TypedTensor<T>,
3061 config: &SliceConfig,
3062 ) -> crate::Result<TypedTensor<T>>
3063 where
3064 T: CubeElement + TensorScalar + CubePrimitive + Clone,
3065 {
3066 let output_shape = validate_slice(input.shape(), config)?;
3067 launch_unary_tensor(
3068 self.runtime(),
3069 input,
3070 &output_shape,
3071 "slice",
3072 |client, count, dim, out, input_arg| unsafe {
3073 indexing::slice_kernel::launch_unchecked::<T, CubeclCudaRuntime>(
3074 client,
3075 count,
3076 dim,
3077 out.into_tensor_arg(),
3078 input_arg.into_tensor_arg(),
3079 runtime_sequence(&config.starts),
3080 comptime_sequence(&config.strides),
3081 );
3082 },
3083 )
3084 }
3085
3086 fn slice_bool(
3087 &self,
3088 input: &TypedTensor<bool>,
3089 config: &SliceConfig,
3090 ) -> crate::Result<TypedTensor<bool>> {
3091 let output_shape = validate_slice(input.shape(), config)?;
3092 launch_unary_bool_tensor(
3093 self.runtime(),
3094 input,
3095 &output_shape,
3096 "slice",
3097 |client, count, dim, out, input_arg| unsafe {
3098 indexing::slice_kernel::launch_unchecked::<u8, CubeclCudaRuntime>(
3099 client,
3100 count,
3101 dim,
3102 out.into_tensor_arg(),
3103 input_arg.into_tensor_arg(),
3104 runtime_sequence(&config.starts),
3105 comptime_sequence(&config.strides),
3106 );
3107 },
3108 )
3109 }
3110
3111 fn dynamic_slice_typed<T, I>(
3112 &self,
3113 input: &TypedTensor<T>,
3114 starts: &TypedTensor<I>,
3115 slice_sizes: &[usize],
3116 ) -> crate::Result<TypedTensor<T>>
3117 where
3118 T: CubeElement + TensorScalar + CubePrimitive + Clone,
3119 I: CubeElement + TensorScalar + CubePrimitive + CubeNumeric + Clone + CudaIndexValidation,
3120 {
3121 ensure_rank("dynamic_slice", input.shape().len(), slice_sizes.len())?;
3122 ensure_rank("dynamic_slice", 1, starts.shape().len())?;
3123 if starts.shape()[0] != input.shape().len() {
3124 return Err(crate::Error::rank_mismatch(
3125 "dynamic_slice",
3126 input.shape().len(),
3127 starts.shape()[0],
3128 ));
3129 }
3130 for (axis, (&window, &dim)) in slice_sizes.iter().zip(input.shape()).enumerate() {
3131 if window > dim {
3132 return Err(crate::Error::invalid_argument(
3133 "dynamic_slice",
3134 "slice_sizes",
3135 format!("slice size exceeds dimension on axis {axis}"),
3136 ));
3137 }
3138 }
3139 let output_len = checked_dim_product("dynamic_slice", "output shape", slice_sizes)?;
3140 if output_len != 0 {
3141 cube_count_for_len(output_len)?;
3142 }
3143 ensure_resident_on_runtime(self.runtime(), input, "dynamic_slice")?;
3144 typed_tensor_binding(input, "dynamic_slice")?;
3145 ensure_resident_on_runtime(self.runtime(), starts, "dynamic_slice")?;
3146 typed_tensor_binding(starts, "dynamic_slice")?;
3147 I::validate(self, starts)?;
3148 launch_binary_tensor(
3149 self.runtime(),
3150 input,
3151 starts,
3152 slice_sizes,
3153 "dynamic_slice",
3154 |client, count, dim, out, input_arg, starts_arg| unsafe {
3155 indexing::dynamic_slice_kernel::launch_unchecked::<T, I, CubeclCudaRuntime>(
3156 client,
3157 count,
3158 dim,
3159 out.into_tensor_arg(),
3160 input_arg.into_tensor_arg(),
3161 starts_arg.into_tensor_arg(),
3162 runtime_sequence(slice_sizes),
3163 slice_sizes.len(),
3164 );
3165 },
3166 )
3167 }
3168
3169 fn dynamic_slice_bool<I>(
3170 &self,
3171 input: &TypedTensor<bool>,
3172 starts: &TypedTensor<I>,
3173 slice_sizes: &[usize],
3174 ) -> crate::Result<TypedTensor<bool>>
3175 where
3176 I: CubeElement + TensorScalar + CubePrimitive + CubeNumeric + Clone + CudaIndexValidation,
3177 {
3178 ensure_rank("dynamic_slice", input.shape().len(), slice_sizes.len())?;
3179 if starts.shape().len() != 1 {
3180 return Err(crate::Error::invalid_argument(
3181 "dynamic_slice",
3182 "starts",
3183 "starts must be a rank-1 tensor",
3184 ));
3185 }
3186 if starts.shape()[0] != input.shape().len() {
3187 return Err(crate::Error::invalid_argument(
3188 "dynamic_slice",
3189 "starts",
3190 format!(
3191 "starts length {} must match input rank {}",
3192 starts.shape()[0],
3193 input.shape().len()
3194 ),
3195 ));
3196 }
3197 for (axis, (&window, &dim)) in slice_sizes.iter().zip(input.shape()).enumerate() {
3198 if window > dim {
3199 return Err(crate::Error::invalid_argument(
3200 "dynamic_slice",
3201 "slice_sizes",
3202 format!("slice size exceeds dimension on axis {axis}"),
3203 ));
3204 }
3205 }
3206 let output_len = checked_dim_product("dynamic_slice", "output shape", slice_sizes)?;
3207 if output_len != 0 {
3208 cube_count_for_len(output_len)?;
3209 }
3210 ensure_resident_on_runtime(self.runtime(), input, "dynamic_slice")?;
3211 bool_tensor_array_arg(input, "dynamic_slice")?;
3212 ensure_resident_on_runtime(self.runtime(), starts, "dynamic_slice")?;
3213 typed_tensor_binding(starts, "dynamic_slice")?;
3214 I::validate(self, starts)?;
3215 launch_binary_bool_tensor(
3216 self.runtime(),
3217 input,
3218 starts,
3219 slice_sizes,
3220 "dynamic_slice",
3221 |client, count, dim, out, input_arg, starts_arg| unsafe {
3222 indexing::dynamic_slice_kernel::launch_unchecked::<u8, I, CubeclCudaRuntime>(
3223 client,
3224 count,
3225 dim,
3226 out.into_tensor_arg(),
3227 input_arg.into_tensor_arg(),
3228 starts_arg.into_tensor_arg(),
3229 runtime_sequence(slice_sizes),
3230 slice_sizes.len(),
3231 );
3232 },
3233 )
3234 }
3235
3236 fn pad_typed<T>(
3237 &self,
3238 input: &TypedTensor<T>,
3239 config: &PadConfig,
3240 ) -> crate::Result<TypedTensor<T>>
3241 where
3242 T: CubeElement + TensorScalar + CubePrimitive + Clone,
3243 {
3244 let output_shape = pad_output_shape(input.shape(), config)?;
3245 launch_unary_tensor(
3246 self.runtime(),
3247 input,
3248 &output_shape,
3249 "pad",
3250 |client, count, dim, out, input_arg| unsafe {
3251 indexing::pad_kernel::launch_unchecked::<T, CubeclCudaRuntime>(
3252 client,
3253 count,
3254 dim,
3255 out.into_tensor_arg(),
3256 input_arg.into_tensor_arg(),
3257 runtime_sequence(&config.edge_padding_low),
3258 runtime_sequence(&config.interior_padding),
3259 config.edge_padding_low.len(),
3260 );
3261 },
3262 )
3263 }
3264
3265 fn pad_bool(
3266 &self,
3267 input: &TypedTensor<bool>,
3268 config: &PadConfig,
3269 ) -> crate::Result<TypedTensor<bool>> {
3270 let output_shape = pad_output_shape(input.shape(), config)?;
3271 launch_unary_bool_tensor(
3272 self.runtime(),
3273 input,
3274 &output_shape,
3275 "pad",
3276 |client, count, dim, out, input_arg| unsafe {
3277 indexing::pad_kernel::launch_unchecked::<u8, CubeclCudaRuntime>(
3278 client,
3279 count,
3280 dim,
3281 out.into_tensor_arg(),
3282 input_arg.into_tensor_arg(),
3283 runtime_sequence(&config.edge_padding_low),
3284 runtime_sequence(&config.interior_padding),
3285 config.edge_padding_low.len(),
3286 );
3287 },
3288 )
3289 }
3290
3291 fn concatenate_typed<T>(
3292 &self,
3293 inputs: &[&TypedTensor<T>],
3294 axis: usize,
3295 ) -> crate::Result<TypedTensor<T>>
3296 where
3297 T: CubeElement + TensorScalar + CubePrimitive + Clone,
3298 {
3299 let output_shape = concatenate_output_shape(inputs, axis)?;
3300 let output = alloc_output::<T>(self.runtime(), &output_shape)?;
3301 let mut offset = 0usize;
3302 for input in inputs {
3303 launch_unary_tensor_into(
3304 self.runtime(),
3305 &output,
3306 input,
3307 "concatenate",
3308 cube_count_for_len(input.n_elements())?,
3309 cube_dim_1d(),
3310 |client, count, dim, out, input_arg| unsafe {
3311 structural::concatenate_copy_kernel::launch_unchecked::<T, CubeclCudaRuntime>(
3312 client,
3313 count,
3314 dim,
3315 out.into_tensor_arg(),
3316 input_arg.into_tensor_arg(),
3317 axis,
3318 offset,
3319 input.shape().len(),
3320 );
3321 },
3322 )?;
3323 offset += input.shape()[axis];
3326 }
3327 Ok(output)
3328 }
3329
3330 fn concatenate_bool(
3331 &self,
3332 inputs: &[&TypedTensor<bool>],
3333 axis: usize,
3334 ) -> crate::Result<TypedTensor<bool>> {
3335 let output_shape = concatenate_output_shape(inputs, axis)?;
3336 for input in inputs {
3337 ensure_resident_on_runtime(self.runtime(), input, "concatenate")?;
3338 typed_tensor_binding(input, "concatenate")?;
3339 }
3340 checked_dim_product("concatenate", "output shape", &output_shape)?;
3341 let launch_counts = inputs
3342 .iter()
3343 .map(|input| cube_count_for_len(input.n_elements()))
3344 .collect::<crate::Result<Vec<_>>>()?;
3345 let output = dispatch::alloc_bool_output(self.runtime(), &output_shape)?;
3346 let mut offset = 0usize;
3347 for (input, launch_count) in inputs.iter().zip(launch_counts) {
3348 launch_bool_tensor_into(
3349 self.runtime(),
3350 &output,
3351 input,
3352 "concatenate",
3353 launch_count,
3354 cube_dim_1d(),
3355 |client, count, dim, out, input_arg| unsafe {
3356 structural::concatenate_copy_kernel::launch_unchecked::<u8, CubeclCudaRuntime>(
3357 client,
3358 count,
3359 dim,
3360 out.into_tensor_arg(),
3361 input_arg.into_tensor_arg(),
3362 axis,
3363 offset,
3364 input.shape().len(),
3365 );
3366 },
3367 )?;
3368 offset += input.shape()[axis];
3369 }
3370 Ok(output)
3371 }
3372
3373 fn gather_typed<T, I>(
3374 &self,
3375 operand: &TypedTensor<T>,
3376 start_indices: &TypedTensor<I>,
3377 config: &GatherConfig,
3378 ) -> crate::Result<TypedTensor<T>>
3379 where
3380 T: CubeElement + TensorScalar + CubePrimitive + Clone,
3381 I: CubeElement + TensorScalar + CubePrimitive + CubeNumeric + Clone + CudaIndexValidation,
3382 {
3383 let meta = gather_launch_meta(operand.shape(), start_indices.shape(), config)?;
3384 let output_len = checked_dim_product("gather", "output shape", &meta.output_shape)?;
3385 if output_len != 0 {
3386 cube_count_for_len(output_len)?;
3387 }
3388 ensure_resident_on_runtime(self.runtime(), operand, "gather")?;
3389 typed_tensor_binding(operand, "gather")?;
3390 ensure_resident_on_runtime(self.runtime(), start_indices, "gather")?;
3391 typed_tensor_binding(start_indices, "gather")?;
3392 I::validate(self, start_indices)?;
3393 launch_binary_tensor(
3394 self.runtime(),
3395 operand,
3396 start_indices,
3397 &meta.output_shape,
3398 "gather",
3399 |client, count, dim, out, operand_arg, indices_arg| unsafe {
3400 indexing::gather_kernel::launch_unchecked::<T, I, CubeclCudaRuntime>(
3401 client,
3402 count,
3403 dim,
3404 out.into_tensor_arg(),
3405 operand_arg.into_tensor_arg(),
3406 indices_arg.into_tensor_arg(),
3407 comptime_sequence(&meta.window_dims),
3408 comptime_sequence(&config.offset_dims),
3409 comptime_sequence(&config.start_index_map),
3410 runtime_sequence(&config.slice_sizes),
3411 config.index_vector_dim,
3412 operand.shape().len(),
3413 meta.output_shape.len(),
3414 start_indices.shape().len(),
3415 );
3416 },
3417 )
3418 }
3419
3420 fn gather_bool<I>(
3421 &self,
3422 operand: &TypedTensor<bool>,
3423 start_indices: &TypedTensor<I>,
3424 config: &GatherConfig,
3425 ) -> crate::Result<TypedTensor<bool>>
3426 where
3427 I: CubeElement + TensorScalar + CubePrimitive + CubeNumeric + Clone + CudaIndexValidation,
3428 {
3429 let meta = gather_launch_meta(operand.shape(), start_indices.shape(), config)?;
3430 let output_len = checked_dim_product("gather", "output shape", &meta.output_shape)?;
3431 if output_len != 0 {
3432 cube_count_for_len(output_len)?;
3433 }
3434 ensure_resident_on_runtime(self.runtime(), operand, "gather")?;
3435 bool_tensor_array_arg(operand, "gather")?;
3436 ensure_resident_on_runtime(self.runtime(), start_indices, "gather")?;
3437 typed_tensor_binding(start_indices, "gather")?;
3438 I::validate(self, start_indices)?;
3439 launch_binary_bool_tensor(
3440 self.runtime(),
3441 operand,
3442 start_indices,
3443 &meta.output_shape,
3444 "gather",
3445 |client, count, dim, out, operand_arg, indices_arg| unsafe {
3446 indexing::gather_kernel::launch_unchecked::<u8, I, CubeclCudaRuntime>(
3447 client,
3448 count,
3449 dim,
3450 out.into_tensor_arg(),
3451 operand_arg.into_tensor_arg(),
3452 indices_arg.into_tensor_arg(),
3453 comptime_sequence(&meta.window_dims),
3454 comptime_sequence(&config.offset_dims),
3455 comptime_sequence(&config.start_index_map),
3456 runtime_sequence(&config.slice_sizes),
3457 config.index_vector_dim,
3458 operand.shape().len(),
3459 meta.output_shape.len(),
3460 start_indices.shape().len(),
3461 );
3462 },
3463 )
3464 }
3465
3466 fn scatter_float_typed<T, I>(
3467 &self,
3468 operand: &TypedTensor<T>,
3469 scatter_indices: &TypedTensor<I>,
3470 updates: &TypedTensor<T>,
3471 config: &ScatterConfig,
3472 ) -> crate::Result<TypedTensor<T>>
3473 where
3474 T: CubeElement + TensorScalar + CubeFloat + Clone,
3475 I: CubeElement + TensorScalar + CubePrimitive + CubeNumeric + Clone + CudaIndexValidation,
3476 {
3477 let meta = scatter_launch_meta(
3478 operand.shape(),
3479 scatter_indices.shape(),
3480 updates.shape(),
3481 config,
3482 )?;
3483 let update_len = scatter_update_len(&meta)?;
3484 let output_len = checked_dim_product("scatter", "output shape", operand.shape())?;
3485 if output_len != 0 {
3486 cube_count_for_len(output_len)?;
3487 }
3488 if update_len != 0 {
3489 cube_count_for_len(update_len)?;
3490 }
3491 let client = self.runtime().client();
3492 ensure_resident_on_runtime(self.runtime(), operand, "scatter")?;
3493 typed_tensor_binding(operand, "scatter")?;
3494 ensure_resident_on_runtime(self.runtime(), scatter_indices, "scatter")?;
3495 typed_tensor_binding(scatter_indices, "scatter")?;
3496 ensure_resident_on_runtime(self.runtime(), updates, "scatter")?;
3497 typed_tensor_binding(updates, "scatter")?;
3498 ensure_atomic_add_supported::<T>(client, "scatter")?;
3499 I::validate(self, scatter_indices)?;
3500 let output = alloc_output::<T>(self.runtime(), operand.shape())?;
3501 if output.n_elements() == 0 {
3502 return Ok(output);
3503 }
3504
3505 launch_unary_tensor_into(
3506 self.runtime(),
3507 &output,
3508 operand,
3509 "scatter",
3510 cube_count_for_len(output.n_elements())?,
3511 cube_dim_1d(),
3512 |client, count, dim, out_arg, operand_arg| unsafe {
3513 indexing::scatter_copy_kernel::launch_unchecked::<T, CubeclCudaRuntime>(
3514 client,
3515 count,
3516 dim,
3517 out_arg.into_tensor_arg(),
3518 operand_arg.into_tensor_arg(),
3519 );
3520 },
3521 )?;
3522
3523 if update_len == 0 {
3524 return Ok(output);
3525 }
3526 let output_parts =
3527 typed_tensor_array_arg_as::<T, T>(&output, output.n_elements(), "scatter")?;
3528 let operand_arg = typed_tensor_binding(operand, "scatter")?;
3529 let scatter_arg = typed_tensor_binding(scatter_indices, "scatter")?;
3530 let updates_arg = typed_tensor_binding(updates, "scatter")?;
3531 unsafe {
3532 indexing::scatter_float_kernel::launch_unchecked::<T, I, CubeclCudaRuntime>(
3540 client,
3541 cube_count_for_len(update_len)?,
3542 cube_dim_1d(),
3543 output_parts,
3544 operand_arg.into_tensor_arg(),
3545 scatter_arg.into_tensor_arg(),
3546 updates_arg.into_tensor_arg(),
3547 comptime_sequence(&meta.window_dims),
3548 comptime_sequence(&config.update_window_dims),
3549 comptime_sequence(&config.scatter_dims_to_operand_dims),
3550 config.index_vector_dim,
3551 operand.shape().len(),
3552 updates.shape().len(),
3553 scatter_indices.shape().len(),
3554 );
3555 }
3556 Ok(output)
3557 }
3558
3559 fn scatter_complex_typed<T, F, I>(
3560 &self,
3561 operand: &TypedTensor<T>,
3562 scatter_indices: &TypedTensor<I>,
3563 updates: &TypedTensor<T>,
3564 config: &ScatterConfig,
3565 ) -> crate::Result<TypedTensor<T>>
3566 where
3567 T: CubeElement + TensorScalar + CubeComplex + Clone,
3568 F: CubeElement + TensorScalar + CubeFloat + Clone,
3569 I: CubeElement + TensorScalar + CubePrimitive + CubeNumeric + Clone + CudaIndexValidation,
3570 {
3571 let meta = scatter_launch_meta(
3572 operand.shape(),
3573 scatter_indices.shape(),
3574 updates.shape(),
3575 config,
3576 )?;
3577 let update_len = scatter_update_len(&meta)?;
3578 let output_len = checked_dim_product("scatter", "output shape", operand.shape())?;
3579 let output_part_len = output_len.checked_mul(2).ok_or_else(|| {
3580 crate::Error::invalid_argument(
3581 "scatter",
3582 "shape",
3583 "complex output part length overflow",
3584 )
3585 })?;
3586 let update_part_len = updates.n_elements().checked_mul(2).ok_or_else(|| {
3587 crate::Error::invalid_argument(
3588 "scatter",
3589 "shape",
3590 "complex update part length overflow",
3591 )
3592 })?;
3593 if output_len != 0 {
3594 cube_count_for_len(output_len)?;
3595 }
3596 if update_len != 0 {
3597 cube_count_for_len(update_len)?;
3598 }
3599 let client = self.runtime().client();
3600 ensure_resident_on_runtime(self.runtime(), operand, "scatter")?;
3601 typed_tensor_binding(operand, "scatter")?;
3602 ensure_resident_on_runtime(self.runtime(), scatter_indices, "scatter")?;
3603 typed_tensor_binding(scatter_indices, "scatter")?;
3604 ensure_resident_on_runtime(self.runtime(), updates, "scatter")?;
3605 typed_tensor_binding(updates, "scatter")?;
3606 typed_tensor_array_arg_as::<T, F>(updates, update_part_len, "scatter")?;
3607 ensure_atomic_add_supported::<F>(client, "scatter")?;
3608 I::validate(self, scatter_indices)?;
3609 let output = alloc_output::<T>(self.runtime(), operand.shape())?;
3610 if output.n_elements() == 0 {
3611 return Ok(output);
3612 }
3613
3614 launch_unary_tensor_into(
3615 self.runtime(),
3616 &output,
3617 operand,
3618 "scatter",
3619 cube_count_for_len(output.n_elements())?,
3620 cube_dim_1d(),
3621 |client, count, dim, out_arg, operand_arg| unsafe {
3622 indexing::scatter_copy_kernel::launch_unchecked::<T, CubeclCudaRuntime>(
3623 client,
3624 count,
3625 dim,
3626 out_arg.into_tensor_arg(),
3627 operand_arg.into_tensor_arg(),
3628 );
3629 },
3630 )?;
3631
3632 if update_len == 0 {
3633 return Ok(output);
3634 }
3635 let output_parts = typed_tensor_array_arg_as::<T, F>(&output, output_part_len, "scatter")?;
3638 let update_parts = typed_tensor_array_arg_as::<T, F>(updates, update_part_len, "scatter")?;
3639 let operand_arg = typed_tensor_binding(operand, "scatter")?;
3640 let scatter_arg = typed_tensor_binding(scatter_indices, "scatter")?;
3641 let updates_arg = typed_tensor_binding(updates, "scatter")?;
3642 unsafe {
3643 indexing::scatter_complex_kernel::launch_unchecked::<T, F, I, CubeclCudaRuntime>(
3651 client,
3652 cube_count_for_len(update_len)?,
3653 cube_dim_1d(),
3654 output_parts,
3655 operand_arg.into_tensor_arg(),
3656 scatter_arg.into_tensor_arg(),
3657 updates_arg.into_tensor_arg(),
3658 update_parts,
3659 comptime_sequence(&meta.window_dims),
3660 comptime_sequence(&config.update_window_dims),
3661 comptime_sequence(&config.scatter_dims_to_operand_dims),
3662 config.index_vector_dim,
3663 operand.shape().len(),
3664 updates.shape().len(),
3665 scatter_indices.shape().len(),
3666 );
3667 }
3668 Ok(output)
3669 }
3670}
3671
3672impl BackendRuntimeCache for CudaBackend {
3673 type RuntimeCache = ();
3674}
3675
3676#[derive(Clone, Copy, Debug)]
3677enum CheckedIntegerDomain {
3678 DivisionByZero,
3679 NegativeExponent,
3680}
3681
3682#[derive(Clone, Copy)]
3683enum CastIntegerTarget {
3684 I32,
3685 I64,
3686}
3687
3688trait CudaCastFloat:
3689 CubeElement
3690 + TensorScalar
3691 + CubeFloat
3692 + CubePrimitive<WithScalar<bool> = bool, WithScalar<Self> = Self>
3693 + Clone
3694 + Send
3695 + Sync
3696 + Copy
3697 + fmt::Display
3698 + 'static
3699{
3700 fn bounds(target: CastIntegerTarget) -> (Self, Self, bool);
3701 fn read_flag(backend: &CudaBackend, flag: &TypedTensor<Self>) -> crate::Result<Self>;
3702 fn invalid_error(self, target: CastIntegerTarget) -> crate::Error;
3703 fn is_nonfinite(self) -> bool;
3704 fn cpu_real_display(self) -> String;
3705}
3706
3707macro_rules! impl_cuda_cast_float {
3708 ($ty:ty, $variant:ident, $i32_max_inclusive:expr, $display:expr) => {
3709 impl CudaCastFloat for $ty {
3710 fn bounds(target: CastIntegerTarget) -> (Self, Self, bool) {
3711 match target {
3712 CastIntegerTarget::I32 => (
3713 i32::MIN as Self,
3714 if $i32_max_inclusive {
3715 i32::MAX as Self
3716 } else {
3717 2_147_483_648.0 as Self
3718 },
3719 $i32_max_inclusive,
3720 ),
3721 CastIntegerTarget::I64 => (
3722 -9_223_372_036_854_775_808.0 as Self,
3723 9_223_372_036_854_775_808.0 as Self,
3724 false,
3725 ),
3726 }
3727 }
3728 fn read_flag(backend: &CudaBackend, flag: &TypedTensor<Self>) -> crate::Result<Self> {
3729 let host = interop::download_typed_tensor(backend.runtime(), flag, "cast")?;
3730 host.as_slice()?.get(1).copied().ok_or_else(|| {
3731 crate::Error::invalid_argument(
3732 "cast",
3733 "validation_flag",
3734 "validation flag was malformed",
3735 )
3736 })
3737 }
3738 fn invalid_error(self, target: CastIntegerTarget) -> crate::Error {
3739 let name = match target {
3740 CastIntegerTarget::I32 => "i32",
3741 CastIntegerTarget::I64 => "i64",
3742 };
3743 let message = if !self.is_finite() {
3744 format!(
3745 "real value must be finite when casting to {name}, got {}",
3746 self.cpu_real_display()
3747 )
3748 } else {
3749 format!(
3750 "real value {} is out of {name} range",
3751 self.cpu_real_display()
3752 )
3753 };
3754 crate::Error::invalid_argument("cast", "value", message)
3755 }
3756 fn is_nonfinite(self) -> bool {
3757 !self.is_finite()
3758 }
3759 fn cpu_real_display(self) -> String {
3760 ($display)(self)
3761 }
3762 }
3763 };
3764}
3765impl_cuda_cast_float!(f32, F32, false, |value: f32| format!("{}", value as f64));
3766impl_cuda_cast_float!(f64, F64, true, |value: f64| format!("{value}"));
3767
3768fn validate_cuda_real_cast<S, F>(
3769 backend: &CudaBackend,
3770 input: &TypedTensor<S>,
3771 stride: usize,
3772 target: CastIntegerTarget,
3773) -> crate::Result<()>
3774where
3775 S: CubeElement + TensorScalar + Clone,
3776 F: CudaCastFloat,
3777{
3778 ensure_resident_on_runtime(backend.runtime(), input, "cast")?;
3779 let n = input.n_elements();
3780 let _validated_input = typed_tensor_array_arg(input, "cast")?;
3781 if n == 0 {
3782 return Ok(());
3783 }
3784 u32::try_from(n).map_err(|_| {
3785 crate::Error::invalid_argument(
3786 "cast",
3787 "shape",
3788 "validation domain exceeds u32::MAX elements",
3789 )
3790 })?;
3791 let count = cube_count_for_len(n)?;
3792 let input_parts = n.checked_mul(stride).ok_or_else(|| {
3793 crate::Error::invalid_argument("cast", "shape", "validation input length overflow")
3794 })?;
3795 let input_arg = typed_tensor_array_arg_as::<S, F>(input, input_parts, "cast")?;
3796 let flag = alloc_output::<F>(backend.runtime(), &[2])?;
3797 let flag_u32_len = std::mem::size_of::<F>()
3798 .checked_mul(2)
3799 .and_then(|x| x.checked_div(std::mem::size_of::<u32>()))
3800 .ok_or_else(|| {
3801 crate::Error::invalid_argument("cast", "shape", "validation flag size overflow")
3802 })?;
3803 let flag_atomic = typed_tensor_array_arg_as::<F, u32>(&flag, flag_u32_len, "cast")?;
3804 let flag_values = typed_tensor_array_arg(&flag, "cast")?;
3805 unsafe {
3806 indexing::init_float_index_validation_flag::launch_unchecked::<F, CubeclCudaRuntime>(
3807 backend.runtime().client(),
3808 CubeCount::Static(1, 1, 1),
3809 cube_dim_1d(),
3810 flag_atomic,
3811 flag_values,
3812 );
3813 }
3814 let flag_atomic = typed_tensor_array_arg_as::<F, u32>(&flag, flag_u32_len, "cast")?;
3815 let (min, max, inclusive) = F::bounds(target);
3816 unsafe {
3817 structural::validate_real_cast::launch_unchecked::<F, CubeclCudaRuntime>(
3818 backend.runtime().client(),
3819 count,
3820 cube_dim_1d(),
3821 input_arg,
3822 flag_atomic,
3823 min,
3824 max,
3825 stride,
3826 inclusive,
3827 );
3828 }
3829 let input_arg = typed_tensor_array_arg_as::<S, F>(input, input_parts, "cast")?;
3830 let flag_atomic = typed_tensor_array_arg_as::<F, u32>(&flag, flag_u32_len, "cast")?;
3831 let flag_values = typed_tensor_array_arg(&flag, "cast")?;
3832 unsafe {
3833 structural::extract_invalid_real_cast::launch_unchecked::<F, CubeclCudaRuntime>(
3834 backend.runtime().client(),
3835 CubeCount::Static(1, 1, 1),
3836 cube_dim_1d(),
3837 input_arg,
3838 flag_atomic,
3839 flag_values,
3840 stride,
3841 );
3842 }
3843 let value = F::read_flag(backend, &flag)?;
3844 let (min, max, inclusive) = F::bounds(target);
3845 if value.is_nonfinite() || value < min || if inclusive { value > max } else { value >= max } {
3846 return Err(value.invalid_error(target));
3847 }
3848 Ok(())
3849}
3850
3851fn checked_integer_domain_error(
3852 domain: CheckedIntegerDomain,
3853 op: &'static str,
3854 dtype: crate::DType,
3855) -> crate::Error {
3856 match domain {
3857 CheckedIntegerDomain::DivisionByZero => error::division_by_zero(op, dtype),
3858 CheckedIntegerDomain::NegativeExponent => error::negative_integer_exponent(op, dtype),
3859 }
3860}
3861
3862fn read_checked_integer_flag(
3863 backend: &CudaBackend,
3864 flag: &TypedTensor<i32>,
3865 op: &'static str,
3866) -> crate::Result<i32> {
3867 let host = interop::download_typed_tensor(backend.runtime(), flag, op)?;
3868 Ok(host.as_slice()?.first().copied().unwrap_or_default())
3869}
3870
3871trait CudaFloatIndex:
3872 CubeElement
3873 + TensorScalar
3874 + CubePrimitive<WithScalar<bool> = bool, WithScalar<Self> = Self>
3875 + CubeFloat
3876 + Clone
3877 + Send
3878 + Sync
3879 + fmt::Display
3880 + Copy
3881 + 'static
3882{
3883 const MAX_EXACT_INTEGER: Self;
3884 fn is_invalid_index(self) -> bool;
3885 fn read_invalid_flag(backend: &CudaBackend, flag: &TypedTensor<Self>) -> crate::Result<Self>;
3886}
3887
3888trait CudaIndexValidation: Sized {
3889 fn validate(backend: &CudaBackend, indices: &TypedTensor<Self>) -> crate::Result<()>;
3890}
3891
3892impl CudaIndexValidation for f32 {
3893 fn validate(backend: &CudaBackend, indices: &TypedTensor<Self>) -> crate::Result<()> {
3894 validate_float_index_tensor(backend, indices)
3895 }
3896}
3897
3898impl CudaIndexValidation for f64 {
3899 fn validate(backend: &CudaBackend, indices: &TypedTensor<Self>) -> crate::Result<()> {
3900 validate_float_index_tensor(backend, indices)
3901 }
3902}
3903
3904impl CudaIndexValidation for i32 {
3905 fn validate(_backend: &CudaBackend, _indices: &TypedTensor<Self>) -> crate::Result<()> {
3906 Ok(())
3907 }
3908}
3909
3910impl CudaIndexValidation for i64 {
3911 fn validate(_backend: &CudaBackend, _indices: &TypedTensor<Self>) -> crate::Result<()> {
3912 Ok(())
3913 }
3914}
3915
3916impl CudaFloatIndex for f32 {
3917 const MAX_EXACT_INTEGER: Self = 16_777_216.0;
3918
3919 fn is_invalid_index(self) -> bool {
3920 !self.is_finite() || self.fract() != 0.0 || self.abs() > 16_777_216.0
3921 }
3922
3923 fn read_invalid_flag(backend: &CudaBackend, flag: &TypedTensor<Self>) -> crate::Result<Self> {
3924 let host = interop::download_typed_tensor(backend.runtime(), flag, "index_tensor")?;
3925 host.as_slice()?.get(1).copied().ok_or_else(|| {
3926 crate::Error::invalid_argument(
3927 "index_tensor",
3928 "validation_flag",
3929 "validation flag was malformed",
3930 )
3931 })
3932 }
3933}
3934
3935impl CudaFloatIndex for f64 {
3936 const MAX_EXACT_INTEGER: Self = 9_007_199_254_740_992.0;
3937
3938 fn is_invalid_index(self) -> bool {
3939 !self.is_finite() || self.fract() != 0.0 || self.abs() > 9_007_199_254_740_992.0
3940 }
3941
3942 fn read_invalid_flag(backend: &CudaBackend, flag: &TypedTensor<Self>) -> crate::Result<Self> {
3943 let host = interop::download_typed_tensor(backend.runtime(), flag, "index_tensor")?;
3944 host.as_slice()?.get(1).copied().ok_or_else(|| {
3945 crate::Error::invalid_argument(
3946 "index_tensor",
3947 "validation_flag",
3948 "validation flag was malformed",
3949 )
3950 })
3951 }
3952}
3953
3954fn validate_float_index_tensor<F>(
3955 backend: &CudaBackend,
3956 indices: &TypedTensor<F>,
3957) -> crate::Result<()>
3958where
3959 F: CudaFloatIndex,
3960{
3961 ensure_resident_on_runtime(backend.runtime(), indices, "index_tensor")?;
3962 let indices_arg = typed_tensor_binding(indices, "index_tensor")?;
3963 if indices.n_elements() == 0 {
3964 return Ok(());
3965 }
3966 u32::try_from(indices.n_elements()).map_err(|_| {
3967 crate::Error::invalid_argument(
3968 "index_tensor",
3969 "shape",
3970 "float index validation domain exceeds u32::MAX elements",
3971 )
3972 })?;
3973 let count = cube_count_for_len(indices.n_elements())?;
3974 let flag_u32_len = std::mem::size_of::<F>()
3975 .checked_mul(2)
3976 .and_then(|bytes| bytes.checked_div(std::mem::size_of::<u32>()))
3977 .ok_or_else(|| {
3978 crate::Error::invalid_argument("index_tensor", "shape", "flag size overflow")
3979 })?;
3980 let flag = alloc_output::<F>(backend.runtime(), &[2])?;
3981 let flag_values = typed_tensor_array_arg(&flag, "index_tensor")?;
3982 let flag_atomic = typed_tensor_array_arg_as::<F, u32>(&flag, flag_u32_len, "index_tensor")?;
3983 unsafe {
3984 indexing::init_float_index_validation_flag::launch_unchecked::<F, CubeclCudaRuntime>(
3987 backend.runtime().client(),
3988 CubeCount::Static(1, 1, 1),
3989 cube_dim_1d(),
3990 flag_atomic,
3991 flag_values,
3992 );
3993 }
3994 let flag_atomic = typed_tensor_array_arg_as::<F, u32>(&flag, flag_u32_len, "index_tensor")?;
3995 unsafe {
3996 indexing::validate_float_indices_kernel::launch_unchecked::<F, CubeclCudaRuntime>(
4000 backend.runtime().client(),
4001 count,
4002 cube_dim_1d(),
4003 indices_arg.into_tensor_arg(),
4004 flag_atomic,
4005 F::MAX_EXACT_INTEGER,
4006 );
4007 }
4008 let indices_arg = typed_tensor_binding(indices, "index_tensor")?;
4009 let flag_atomic = typed_tensor_array_arg_as::<F, u32>(&flag, flag_u32_len, "index_tensor")?;
4010 let flag_values = typed_tensor_array_arg(&flag, "index_tensor")?;
4011 unsafe {
4012 indexing::extract_invalid_float_index_kernel::launch_unchecked::<F, CubeclCudaRuntime>(
4015 backend.runtime().client(),
4016 CubeCount::Static(1, 1, 1),
4017 cube_dim_1d(),
4018 indices_arg.into_tensor_arg(),
4019 flag_atomic,
4020 flag_values,
4021 );
4022 }
4023 let invalid = F::read_invalid_flag(backend, &flag)?;
4024 if invalid.is_invalid_index() {
4025 return Err(crate::Error::invalid_argument(
4026 "index_tensor",
4027 "index",
4028 format!("index value {invalid} is not an exactly representable i64"),
4029 ));
4030 }
4031 Ok(())
4032}
4033
4034fn launch_checked_integer_binary<I>(
4035 backend: &CudaBackend,
4036 lhs: &TypedTensor<I>,
4037 rhs: &TypedTensor<I>,
4038 op: &'static str,
4039 dtype: crate::DType,
4040 domain: CheckedIntegerDomain,
4041 launch: impl FnOnce(
4042 &ComputeClient<CubeclCudaRuntime>,
4043 CubeCount,
4044 CubeDim,
4045 ArrayArg<CubeclCudaRuntime>,
4046 ArrayArg<CubeclCudaRuntime>,
4047 ArrayArg<CubeclCudaRuntime>,
4048 ArrayArg<CubeclCudaRuntime>,
4049 ),
4050) -> crate::Result<TypedTensor<I>>
4051where
4052 I: CubeElement + TensorScalar + CubePrimitive + Clone + Send + Sync + 'static,
4053{
4054 dispatch::ensure_same_shape(op, lhs.shape(), rhs.shape())?;
4055 ensure_resident_on_runtime(backend.runtime(), lhs, op)?;
4056 ensure_resident_on_runtime(backend.runtime(), rhs, op)?;
4057
4058 let output = alloc_output::<I>(backend.runtime(), lhs.shape())?;
4059 if output.n_elements() == 0 {
4060 return Ok(output);
4061 }
4062
4063 let flag = alloc_output::<i32>(backend.runtime(), &[1])?;
4064 launch_nullary_into(
4065 backend.runtime(),
4066 &flag,
4067 op,
4068 cube_count_for_len(flag.n_elements())?,
4069 cube_dim_1d(),
4070 |client, count, dim, out| unsafe {
4071 structural::fill_zero_kernel::launch_unchecked::<i32, CubeclCudaRuntime>(
4072 client, count, dim, out,
4073 );
4074 },
4075 )?;
4076
4077 let output_arg = typed_tensor_array_arg(&output, op)?;
4078 let lhs_arg = typed_tensor_array_arg(lhs, op)?;
4079 let rhs_arg = typed_tensor_array_arg(rhs, op)?;
4080 let flag_arg = typed_tensor_array_arg(&flag, op)?;
4081 launch(
4082 backend.runtime().client(),
4083 cube_count_for_len(output.n_elements())?,
4084 cube_dim_1d(),
4085 output_arg,
4086 lhs_arg,
4087 rhs_arg,
4088 flag_arg,
4089 );
4090
4091 if read_checked_integer_flag(backend, &flag, op)? != 0 {
4092 return Err(checked_integer_domain_error(domain, op, dtype));
4093 }
4094 Ok(output)
4095}
4096
4097fn launch_scalar_binary<I>(
4098 backend: &CudaBackend,
4099 lhs: &TypedTensor<I>,
4100 rhs: &TypedTensor<I>,
4101 op: &'static str,
4102 launch: impl FnOnce(
4103 &ComputeClient<CubeclCudaRuntime>,
4104 CubeCount,
4105 CubeDim,
4106 ArrayArg<CubeclCudaRuntime>,
4107 ArrayArg<CubeclCudaRuntime>,
4108 ArrayArg<CubeclCudaRuntime>,
4109 bool,
4110 ),
4111) -> crate::Result<TypedTensor<I>>
4112where
4113 I: CubeElement + TensorScalar + CubePrimitive + Clone + Send + Sync + 'static,
4114{
4115 if !(lhs.shape().is_empty() ^ rhs.shape().is_empty()) {
4116 return Err(crate::Error::shape_mismatch(
4117 op,
4118 lhs.shape().to_vec(),
4119 rhs.shape().to_vec(),
4120 ));
4121 }
4122 ensure_resident_on_runtime(backend.runtime(), lhs, op)?;
4123 ensure_resident_on_runtime(backend.runtime(), rhs, op)?;
4124
4125 let lhs_scalar = lhs.shape().is_empty();
4126 let output_shape = if lhs_scalar { rhs.shape() } else { lhs.shape() };
4127 let output = alloc_output::<I>(backend.runtime(), output_shape)?;
4128 let output_arg = typed_tensor_array_arg(&output, op)?;
4129 let lhs_arg = typed_tensor_array_arg(lhs, op)?;
4130 let rhs_arg = typed_tensor_array_arg(rhs, op)?;
4131 if output.n_elements() == 0 {
4132 return Ok(output);
4133 }
4134 launch(
4135 backend.runtime().client(),
4136 cube_count_for_len(output.n_elements())?,
4137 cube_dim_1d(),
4138 output_arg,
4139 lhs_arg,
4140 rhs_arg,
4141 lhs_scalar,
4142 );
4143 Ok(output)
4144}
4145
4146fn launch_real_complex_scalar_binary<R, C>(
4147 backend: &CudaBackend,
4148 real: &TypedTensor<R>,
4149 complex: &TypedTensor<C>,
4150 op: &'static str,
4151 real_lhs: bool,
4152 mode: usize,
4153) -> crate::Result<TypedTensor<C>>
4154where
4155 R: TensorScalar + CubeFloat + CubeElement + CubePrimitive + Clone + Send + Sync + 'static,
4156 C: CubeComplex<FloatElem = R>
4157 + TensorScalar
4158 + CubeElement
4159 + CubePrimitive
4160 + Clone
4161 + Send
4162 + Sync
4163 + 'static,
4164{
4165 if !real.shape().is_empty() {
4166 return Err(crate::Error::shape_mismatch(
4167 op,
4168 if real_lhs {
4169 real.shape().to_vec()
4170 } else {
4171 complex.shape().to_vec()
4172 },
4173 if real_lhs {
4174 complex.shape().to_vec()
4175 } else {
4176 real.shape().to_vec()
4177 },
4178 ));
4179 }
4180 ensure_resident_on_runtime(backend.runtime(), real, op)?;
4181 ensure_resident_on_runtime(backend.runtime(), complex, op)?;
4182 let component_len = complex.n_elements().checked_mul(2).ok_or_else(|| {
4183 crate::Error::invalid_argument(op, "shape", "complex component length overflow")
4184 })?;
4185 let real_arg = typed_tensor_array_arg(real, op)?;
4186 let complex_arg = typed_tensor_array_arg_as::<C, R>(complex, component_len, op)?;
4190
4191 let output = alloc_output::<C>(backend.runtime(), complex.shape())?;
4192 let output_arg = typed_tensor_array_arg_as::<C, R>(&output, component_len, op)?;
4193 if output.n_elements() == 0 {
4194 return Ok(output);
4195 }
4196 unsafe {
4197 elementwise::scalar_real_complex_binary::launch_unchecked::<R, CubeclCudaRuntime>(
4198 backend.runtime().client(),
4199 cube_count_for_len(output.n_elements())?,
4200 cube_dim_1d(),
4201 output_arg,
4202 real_arg,
4203 complex_arg,
4204 real_lhs,
4205 mode,
4206 );
4207 }
4208 Ok(output)
4209}
4210
4211fn typed_or_unsupported<'a, T: tenferro_tensor::TensorScalar>(
4218 tensor: &'a Tensor,
4219 op: &'static str,
4220) -> crate::Result<&'a tenferro_tensor::TypedTensor<T>> {
4221 tensor.as_typed::<T>().ok_or_else(|| {
4222 crate::Error::unsupported(op, "the tensor does not carry the scalar its tag names")
4223 })
4224}
4225
4226fn promoted_real_complex_scalar_binary(
4227 backend: &CudaBackend,
4228 lhs: &Tensor,
4229 rhs: &Tensor,
4230 op: &'static str,
4231 mode: usize,
4232) -> Option<crate::Result<Tensor>> {
4233 match (lhs.dtype(), rhs.dtype()) {
4236 (DType::F32, DType::C32) if lhs.shape().is_empty() => Some((|| {
4237 let real = typed_or_unsupported::<f32>(lhs, op)?;
4238 let complex = typed_or_unsupported::<Complex32>(rhs, op)?;
4239 launch_real_complex_scalar_binary(backend, real, complex, op, true, mode)
4240 .map(Tensor::from_typed::<num_complex::Complex32>)
4241 })()),
4242 (DType::C32, DType::F32) if rhs.shape().is_empty() => Some((|| {
4243 let complex = typed_or_unsupported::<Complex32>(lhs, op)?;
4244 let real = typed_or_unsupported::<f32>(rhs, op)?;
4245 launch_real_complex_scalar_binary(backend, real, complex, op, false, mode)
4246 .map(Tensor::from_typed::<num_complex::Complex32>)
4247 })()),
4248 (DType::F64, DType::C64) if lhs.shape().is_empty() => Some((|| {
4249 let real = typed_or_unsupported::<f64>(lhs, op)?;
4250 let complex = typed_or_unsupported::<Complex64>(rhs, op)?;
4251 launch_real_complex_scalar_binary(backend, real, complex, op, true, mode)
4252 .map(Tensor::from_typed::<num_complex::Complex64>)
4253 })()),
4254 (DType::C64, DType::F64) if rhs.shape().is_empty() => Some((|| {
4255 let complex = typed_or_unsupported::<Complex64>(lhs, op)?;
4256 let real = typed_or_unsupported::<f64>(rhs, op)?;
4257 launch_real_complex_scalar_binary(backend, real, complex, op, false, mode)
4258 .map(Tensor::from_typed::<num_complex::Complex64>)
4259 })()),
4260 _ => None,
4261 }
4262}
4263
4264fn launch_checked_integer_scalar_binary<I>(
4265 backend: &CudaBackend,
4266 lhs: &TypedTensor<I>,
4267 rhs: &TypedTensor<I>,
4268 op: &'static str,
4269 dtype: crate::DType,
4270 domain: CheckedIntegerDomain,
4271 launch: impl FnOnce(
4272 &ComputeClient<CubeclCudaRuntime>,
4273 CubeCount,
4274 CubeDim,
4275 ArrayArg<CubeclCudaRuntime>,
4276 ArrayArg<CubeclCudaRuntime>,
4277 ArrayArg<CubeclCudaRuntime>,
4278 ArrayArg<CubeclCudaRuntime>,
4279 bool,
4280 ),
4281) -> crate::Result<TypedTensor<I>>
4282where
4283 I: CubeElement + TensorScalar + CubePrimitive + Clone + Send + Sync + 'static,
4284{
4285 if !(lhs.shape().is_empty() ^ rhs.shape().is_empty()) {
4286 return Err(crate::Error::shape_mismatch(
4287 op,
4288 lhs.shape().to_vec(),
4289 rhs.shape().to_vec(),
4290 ));
4291 }
4292 ensure_resident_on_runtime(backend.runtime(), lhs, op)?;
4293 ensure_resident_on_runtime(backend.runtime(), rhs, op)?;
4294
4295 let lhs_scalar = lhs.shape().is_empty();
4296 let output_shape = if lhs_scalar { rhs.shape() } else { lhs.shape() };
4297 let output = alloc_output::<I>(backend.runtime(), output_shape)?;
4298 let output_arg = typed_tensor_array_arg(&output, op)?;
4299 let lhs_arg = typed_tensor_array_arg(lhs, op)?;
4300 let rhs_arg = typed_tensor_array_arg(rhs, op)?;
4301 if output.n_elements() == 0 {
4302 return Ok(output);
4303 }
4304 let flag = alloc_output::<i32>(backend.runtime(), &[1])?;
4305 let flag_arg = typed_tensor_array_arg(&flag, op)?;
4306 launch_nullary_into(
4307 backend.runtime(),
4308 &flag,
4309 op,
4310 cube_count_for_len(flag.n_elements())?,
4311 cube_dim_1d(),
4312 |client, count, dim, out| unsafe {
4313 structural::fill_zero_kernel::launch_unchecked::<i32, CubeclCudaRuntime>(
4314 client, count, dim, out,
4315 );
4316 },
4317 )?;
4318
4319 launch(
4320 backend.runtime().client(),
4321 cube_count_for_len(output.n_elements())?,
4322 cube_dim_1d(),
4323 output_arg,
4324 lhs_arg,
4325 rhs_arg,
4326 flag_arg,
4327 lhs_scalar,
4328 );
4329 if read_checked_integer_flag(backend, &flag, op)? != 0 {
4330 return Err(checked_integer_domain_error(domain, op, dtype));
4331 }
4332 Ok(output)
4333}
4334
4335#[derive(Clone, Copy)]
4336enum UnaryReadOp {
4337 Neg,
4338 Exp,
4339 Log,
4340 Sin,
4341 Cos,
4342 Tanh,
4343 Sqrt,
4344 Rsqrt,
4345 Expm1,
4346 Log1p,
4347 Erf,
4348}
4349
4350enum CudaReadInput<'a> {
4356 Borrowed(&'a Tensor),
4358 Materialized(Box<Tensor>),
4360}
4361
4362impl CudaReadInput<'_> {
4363 fn as_tensor(&self) -> &Tensor {
4364 match self {
4365 Self::Borrowed(tensor) => tensor,
4366 Self::Materialized(tensor) => tensor,
4367 }
4368 }
4369}
4370
4371fn launch_elementwise_binary_into<T>(
4372 backend: &CudaBackend,
4373 lhs: &Tensor,
4374 rhs: &Tensor,
4375 out: &mut Tensor,
4376 op: &'static str,
4377 launch: impl FnOnce(
4378 &ComputeClient<CubeclCudaRuntime>,
4379 CubeCount,
4380 CubeDim,
4381 ArrayArg<CubeclCudaRuntime>,
4382 ArrayArg<CubeclCudaRuntime>,
4383 ArrayArg<CubeclCudaRuntime>,
4384 ),
4385) -> crate::Result<()>
4386where
4387 T: CubeElement + TensorScalar + Clone,
4388{
4389 let lhs = lhs.as_typed::<T>().ok_or_else(|| {
4390 crate::Error::unsupported(op, "the GPU dispatch requires a preset scalar")
4391 })?;
4392 let rhs = rhs.as_typed::<T>().ok_or_else(|| {
4393 crate::Error::unsupported(op, "the GPU dispatch requires a preset scalar")
4394 })?;
4395 let out = out.as_typed_mut::<T>().ok_or_else(|| {
4396 crate::Error::unsupported(op, "the GPU dispatch requires a preset scalar")
4397 })?;
4398 ensure_resident_on_runtime(backend.runtime(), lhs, op)?;
4399 ensure_resident_on_runtime(backend.runtime(), rhs, op)?;
4400 let lhs_arg = typed_tensor_array_arg(lhs, op)?;
4401 let rhs_arg = typed_tensor_array_arg(rhs, op)?;
4402 let out_len = out.n_elements();
4403 ensure_resident_on_runtime(backend.runtime(), out, op)?;
4404 let out_arg = typed_tensor_mut_array_arg(out, op)?;
4405 if out_len == 0 {
4406 return Ok(());
4407 }
4408 launch(
4409 backend.runtime().client(),
4410 cube_count_for_len(out_len)?,
4411 cube_dim_1d(),
4412 out_arg,
4413 lhs_arg,
4414 rhs_arg,
4415 );
4416 Ok(())
4417}
4418
4419fn launch_elementwise_unary_into<T>(
4420 backend: &CudaBackend,
4421 input: &Tensor,
4422 out: &mut Tensor,
4423 op: &'static str,
4424 launch: impl FnOnce(
4425 &ComputeClient<CubeclCudaRuntime>,
4426 CubeCount,
4427 CubeDim,
4428 ArrayArg<CubeclCudaRuntime>,
4429 ArrayArg<CubeclCudaRuntime>,
4430 ),
4431) -> crate::Result<()>
4432where
4433 T: CubeElement + TensorScalar + Clone,
4434{
4435 let input = input.as_typed::<T>().ok_or_else(|| {
4436 crate::Error::unsupported(op, "the GPU dispatch requires a preset scalar")
4437 })?;
4438 let out = out.as_typed_mut::<T>().ok_or_else(|| {
4439 crate::Error::unsupported(op, "the GPU dispatch requires a preset scalar")
4440 })?;
4441 ensure_resident_on_runtime(backend.runtime(), input, op)?;
4442 let input_arg = typed_tensor_array_arg(input, op)?;
4443 let out_len = out.n_elements();
4444 ensure_resident_on_runtime(backend.runtime(), out, op)?;
4445 let out_arg = typed_tensor_mut_array_arg(out, op)?;
4446 if out_len == 0 {
4447 return Ok(());
4448 }
4449 launch(
4450 backend.runtime().client(),
4451 cube_count_for_len(out_len)?,
4452 cube_dim_1d(),
4453 out_arg,
4454 input_arg,
4455 );
4456 Ok(())
4457}
4458
4459impl CudaBackend {
4460 fn elementwise_read_into_native(
4464 &mut self,
4465 op: ElementwiseReadOp,
4466 inputs: &[TensorRead<'_>],
4467 out: &mut TensorWrite<'_>,
4468 ) -> Option<crate::Result<()>> {
4469 if inputs.len() != op.arity() || inputs.iter().any(|input| input.as_tensor().is_none()) {
4470 return None;
4471 }
4472 let TensorWrite::Tensor(output) = out else {
4473 return None;
4474 };
4475 if inputs
4476 .iter()
4477 .any(|input| input.shape() != output.shape() || input.dtype() != output.dtype())
4478 {
4479 return None;
4480 }
4481 let first = inputs.first().and_then(|input| input.as_tensor())?;
4482 let second = inputs.get(1).and_then(|input| input.as_tensor());
4483 let dtype = output.dtype();
4484
4485 macro_rules! binary {
4486 ($ty:ty, $kernel:ident) => {
4487 launch_elementwise_binary_into::<$ty>(
4488 self,
4489 first,
4490 second?,
4491 output,
4492 op.label(),
4493 |client, count, dim, out, lhs, rhs| unsafe {
4497 elementwise::$kernel::launch_unchecked::<$ty, CubeclCudaRuntime>(
4498 client, count, dim, out, lhs, rhs,
4499 );
4500 },
4501 )
4502 };
4503 }
4504 macro_rules! unary {
4505 ($ty:ty, $kernel:ident) => {
4506 launch_elementwise_unary_into::<$ty>(
4507 self,
4508 first,
4509 output,
4510 op.label(),
4511 |client, count, dim, out, input| unsafe {
4515 elementwise::$kernel::launch_unchecked::<$ty, CubeclCudaRuntime>(
4516 client, count, dim, out, input,
4517 );
4518 },
4519 )
4520 };
4521 }
4522
4523 let result = match op {
4524 ElementwiseReadOp::Add => match dtype {
4525 DType::F32 => binary!(f32, add_float),
4526 DType::F64 => binary!(f64, add_float),
4527 DType::I32 => binary!(i32, add_int),
4528 DType::I64 => binary!(i64, add_int),
4529 DType::C32 => binary!(Complex32, add_complex),
4530 DType::C64 => binary!(Complex64, add_complex),
4531 _ => return None,
4532 },
4533 ElementwiseReadOp::Subtract => match dtype {
4534 DType::F32 => binary!(f32, sub_float),
4535 DType::F64 => binary!(f64, sub_float),
4536 DType::I32 => binary!(i32, sub_int),
4537 DType::I64 => binary!(i64, sub_int),
4538 DType::C32 => binary!(Complex32, sub_complex),
4539 DType::C64 => binary!(Complex64, sub_complex),
4540 _ => return None,
4541 },
4542 ElementwiseReadOp::Multiply => match dtype {
4543 DType::F32 => binary!(f32, mul_float),
4544 DType::F64 => binary!(f64, mul_float),
4545 DType::I32 => binary!(i32, mul_int),
4546 DType::I64 => binary!(i64, mul_int),
4547 DType::C32 => binary!(Complex32, mul_complex),
4548 DType::C64 => binary!(Complex64, mul_complex),
4549 _ => return None,
4550 },
4551 ElementwiseReadOp::Negate => match dtype {
4552 DType::F32 => unary!(f32, neg_float),
4553 DType::F64 => unary!(f64, neg_float),
4554 DType::I32 => unary!(i32, neg_int),
4555 DType::I64 => unary!(i64, neg_int),
4556 DType::C32 => unary!(Complex32, neg_complex),
4557 DType::C64 => unary!(Complex64, neg_complex),
4558 _ => return None,
4559 },
4560 ElementwiseReadOp::Conj | ElementwiseReadOp::Divide => return None,
4561 _ => return None,
4562 };
4563 Some(result)
4564 }
4565
4566 fn binary_read_native(
4567 &self,
4568 op: ElementwiseReadOp,
4569 lhs: TensorRead<'_>,
4570 rhs: TensorRead<'_>,
4571 ) -> Option<crate::Result<Tensor>> {
4572 let lhs = lhs.tensor_view();
4573 let rhs = rhs.tensor_view();
4574 if lhs.dtype() != rhs.dtype() || lhs.shape() != rhs.shape() {
4575 return None;
4576 }
4577 let compact = |view: &TensorView| -> crate::Result<bool> {
4578 Ok(view.offset() == 0 && view.is_col_major_contiguous()?)
4579 };
4580 match (compact(&lhs), compact(&rhs)) {
4581 (Ok(true), Ok(true)) => {}
4582 (Ok(false), _) | (_, Ok(false)) => return None,
4583 (Err(error), _) | (_, Err(error)) => return Some(Err(error)),
4584 }
4585
4586 macro_rules! binary {
4587 ($ty:ty, $lhs:expr, $rhs:expr, $kernel:ident) => {
4588 dispatch::launch_binary_views(
4589 self.runtime(),
4590 $lhs,
4591 $rhs,
4592 lhs.shape(),
4593 op.label(),
4594 |client, count, dim, out, lhs_arg, rhs_arg| unsafe {
4597 elementwise::$kernel::launch_unchecked::<$ty, CubeclCudaRuntime>(
4598 client, count, dim, out, lhs_arg, rhs_arg,
4599 );
4600 },
4601 )
4602 .map(Tensor::from_typed::<$ty>)
4603 };
4604 }
4605 macro_rules! dispatch_binary {
4606 ($variant:ident, $ty:ty, $kernel:ident) => {
4607 match (&lhs, &rhs) {
4608 (TensorView::$variant(lhs), TensorView::$variant(rhs)) => {
4609 Some(binary!($ty, lhs, rhs, $kernel))
4610 }
4611 _ => None,
4612 }
4613 };
4614 }
4615 macro_rules! binary_parts {
4616 ($ty:ty, $float:ty, $lhs:expr, $rhs:expr) => {
4617 dispatch::launch_binary_views_parts::<$ty, $ty, $ty, $float>(
4618 self.runtime(),
4619 $lhs,
4620 $rhs,
4621 $lhs.shape(),
4622 op.label(),
4623 |client, count, dim, out, lhs_arg, rhs_arg| unsafe {
4627 elementwise::div_complex_parts::launch_unchecked::<$float, CubeclCudaRuntime>(
4628 client, count, dim, out, lhs_arg, rhs_arg,
4629 );
4630 },
4631 )
4632 .map(Tensor::from_typed::<$ty>)
4633 };
4634 }
4635
4636 match op {
4637 ElementwiseReadOp::Add => match lhs.dtype() {
4638 DType::F32 => dispatch_binary!(F32, f32, add_float),
4639 DType::F64 => dispatch_binary!(F64, f64, add_float),
4640 DType::I32 => dispatch_binary!(I32, i32, add_int),
4641 DType::I64 => dispatch_binary!(I64, i64, add_int),
4642 DType::C32 => dispatch_binary!(C32, Complex32, add_complex),
4643 DType::C64 => dispatch_binary!(C64, Complex64, add_complex),
4644 _ => None,
4645 },
4646 ElementwiseReadOp::Subtract => match lhs.dtype() {
4647 DType::F32 => dispatch_binary!(F32, f32, sub_float),
4648 DType::F64 => dispatch_binary!(F64, f64, sub_float),
4649 DType::I32 => dispatch_binary!(I32, i32, sub_int),
4650 DType::I64 => dispatch_binary!(I64, i64, sub_int),
4651 DType::C32 => dispatch_binary!(C32, Complex32, sub_complex),
4652 DType::C64 => dispatch_binary!(C64, Complex64, sub_complex),
4653 _ => None,
4654 },
4655 ElementwiseReadOp::Multiply => match lhs.dtype() {
4656 DType::F32 => dispatch_binary!(F32, f32, mul_float),
4657 DType::F64 => dispatch_binary!(F64, f64, mul_float),
4658 DType::I32 => dispatch_binary!(I32, i32, mul_int),
4659 DType::I64 => dispatch_binary!(I64, i64, mul_int),
4660 DType::C32 => dispatch_binary!(C32, Complex32, mul_complex),
4661 DType::C64 => dispatch_binary!(C64, Complex64, mul_complex),
4662 _ => None,
4663 },
4664 ElementwiseReadOp::Divide => match lhs.dtype() {
4665 DType::F32 => dispatch_binary!(F32, f32, div_float),
4666 DType::F64 => dispatch_binary!(F64, f64, div_float),
4667 DType::C32 => match (&lhs, &rhs) {
4668 (TensorView::C32(lhs), TensorView::C32(rhs)) => {
4669 Some(binary_parts!(Complex32, f32, lhs, rhs))
4670 }
4671 _ => None,
4672 },
4673 DType::C64 => match (&lhs, &rhs) {
4674 (TensorView::C64(lhs), TensorView::C64(rhs)) => {
4675 Some(binary_parts!(Complex64, f64, lhs, rhs))
4676 }
4677 _ => None,
4678 },
4679 _ => None,
4680 },
4681 _ => None,
4682 }
4683 }
4684
4685 fn unary_read_native(
4686 &self,
4687 op: UnaryReadOp,
4688 input: TensorRead<'_>,
4689 ) -> Option<crate::Result<Tensor>> {
4690 let input = input.tensor_view();
4691 let compact = if input.offset() != 0 {
4692 false
4693 } else {
4694 match input.is_col_major_contiguous() {
4695 Ok(compact) => compact,
4696 Err(error) => return Some(Err(error)),
4697 }
4698 };
4699 if !compact {
4700 return None;
4701 }
4702
4703 macro_rules! unary {
4704 ($variant:ident, $ty:ty, $kernel:ident) => {
4705 match &input {
4706 TensorView::$variant(input) => Some(
4707 dispatch::launch_unary_view(
4708 self.runtime(),
4709 input,
4710 input.shape(),
4711 match op {
4712 UnaryReadOp::Neg => "neg",
4713 UnaryReadOp::Exp => "exp",
4714 UnaryReadOp::Log => "log",
4715 UnaryReadOp::Sin => "sin",
4716 UnaryReadOp::Cos => "cos",
4717 UnaryReadOp::Tanh => "tanh",
4718 UnaryReadOp::Sqrt => "sqrt",
4719 UnaryReadOp::Rsqrt => "rsqrt",
4720 UnaryReadOp::Expm1 => "expm1",
4721 UnaryReadOp::Log1p => "log1p",
4722 UnaryReadOp::Erf => "erf",
4723 },
4724 |client, count, dim, out, input_arg| unsafe {
4727 elementwise::$kernel::launch_unchecked::<$ty, CubeclCudaRuntime>(
4728 client, count, dim, out, input_arg,
4729 );
4730 },
4731 )
4732 .map(Tensor::from_typed::<$ty>),
4733 ),
4734 _ => None,
4735 }
4736 };
4737 }
4738
4739 match (input.dtype(), op) {
4740 (DType::F32, UnaryReadOp::Neg) => unary!(F32, f32, neg_float),
4741 (DType::F64, UnaryReadOp::Neg) => unary!(F64, f64, neg_float),
4742 (DType::I32, UnaryReadOp::Neg) => unary!(I32, i32, neg_int),
4743 (DType::I64, UnaryReadOp::Neg) => unary!(I64, i64, neg_int),
4744 (DType::C32, UnaryReadOp::Neg) => unary!(C32, Complex32, neg_complex),
4745 (DType::C64, UnaryReadOp::Neg) => unary!(C64, Complex64, neg_complex),
4746 (DType::F32, UnaryReadOp::Exp) => unary!(F32, f32, exp_float),
4747 (DType::F64, UnaryReadOp::Exp) => unary!(F64, f64, exp_float),
4748 (DType::F32, UnaryReadOp::Log) => unary!(F32, f32, log_float),
4749 (DType::F64, UnaryReadOp::Log) => unary!(F64, f64, log_float),
4750 (DType::F32, UnaryReadOp::Sin) => unary!(F32, f32, sin_float),
4751 (DType::F64, UnaryReadOp::Sin) => unary!(F64, f64, sin_float),
4752 (DType::F32, UnaryReadOp::Cos) => unary!(F32, f32, cos_float),
4753 (DType::F64, UnaryReadOp::Cos) => unary!(F64, f64, cos_float),
4754 (DType::F32, UnaryReadOp::Tanh) => unary!(F32, f32, tanh_float),
4755 (DType::F64, UnaryReadOp::Tanh) => unary!(F64, f64, tanh_float),
4756 (DType::F32, UnaryReadOp::Sqrt) => unary!(F32, f32, sqrt_float),
4757 (DType::F64, UnaryReadOp::Sqrt) => unary!(F64, f64, sqrt_float),
4758 (DType::F32, UnaryReadOp::Rsqrt) => unary!(F32, f32, rsqrt_float),
4759 (DType::F64, UnaryReadOp::Rsqrt) => unary!(F64, f64, rsqrt_float),
4760 (DType::F32, UnaryReadOp::Expm1) => unary!(F32, f32, expm1_float),
4761 (DType::F64, UnaryReadOp::Expm1) => unary!(F64, f64, expm1_float),
4762 (DType::F32, UnaryReadOp::Log1p) => unary!(F32, f32, log1p_float),
4763 (DType::F64, UnaryReadOp::Log1p) => unary!(F64, f64, log1p_float),
4764 (DType::F32, UnaryReadOp::Erf) => unary!(F32, f32, erf_float),
4765 (DType::F64, UnaryReadOp::Erf) => unary!(F64, f64, erf_float),
4766 _ => None,
4767 }
4768 }
4769
4770 fn read_input<'a>(&mut self, input: TensorRead<'a>) -> crate::Result<CudaReadInput<'a>> {
4773 match input.as_tensor() {
4774 Some(tensor) => Ok(CudaReadInput::Borrowed(tensor)),
4775 None => Ok(CudaReadInput::Materialized(Box::new(
4776 ops::to_contiguous_read(self, input)?,
4777 ))),
4778 }
4779 }
4780}
4781
4782fn contiguous_read_typed<T: TensorScalar>(tensor: &Tensor) -> crate::Result<&TypedTensor<T>> {
4784 tensor.as_typed::<T>().ok_or_else(|| {
4785 crate::Error::unsupported(
4786 "CudaBackend::to_contiguous_read",
4787 "an externally defined payload is not supported by this GPU operation",
4788 )
4789 })
4790}
4791
4792fn copy_read_typed<T: TensorScalar>(tensor: &Tensor) -> crate::Result<&TypedTensor<T>> {
4794 tensor.as_typed::<T>().ok_or_else(|| {
4795 crate::Error::unsupported(
4796 "copy_read_into",
4797 "an externally defined payload is not supported by this GPU operation",
4798 )
4799 })
4800}
4801
4802impl TensorDeviceTransfer for CudaBackend {
4803 fn download_to_host(&mut self, tensor: TensorRead<'_>) -> crate::Result<Tensor> {
4804 let tensor = tensor.as_tensor().ok_or_else(|| {
4805 crate::Error::unsupported(
4806 "CudaBackend::download_to_host",
4807 "CUDA transfer currently requires an owned tensor; materialize a view explicitly first",
4808 )
4809 })?;
4810 download_tensor(self.runtime(), tensor)
4811 }
4812
4813 fn upload_host_tensor(&mut self, tensor: TensorRead<'_>) -> crate::Result<Tensor> {
4814 let tensor = tensor.as_tensor().ok_or_else(|| {
4815 crate::Error::unsupported(
4816 "CudaBackend::upload_host_tensor",
4817 "CUDA transfer currently requires an owned tensor; materialize a view explicitly first",
4818 )
4819 })?;
4820 upload_tensor(self.runtime(), tensor)
4821 }
4822}
4823
4824#[allow(clippy::large_enum_variant)]
4832enum CompactOperand<'a, T> {
4833 Tensor(&'a TypedTensor<T>),
4834 View(TypedTensorView<'a, T>),
4835}
4836
4837impl<T: TensorScalar + Clone + 'static> CompactOperand<'_, T> {
4838 fn shape(&self) -> &[usize] {
4839 match self {
4840 Self::Tensor(tensor) => tensor.shape(),
4841 Self::View(view) => view.shape(),
4842 }
4843 }
4844
4845 fn ensure_resident(&self, rt: &CudaRuntime, op: &'static str) -> crate::Result<()> {
4846 match self {
4847 Self::Tensor(tensor) => dispatch::ensure_resident_on_runtime(rt, tensor, op),
4848 Self::View(view) => dispatch::ensure_view_resident_on_runtime(rt, view, op),
4849 }
4850 }
4851
4852 fn binding(&self, op: &'static str) -> crate::Result<TensorBinding<CubeclCudaRuntime>> {
4853 match self {
4854 Self::Tensor(tensor) => dispatch::typed_tensor_binding(tensor, op),
4855 Self::View(view) => dispatch::typed_view_binding(view, op),
4856 }
4857 }
4858}
4859
4860enum BroadcastMultiplyView<'a> {
4863 F32(TypedTensorView<'a, f32>),
4864 F64(TypedTensorView<'a, f64>),
4865 I32(TypedTensorView<'a, i32>),
4866 I64(TypedTensorView<'a, i64>),
4867 C32(TypedTensorView<'a, Complex32>),
4868 C64(TypedTensorView<'a, Complex64>),
4869}
4870
4871fn compact_view(view: TensorView<'_>) -> crate::Result<Option<BroadcastMultiplyView<'_>>> {
4877 fn usable<T: TensorScalar + Clone + 'static>(
4878 view: TypedTensorView<'_, T>,
4879 ) -> crate::Result<Option<TypedTensorView<'_, T>>> {
4880 Ok((view.offset() == 0 && view.is_col_major_contiguous()?).then_some(view))
4881 }
4882
4883 Ok(match view {
4884 TensorView::F32(view) => usable(view)?.map(BroadcastMultiplyView::F32),
4885 TensorView::F64(view) => usable(view)?.map(BroadcastMultiplyView::F64),
4886 TensorView::I32(view) => usable(view)?.map(BroadcastMultiplyView::I32),
4887 TensorView::I64(view) => usable(view)?.map(BroadcastMultiplyView::I64),
4888 TensorView::C32(view) => usable(view)?.map(BroadcastMultiplyView::C32),
4889 TensorView::C64(view) => usable(view)?.map(BroadcastMultiplyView::C64),
4890 TensorView::Bool(_) => None,
4891 })
4892}
4893
4894impl TensorBackend for CudaBackend {}
4895
4896fn validate_permutation(op: &'static str, perm: &[usize], rank: usize) -> crate::Result<()> {
4897 ensure_rank(op, rank, perm.len())?;
4898 ensure_axes_unique(op, "perm", perm, rank)
4899}
4900
4901fn ensure_same_shape_for_broadcast_multiply(
4902 lhs_shape: &[usize],
4903 rhs_shape: &[usize],
4904) -> crate::Result<()> {
4905 if lhs_shape != rhs_shape {
4906 return Err(crate::Error::shape_mismatch(
4907 "broadcast_multiply",
4908 lhs_shape.to_vec(),
4909 rhs_shape.to_vec(),
4910 ));
4911 }
4912 Ok(())
4913}
4914
4915fn launch_broadcast_multiply_typed<T>(
4916 backend: &CudaBackend,
4917 lhs: &CompactOperand<'_, T>,
4918 lhs_shape: &[usize],
4919 lhs_dims: &[usize],
4920 rhs: &CompactOperand<'_, T>,
4921 rhs_shape: &[usize],
4922 rhs_dims: &[usize],
4923) -> crate::Result<TypedTensor<T>>
4924where
4925 T: CubeElement + TensorScalar + CubePrimitive + CubeFloat + Clone,
4926{
4927 ensure_same_shape_for_broadcast_multiply(lhs_shape, rhs_shape)?;
4928 validate_broadcast_in_dim(lhs.shape(), lhs_shape, lhs_dims)?;
4929 validate_broadcast_in_dim(rhs.shape(), rhs_shape, rhs_dims)?;
4930 lhs.ensure_resident(backend.runtime(), "broadcast_multiply")?;
4931 rhs.ensure_resident(backend.runtime(), "broadcast_multiply")?;
4932 dispatch::launch_binary_bindings(
4933 backend.runtime(),
4934 lhs.binding("broadcast_multiply")?,
4935 rhs.binding("broadcast_multiply")?,
4936 lhs_shape,
4937 "broadcast_multiply",
4938 |client, count, dim, out, lhs_arg, rhs_arg| unsafe {
4939 elementwise::broadcast_multiply_float::launch_unchecked::<T, CubeclCudaRuntime>(
4940 client,
4941 count,
4942 dim,
4943 out.into_tensor_arg(),
4944 lhs_arg.into_tensor_arg(),
4945 rhs_arg.into_tensor_arg(),
4946 comptime_sequence(lhs_dims),
4947 comptime_sequence(rhs_dims),
4948 lhs_shape.len(),
4949 );
4950 },
4951 )
4952}
4953
4954fn launch_broadcast_multiply_int_typed<T>(
4955 backend: &CudaBackend,
4956 lhs: &CompactOperand<'_, T>,
4957 lhs_shape: &[usize],
4958 lhs_dims: &[usize],
4959 rhs: &CompactOperand<'_, T>,
4960 rhs_shape: &[usize],
4961 rhs_dims: &[usize],
4962) -> crate::Result<TypedTensor<T>>
4963where
4964 T: CubeElement + TensorScalar + CubePrimitive + CubeInt + Clone,
4965{
4966 ensure_same_shape_for_broadcast_multiply(lhs_shape, rhs_shape)?;
4967 validate_broadcast_in_dim(lhs.shape(), lhs_shape, lhs_dims)?;
4968 validate_broadcast_in_dim(rhs.shape(), rhs_shape, rhs_dims)?;
4969 lhs.ensure_resident(backend.runtime(), "broadcast_multiply")?;
4970 rhs.ensure_resident(backend.runtime(), "broadcast_multiply")?;
4971 dispatch::launch_binary_bindings(
4972 backend.runtime(),
4973 lhs.binding("broadcast_multiply")?,
4974 rhs.binding("broadcast_multiply")?,
4975 lhs_shape,
4976 "broadcast_multiply",
4977 |client, count, dim, out, lhs_arg, rhs_arg| unsafe {
4978 elementwise::broadcast_multiply_int::launch_unchecked::<T, CubeclCudaRuntime>(
4979 client,
4980 count,
4981 dim,
4982 out.into_tensor_arg(),
4983 lhs_arg.into_tensor_arg(),
4984 rhs_arg.into_tensor_arg(),
4985 comptime_sequence(lhs_dims),
4986 comptime_sequence(rhs_dims),
4987 lhs_shape.len(),
4988 );
4989 },
4990 )
4991}
4992
4993fn launch_broadcast_multiply_complex_typed<T>(
4994 backend: &CudaBackend,
4995 lhs: &CompactOperand<'_, T>,
4996 lhs_shape: &[usize],
4997 lhs_dims: &[usize],
4998 rhs: &CompactOperand<'_, T>,
4999 rhs_shape: &[usize],
5000 rhs_dims: &[usize],
5001) -> crate::Result<TypedTensor<T>>
5002where
5003 T: CubeElement + TensorScalar + CubePrimitive + CubeComplex + Clone,
5004{
5005 ensure_same_shape_for_broadcast_multiply(lhs_shape, rhs_shape)?;
5006 validate_broadcast_in_dim(lhs.shape(), lhs_shape, lhs_dims)?;
5007 validate_broadcast_in_dim(rhs.shape(), rhs_shape, rhs_dims)?;
5008 lhs.ensure_resident(backend.runtime(), "broadcast_multiply")?;
5009 rhs.ensure_resident(backend.runtime(), "broadcast_multiply")?;
5010 dispatch::launch_binary_bindings(
5011 backend.runtime(),
5012 lhs.binding("broadcast_multiply")?,
5013 rhs.binding("broadcast_multiply")?,
5014 lhs_shape,
5015 "broadcast_multiply",
5016 |client, count, dim, out, lhs_arg, rhs_arg| unsafe {
5017 elementwise::broadcast_multiply_complex::launch_unchecked::<T, CubeclCudaRuntime>(
5018 client,
5019 count,
5020 dim,
5021 out.into_tensor_arg(),
5022 lhs_arg.into_tensor_arg(),
5023 rhs_arg.into_tensor_arg(),
5024 comptime_sequence(lhs_dims),
5025 comptime_sequence(rhs_dims),
5026 lhs_shape.len(),
5027 );
5028 },
5029 )
5030}
5031
5032fn validate_broadcast_in_dim(
5033 input_shape: &[usize],
5034 shape: &[usize],
5035 dims: &[usize],
5036) -> crate::Result<()> {
5037 ensure_rank("broadcast_in_dim", input_shape.len(), dims.len())?;
5038 let mut seen = vec![false; shape.len()];
5039 for (src_axis, &dst_axis) in dims.iter().enumerate() {
5040 ensure_axis("broadcast_in_dim", dst_axis, shape.len())?;
5041 if seen[dst_axis] {
5042 return Err(crate::Error::duplicate_axis(
5043 "broadcast_in_dim",
5044 dst_axis,
5045 "dims",
5046 ));
5047 }
5048 seen[dst_axis] = true;
5049 let src = input_shape[src_axis];
5050 let dst = shape[dst_axis];
5051 if src != dst && src != 1 {
5052 return Err(crate::Error::shape_mismatch(
5053 "broadcast_in_dim",
5054 input_shape.to_vec(),
5055 shape.to_vec(),
5056 ));
5057 }
5058 }
5059 Ok(())
5060}
5061
5062fn extract_diagonal_shape(
5063 input_shape: &[usize],
5064 axis_a: usize,
5065 axis_b: usize,
5066) -> crate::Result<(Vec<usize>, usize)> {
5067 ensure_axis("extract_diagonal", axis_a, input_shape.len())?;
5068 ensure_axis("extract_diagonal", axis_b, input_shape.len())?;
5069 if axis_a == axis_b {
5070 return Err(crate::Error::duplicate_axis(
5071 "extract_diagonal",
5072 axis_a,
5073 "axes",
5074 ));
5075 }
5076 let diag_output_axis = if axis_a < axis_b { axis_a } else { axis_a - 1 };
5077 let diag_dim = input_shape[axis_a].min(input_shape[axis_b]);
5078 let mut output_shape = input_shape.to_vec();
5079 output_shape.remove(axis_b);
5080 output_shape[diag_output_axis] = diag_dim;
5081 Ok((output_shape, diag_output_axis))
5082}
5083
5084fn embed_diagonal_shape(
5085 input_shape: &[usize],
5086 axis_a: usize,
5087 axis_b: usize,
5088) -> crate::Result<Vec<usize>> {
5089 ensure_axis("embed_diagonal", axis_a, input_shape.len())?;
5090 if axis_b > input_shape.len() {
5091 return Err(crate::Error::axis_out_of_bounds(
5092 "embed_diagonal",
5093 axis_b,
5094 input_shape.len(),
5095 ));
5096 }
5097 let mut output_shape = input_shape.to_vec();
5098 output_shape.insert(axis_b, input_shape[axis_a]);
5099 Ok(output_shape)
5100}
5101
5102fn reduction_output_shape(input_shape: &[usize], axes: &[usize]) -> Vec<usize> {
5103 input_shape
5104 .iter()
5105 .enumerate()
5106 .filter_map(|(axis, &dim)| (!axes.contains(&axis)).then_some(dim))
5107 .collect()
5108}
5109
5110fn reduction_keepdims_shape(input_shape: &[usize], axis: usize) -> Vec<usize> {
5111 let mut output_shape = input_shape.to_vec();
5112 output_shape[axis] = 1;
5113 output_shape
5114}
5115
5116fn cubecl_reshape_metadata<T: crate::TensorScalar + Clone>(
5117 tensor: TypedTensor<T>,
5118 shape: Vec<usize>,
5119 op: &'static str,
5120) -> crate::Result<TypedTensor<T>> {
5121 let len = shape
5122 .iter()
5123 .try_fold(1usize, |acc, &dim| acc.checked_mul(dim))
5124 .ok_or_else(|| {
5125 crate::Error::invalid_argument(
5126 op,
5127 "shape",
5128 format!("shape product overflow for CubeCL reshape shape {shape:?}"),
5129 )
5130 })?;
5131 let tensor_len = tensor.n_elements();
5132 if len != tensor_len {
5133 return Err(crate::Error::validation(
5134 op,
5135 tenferro_tensor::ShapeMismatch::ReshapeElementCount {
5136 from: tensor_len,
5137 to: len,
5138 }
5139 .into(),
5140 ));
5141 }
5142
5143 let value = crate::TensorValue::from_tensor(
5148 <T as crate::TensorScalar>::typed_tensor_into_tensor(tensor),
5149 )
5150 .reshape_view(shape)?;
5151 let (group, slot, _, _) = value.try_into_group_parts().map_err(|_| {
5152 crate::Error::runtime_state(
5153 op,
5154 "failed to publish the reshaped tensor descriptor without copying",
5155 )
5156 })?;
5157 let tensor = group.into_tensor(slot).map_err(|(_, error)| {
5158 crate::Error::runtime_state(
5159 op,
5160 format!("failed to detach the reshaped tensor owner: {error}"),
5161 )
5162 })?;
5163 match <T as crate::TensorScalar>::into_typed(tensor) {
5164 Ok(typed) => Ok(typed),
5165 Err(failure) => Err(crate::Error::runtime_state(
5168 op,
5169 format!("reshaped owner lost its dtype guard: {}", failure.error()),
5170 )),
5171 }
5172}
5173
5174fn validate_slice(input_shape: &[usize], config: &SliceConfig) -> crate::Result<Vec<usize>> {
5175 let rank = input_shape.len();
5176 ensure_rank("slice", rank, config.starts.len())?;
5177 ensure_rank("slice", rank, config.limits.len())?;
5178 ensure_rank("slice", rank, config.strides.len())?;
5179 input_shape
5180 .iter()
5181 .enumerate()
5182 .map(|(axis, &dim)| {
5183 let start = config.starts[axis];
5184 let limit = config.limits[axis];
5185 let stride = config.strides[axis];
5186 if start > limit {
5187 return Err(crate::Error::invalid_argument(
5188 "slice",
5189 "bounds",
5190 format!("start exceeds limit on axis {axis}"),
5191 ));
5192 }
5193 if limit > dim {
5198 return Err(crate::Error::invalid_argument(
5199 "slice",
5200 "configuration",
5201 format!("limit {limit} on axis {axis} exceeds dimension size {dim}"),
5202 ));
5203 }
5204 if stride == 0 {
5205 return Err(crate::Error::invalid_argument(
5206 "slice",
5207 "strides",
5208 format!("stride must be positive on axis {axis}"),
5209 ));
5210 }
5211 let span = limit - start;
5212 Ok(span.div_ceil(stride))
5213 })
5214 .collect()
5215}
5216
5217fn pad_output_shape(input_shape: &[usize], config: &PadConfig) -> crate::Result<Vec<usize>> {
5218 let rank = input_shape.len();
5219 ensure_rank("pad", rank, config.edge_padding_low.len())?;
5220 ensure_rank("pad", rank, config.edge_padding_high.len())?;
5221 ensure_rank("pad", rank, config.interior_padding.len())?;
5222 let mut out_shape = Vec::with_capacity(rank);
5223 for (axis, &input_dim_raw) in input_shape.iter().enumerate().take(rank) {
5224 if config.interior_padding[axis] < 0 {
5225 return Err(crate::Error::invalid_argument(
5226 "pad",
5227 "interior_padding",
5228 format!("interior padding must be non-negative on axis {axis}"),
5229 ));
5230 }
5231 let input_dim = i64::try_from(input_dim_raw).map_err(|_| {
5232 crate::Error::invalid_argument(
5233 "pad",
5234 "input_shape",
5235 format!("input dimension on axis {axis} must fit in i64"),
5236 )
5237 })?;
5238 let base = if input_dim == 0 {
5239 0
5240 } else {
5241 let spacing = config.interior_padding[axis]
5242 .checked_add(1)
5243 .ok_or_else(|| {
5244 crate::Error::invalid_argument(
5245 "pad",
5246 "interior_padding",
5247 format!("interior padding overflow on axis {axis}"),
5248 )
5249 })?;
5250 input_dim
5251 .checked_sub(1)
5252 .and_then(|extent| extent.checked_mul(spacing))
5253 .and_then(|extent| extent.checked_add(1))
5254 .ok_or_else(|| {
5255 crate::Error::invalid_argument(
5256 "pad",
5257 "interior_padding",
5258 format!("padded interior extent overflow on axis {axis}"),
5259 )
5260 })?
5261 };
5262 let dim = config.edge_padding_low[axis]
5263 .checked_add(config.edge_padding_high[axis])
5264 .and_then(|edge| edge.checked_add(base))
5265 .ok_or_else(|| {
5266 crate::Error::invalid_argument(
5267 "pad",
5268 "padding",
5269 format!("output dimension overflow on axis {axis}"),
5270 )
5271 })?;
5272 out_shape.push(usize::try_from(dim).map_err(|_| {
5273 crate::Error::invalid_argument(
5274 "pad",
5275 "padding",
5276 format!("negative output dimension on axis {axis}"),
5277 )
5278 })?);
5279 }
5280 Ok(out_shape)
5281}
5282
5283fn validate_slice_sizes_within_operand(
5284 op: &'static str,
5285 operand_shape: &[usize],
5286 slice_sizes: &[usize],
5287) -> crate::Result<()> {
5288 ensure_rank(op, operand_shape.len(), slice_sizes.len())?;
5289 for (axis, (&slice_size, &dim_size)) in slice_sizes.iter().zip(operand_shape).enumerate() {
5290 if slice_size > dim_size {
5291 return Err(crate::Error::invalid_argument(
5292 op,
5293 "slice_sizes",
5294 format!("slice_sizes[{axis}]={slice_size} exceeds operand dimension {dim_size}"),
5295 ));
5296 }
5297 }
5298 Ok(())
5299}
5300
5301fn index_vector_size(shape: &[usize], index_vector_dim: usize) -> usize {
5302 if index_vector_dim == shape.len() {
5303 1
5304 } else {
5305 shape[index_vector_dim]
5306 }
5307}
5308
5309fn index_batch_shape(shape: &[usize], index_vector_dim: usize) -> Vec<usize> {
5310 if index_vector_dim == shape.len() {
5311 return shape.to_vec();
5312 }
5313 shape
5314 .iter()
5315 .enumerate()
5316 .filter_map(|(axis, &dim)| (axis != index_vector_dim).then_some(dim))
5317 .collect()
5318}
5319
5320fn operand_window_dims(rank: usize, collapsed_or_inserted: &[usize]) -> Vec<usize> {
5321 (0..rank)
5322 .filter(|dim| !collapsed_or_inserted.contains(dim))
5323 .collect()
5324}
5325
5326#[derive(Debug)]
5327struct GatherLaunchMeta {
5328 output_shape: Vec<usize>,
5329 window_dims: Vec<usize>,
5330}
5331
5332fn gather_launch_meta(
5333 operand_shape: &[usize],
5334 start_indices_shape: &[usize],
5335 config: &GatherConfig,
5336) -> crate::Result<GatherLaunchMeta> {
5337 ensure_rank("gather", operand_shape.len(), config.slice_sizes.len())?;
5338 validate_slice_sizes_within_operand("gather", operand_shape, &config.slice_sizes)?;
5339 if config.index_vector_dim > start_indices_shape.len() {
5340 return Err(crate::Error::axis_out_of_bounds(
5341 "gather",
5342 config.index_vector_dim,
5343 start_indices_shape.len(),
5344 ));
5345 }
5346 let index_size = index_vector_size(start_indices_shape, config.index_vector_dim);
5347 if index_size != config.start_index_map.len() {
5348 return Err(crate::Error::invalid_argument(
5349 "gather",
5350 "start_index_map",
5351 "start_index_map length mismatch",
5352 ));
5353 }
5354 ensure_axes_unique(
5355 "gather",
5356 "collapsed_slice_dims",
5357 &config.collapsed_slice_dims,
5358 operand_shape.len(),
5359 )?;
5360 for &dim in &config.collapsed_slice_dims {
5361 if config.slice_sizes[dim] != 1 {
5362 return Err(crate::Error::invalid_argument(
5363 "gather",
5364 "collapsed_slice_dims",
5365 format!(
5366 "collapsed slice dimension {dim} must have slice_size == 1, got {}",
5367 config.slice_sizes[dim]
5368 ),
5369 ));
5370 }
5371 }
5372 ensure_axes_unique(
5373 "gather",
5374 "start_index_map",
5375 &config.start_index_map,
5376 operand_shape.len(),
5377 )?;
5378 let window_dims = operand_window_dims(operand_shape.len(), &config.collapsed_slice_dims);
5379 if config.offset_dims.len() != window_dims.len() {
5380 return Err(crate::Error::invalid_argument(
5381 "gather",
5382 "offset_dims",
5383 "offset_dims length mismatch",
5384 ));
5385 }
5386 let batch_shape = index_batch_shape(start_indices_shape, config.index_vector_dim);
5387 let out_rank = batch_shape.len() + config.offset_dims.len();
5388 ensure_axes_unique("gather", "offset_dims", &config.offset_dims, out_rank)?;
5389 let mut output_shape = vec![0usize; out_rank];
5390 let mut out_axis_to_operand_dim = vec![None; out_rank];
5391 for (offset_axis, &out_axis) in config.offset_dims.iter().enumerate() {
5392 out_axis_to_operand_dim[out_axis] = Some(window_dims[offset_axis]);
5393 }
5394 let mut batch_axis = 0usize;
5395 for out_axis in 0..out_rank {
5396 if let Some(operand_dim) = out_axis_to_operand_dim[out_axis] {
5397 output_shape[out_axis] = config.slice_sizes[operand_dim];
5398 } else {
5399 output_shape[out_axis] = batch_shape[batch_axis];
5400 batch_axis += 1;
5401 }
5402 }
5403 Ok(GatherLaunchMeta {
5404 output_shape,
5405 window_dims,
5406 })
5407}
5408
5409#[derive(Debug)]
5410struct ScatterLaunchMeta {
5411 batch_shape: Vec<usize>,
5412 window_dims: Vec<usize>,
5413 window_shape_updates: Vec<usize>,
5414}
5415
5416fn scatter_launch_meta(
5417 operand_shape: &[usize],
5418 scatter_indices_shape: &[usize],
5419 updates_shape: &[usize],
5420 config: &ScatterConfig,
5421) -> crate::Result<ScatterLaunchMeta> {
5422 if config.index_vector_dim > scatter_indices_shape.len() {
5423 return Err(crate::Error::axis_out_of_bounds(
5424 "scatter",
5425 config.index_vector_dim,
5426 scatter_indices_shape.len(),
5427 ));
5428 }
5429 let index_size = index_vector_size(scatter_indices_shape, config.index_vector_dim);
5430 if index_size != config.scatter_dims_to_operand_dims.len() {
5431 return Err(crate::Error::invalid_argument(
5432 "scatter",
5433 "scatter_dims_to_operand_dims",
5434 "scatter_dims_to_operand_dims length mismatch",
5435 ));
5436 }
5437 ensure_axes_unique(
5438 "scatter",
5439 "inserted_window_dims",
5440 &config.inserted_window_dims,
5441 operand_shape.len(),
5442 )?;
5443 ensure_axes_unique(
5444 "scatter",
5445 "scatter_dims_to_operand_dims",
5446 &config.scatter_dims_to_operand_dims,
5447 operand_shape.len(),
5448 )?;
5449 ensure_axes_unique(
5450 "scatter",
5451 "update_window_dims",
5452 &config.update_window_dims,
5453 updates_shape.len(),
5454 )?;
5455 let batch_shape = index_batch_shape(scatter_indices_shape, config.index_vector_dim);
5456 let window_dims = operand_window_dims(operand_shape.len(), &config.inserted_window_dims);
5457 if config.update_window_dims.len() != window_dims.len() {
5458 return Err(crate::Error::invalid_argument(
5459 "scatter",
5460 "update_window_dims",
5461 "update_window_dims length mismatch",
5462 ));
5463 }
5464 let updates_batch_rank = updates_shape.len() - config.update_window_dims.len();
5465 if updates_batch_rank != batch_shape.len() {
5466 return Err(crate::Error::rank_mismatch(
5467 "scatter",
5468 batch_shape.len(),
5469 updates_batch_rank,
5470 ));
5471 }
5472 let mut is_update_window_dim = vec![false; updates_shape.len()];
5473 for &axis in &config.update_window_dims {
5474 is_update_window_dim[axis] = true;
5475 }
5476 let mut batch_axis = 0usize;
5477 for (axis, &actual) in updates_shape.iter().enumerate() {
5478 if is_update_window_dim[axis] {
5479 continue;
5480 }
5481 let expected = batch_shape[batch_axis];
5482 if actual != expected {
5483 return Err(crate::Error::shape_mismatch(
5484 "scatter",
5485 vec![expected],
5486 vec![actual],
5487 ));
5488 }
5489 batch_axis += 1;
5490 }
5491 let window_shape_updates = config
5492 .update_window_dims
5493 .iter()
5494 .map(|&axis| updates_shape[axis])
5495 .collect();
5496 Ok(ScatterLaunchMeta {
5497 batch_shape,
5498 window_dims,
5499 window_shape_updates,
5500 })
5501}
5502
5503fn concatenate_output_shape<T>(
5504 inputs: &[&TypedTensor<T>],
5505 axis: usize,
5506) -> crate::Result<Vec<usize>> {
5507 let first = inputs[0];
5508 let rank = first.shape().len();
5509 ensure_axis("concatenate", axis, rank)?;
5510 let mut out_shape = first.shape().to_vec();
5511 let mut axis_extent = 0usize;
5512 for input in inputs {
5513 ensure_rank("concatenate", rank, input.shape().len())?;
5514 for dim in 0..rank {
5515 if dim == axis {
5516 axis_extent = axis_extent.checked_add(input.shape()[dim]).ok_or_else(|| {
5517 crate::Error::invalid_argument(
5518 "concatenate",
5519 "shape",
5520 "concatenate axis extent overflows usize",
5521 )
5522 })?;
5523 } else if input.shape()[dim] != first.shape()[dim] {
5524 return Err(crate::Error::shape_mismatch(
5525 "concatenate",
5526 first.shape().to_vec(),
5527 input.shape().to_vec(),
5528 ));
5529 }
5530 }
5531 }
5532 out_shape[axis] = axis_extent;
5533 Ok(out_shape)
5534}
5535
5536#[cfg(test)]
5537mod tests;