Skip to main content

tenferro_gpu/cubecl/
mod.rs

1//! CubeCL-based GPU backend for tenferro tensors.
2//!
3//! This module provides GPU acceleration via [CubeCL](https://github.com/tracel-ai/cubecl)
4//! running on NVIDIA CUDA devices. It is gated behind the `cuda` feature flag and
5//! requires **CUDA 12.8+** with a compatible NVIDIA GPU.
6//!
7//! # Enabling the feature
8//!
9//! Add to your `Cargo.toml`:
10//!
11//! ```toml
12//! tenferro-gpu = { version = "...", features = ["cuda"] }
13//! ```
14//!
15//! You must also enable a CPU backend (`cpu-faer` or `cpu-blas`); the CubeCL backend
16//! complements the CPU path but does not replace it.
17//!
18//! # Prerequisites
19//!
20//! - NVIDIA GPU with CUDA compute capability ≥ 7.0
21//! - CUDA Toolkit 12.8 or newer installed (provides NVRTC for JIT kernel compilation)
22//! - cuTENSOR shared library available on `LD_LIBRARY_PATH`
23//!
24//! ## Environment variables
25//!
26//! | Variable | Purpose |
27//! |----------|---------|
28//! | `CUDA_PATH` | CUDA toolkit root (e.g. `/usr/local/cuda-12.8`) |
29//! | `CUBECL_DEBUG_LOG` | Set to `0` to suppress verbose JIT logs |
30//! | `TENFERRO_CUTENSOR_PATH` | Override cuTENSOR library search path |
31//!
32//! # Basic usage
33//!
34//! GPU tensors must be explicitly uploaded before use on the device and downloaded
35//! back to the host afterwards (no implicit CPU↔GPU transfer, following the PyTorch
36//! convention).
37//!
38//! ```rust
39//! use tenferro_gpu::{cuda::cuda_devices, cuda::CudaBackend, cuda::CudaDeviceError};
40//!
41//! fn first_cuda_backend() -> Result<Option<CudaBackend>, CudaDeviceError> {
42//!     let devices = cuda_devices()?;
43//!     let Some(device) = devices.first() else {
44//!         return Ok(None);
45//!     };
46//!     Ok(Some(CudaBackend::new(device.id())?))
47//! }
48//!
49//! let _example: fn() -> Result<Option<CudaBackend>, CudaDeviceError> = first_cuda_backend;
50//! ```
51//!
52//! # Running GPU tests
53//!
54//! All GPU tests are marked `#[ignore]` so that `cargo test --features cuda`
55//! passes on machines without a GPU. To actually run them:
56//!
57//! ```sh
58//! CUBECL_DEBUG_LOG=0 \
59//! CUDA_PATH=/usr/local/cuda-12.8 \
60//! cargo test -p tenferro-gpu --features cuda -- --ignored
61//! ```
62
63use 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
102/// The Rust scalar type behind a preset variant name a macro received.
103macro_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                    // SAFETY: The tiled classification proves a compact 2D
275                    // transpose. Bounds guards cover edge tiles and every unit
276                    // reaches the shared-memory barrier.
277                    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        // SAFETY: `NativePermutationPlan` validated both allocation ranges,
301        // destination non-overlap, and disjoint source/destination storage.
302        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    // CUDA library handles are dropped before `rt`; Rust drops fields in
343    // declaration order, so cache-owned handles release while the CUDA primary
344    // context is still retained by `CudaRuntime`.
345    cutensor: OnceLock<ffi::cutensor::CutensorHandle>,
346    extension_cache: CudaExtensionCache,
347    // Backend-level so the configured cap survives clearing or evicting the
348    // extension-cache entry that owns the shared scratch pool itself.
349    cutensor_workspace_max_retained_bytes: AtomicU64,
350    // Backend-level for the same reason, and cumulative across cache clears:
351    // it is a diagnostic for whether the cap is set below the workload's real
352    // requirement, not a per-cache statistic.
353    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/// Type-indexed cache for CUDA extension-owned backend state.
368#[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
385/// Default cap on retained shared cuTENSOR contraction scratch, in bytes.
386///
387/// This bounds only the scratch the backend keeps for reuse, not total device
388/// memory: a contraction whose requirement exceeds the remaining cap still runs
389/// in a temporary workspace. It is deliberately permissive, because there is no
390/// single optimal cap across workloads and the cap never affects correctness;
391/// callers that need to bound retained device memory configure a smaller value.
392/// On a device with less free memory than this the cap simply never binds.
393const 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    /// Create an empty extension cache.
481    ///
482    /// # Examples
483    ///
484    /// ```
485    /// use tenferro_gpu::cuda::CudaExtensionCache;
486    ///
487    /// let cache = CudaExtensionCache::new();
488    /// assert!(cache.is_empty()?);
489    /// # Ok::<(), tenferro_tensor::Error>(())
490    /// ```
491    ///
492    /// The cache retains at most 16 extension states by default. Use
493    /// [`Self::with_max_entries`] to choose a different bound. Later cache
494    /// operations return [`crate::Error::RuntimeState`] if the cache mutex is
495    /// poisoned.
496    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    /// Create an empty extension cache with an explicit entry bound.
503    pub fn with_max_entries(max_entries: NonZeroUsize) -> Self {
504        Self {
505            inner: Mutex::new(CudaExtensionCacheInner::new(max_entries)),
506        }
507    }
508
509    /// Returns `true` when no extension state has been initialized.
510    ///
511    /// # Examples
512    ///
513    /// ```
514    /// use tenferro_gpu::cuda::CudaExtensionCache;
515    ///
516    /// assert!(CudaExtensionCache::new().is_empty()?);
517    /// # Ok::<(), tenferro_tensor::Error>(())
518    /// ```
519    /// # Errors
520    ///
521    /// Returns [`crate::Error::RuntimeState`] if the cache mutex is poisoned.
522    pub fn is_empty(&self) -> crate::Result<bool> {
523        Ok(self.lock_inner()?.entries.is_empty())
524    }
525
526    /// Remove every cached CUDA extension state value.
527    ///
528    /// This operation returns a runtime-state error if the cache mutex is
529    /// poisoned.
530    ///
531    /// # Errors
532    ///
533    /// Returns [`crate::Error::RuntimeState`] if the cache mutex is poisoned.
534    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    /// Snapshot the number of retained entries and logical retained bytes.
548    /// # Errors
549    ///
550    /// Returns [`crate::Error::RuntimeState`] if the cache mutex is poisoned.
551    pub fn stats(&self) -> crate::Result<CacheStats> {
552        let inner = self.lock_inner()?;
553        Ok(inner.snapshot_stats())
554    }
555
556    /// Return the configured entry bound.
557    /// # Errors
558    ///
559    /// Returns [`crate::Error::RuntimeState`] if the cache mutex is poisoned.
560    pub fn max_entries(&self) -> crate::Result<NonZeroUsize> {
561        Ok(self.lock_inner()?.max_entries)
562    }
563
564    /// Return the configured logical retained-byte bound.
565    /// # Errors
566    ///
567    /// Returns [`crate::Error::RuntimeState`] if the cache mutex is poisoned.
568    pub fn max_retained_bytes(&self) -> crate::Result<NonZeroUsize> {
569        Ok(self.lock_inner()?.max_retained_bytes)
570    }
571
572    /// Replace the entry bound and evict oldest entries if needed.
573    /// # Errors
574    ///
575    /// Returns [`crate::Error::RuntimeState`] if the cache mutex is poisoned
576    /// while changing the bound.
577    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    /// Configure the logical retained-byte bound and evict oldest entries if
585    /// needed.
586    /// # Errors
587    ///
588    /// Returns [`crate::Error::RuntimeState`] if the cache mutex is poisoned
589    /// while changing the bound.
590    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    /// Get or lazily initialize one cache entry keyed by `T`.
598    ///
599    /// # Examples
600    ///
601    /// ```
602    /// use tenferro_gpu::cuda::CudaExtensionCache;
603    ///
604    /// let cache = CudaExtensionCache::new();
605    /// let value = cache.get_or_try_init::<usize>(|| Ok(3)).unwrap();
606    /// assert_eq!(*value, 3);
607    /// ```
608    /// # Errors
609    ///
610    /// Propagates the initializer's typed error, returns
611    /// [`crate::Error::RuntimeState`] for a poisoned cache or a missing/wrongly
612    /// typed entry, and preserves backend errors from initialization.
613    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    /// Update the logical retained-byte estimate for an existing typed entry.
673    ///
674    /// This supports extension states whose own internal cache grows after the
675    /// top-level entry is initialized. If another thread clears or evicts the
676    /// typed entry before the update, the update is treated as a no-op.
677    /// # Errors
678    ///
679    /// Returns [`crate::Error::RuntimeState`] if the cache mutex is poisoned.
680    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/// Borrow guard for one cached CUDA extension state value.
708#[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        // SAFETY: get_or_try_init validates the downcast while holding this
736        // same mutex guard. The entry cannot move or be evicted while this
737        // guard owns the mutex.
738        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        // Same-dtype casts are explicit copies. Use the native materialization
748        // path so an identity copy preserves NaN payloads instead of routing
749        // through cuTENSOR's alpha-scaled permutation operation.
750        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    /// Create a new CubeCL backend for the caller-selected CUDA device.
776    ///
777    /// # Examples
778    ///
779    /// ```
780    /// use tenferro_gpu::{cuda::CudaBackend, cuda::CudaDeviceError, cuda::CudaDeviceId};
781    ///
782    /// let _ctor: fn(CudaDeviceId) -> Result<CudaBackend, CudaDeviceError> = CudaBackend::new;
783    /// ```
784    /// # Errors
785    ///
786    /// Returns [`CudaDeviceError::Discovery`] when device discovery fails,
787    /// [`CudaDeviceError::Unavailable`] when the selected device is not
788    /// discovered, or [`CudaDeviceError::Initialization`] when CUDA runtime,
789    /// context, or CubeCL client initialization fails.
790    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    /// Borrow the underlying CubeCL runtime.
805    ///
806    /// # Examples
807    ///
808    /// ```
809    /// use tenferro_gpu::{cuda::CudaBackend, cuda::CudaRuntime};
810    ///
811    /// let _runtime: fn(&CudaBackend) -> &CudaRuntime = CudaBackend::runtime;
812    /// ```
813    pub fn runtime(&self) -> &CudaRuntime {
814        &self.inner.rt
815    }
816
817    /// Return the caller-selected CUDA device identity used by this backend.
818    ///
819    /// # Examples
820    ///
821    /// ```
822    /// use tenferro_gpu::{cuda::CudaBackend, cuda::CudaDeviceId};
823    ///
824    /// let _device_id: fn(&CudaBackend) -> CudaDeviceId = CudaBackend::device_id;
825    /// ```
826    pub fn device_id(&self) -> CudaDeviceId {
827        self.inner.rt.device_id()
828    }
829
830    /// Return the opaque identity of this exact executable backend instance.
831    ///
832    /// Clones of a backend return the same identity. Independently constructed
833    /// backends return different identities even when they target the same
834    /// CUDA device ordinal.
835    ///
836    /// # Examples
837    ///
838    /// ```
839    /// use tenferro_gpu::cuda::CudaBackend;
840    ///
841    /// let _identity = CudaBackend::runtime_identity;
842    /// ```
843    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    /// Clear CUDA extension-owned backend state.
868    ///
869    /// # Errors
870    ///
871    /// Returns [`crate::Error::RuntimeState`] if the extension cache mutex is
872    /// poisoned.
873    pub fn clear_cuda_extension_cache(&self) -> crate::Result<()> {
874        self.inner.extension_cache.clear()
875    }
876
877    /// Return CUDA extension cache stats.
878    ///
879    /// # Errors
880    ///
881    /// Returns [`crate::Error::RuntimeState`] if the extension cache mutex is
882    /// poisoned.
883    pub fn cuda_extension_cache_stats(&self) -> crate::Result<CacheStats> {
884        self.inner.extension_cache.stats()
885    }
886
887    /// Return the CUDA extension cache entry bound.
888    ///
889    /// # Errors
890    ///
891    /// Returns [`crate::Error::RuntimeState`] if the extension cache mutex is
892    /// poisoned.
893    pub fn cuda_extension_cache_max_entries(&self) -> crate::Result<NonZeroUsize> {
894        self.inner.extension_cache.max_entries()
895    }
896
897    /// Return the CUDA extension cache logical retained-byte bound.
898    ///
899    /// # Errors
900    ///
901    /// Returns [`crate::Error::RuntimeState`] if the extension cache mutex is
902    /// poisoned.
903    pub fn cuda_extension_cache_max_retained_bytes(&self) -> crate::Result<NonZeroUsize> {
904        self.inner.extension_cache.max_retained_bytes()
905    }
906
907    /// Configure the CUDA extension cache entry bound.
908    ///
909    /// # Errors
910    ///
911    /// Returns [`crate::Error::RuntimeState`] if the extension cache mutex is
912    /// poisoned while changing the bound.
913    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    /// Configure the CUDA extension cache logical retained-byte bound.
921    ///
922    /// # Errors
923    ///
924    /// Returns [`crate::Error::RuntimeState`] if the extension cache mutex is
925    /// poisoned while changing the bound.
926    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    /// Return cuTENSOR contraction plan cache stats.
936    ///
937    /// The returned entry count is the number of retained cuTENSOR contraction
938    /// plans inside the CUDA backend's extension cache entry. Logical retained
939    /// bytes cover plan metadata, not the shared per-stream device scratch;
940    /// [`CudaBackend::cutensor_workspace_stats`] reports that separately. The
941    /// cache byte limit is therefore not a total device-memory limit.
942    /// # Errors
943    ///
944    /// Returns [`crate::Error::RuntimeState`] if the cache mutex is poisoned.
945    pub fn cutensor_plan_cache_stats(&self) -> crate::Result<CacheStats> {
946        gemm::cutensor_plan_cache_stats(self)
947    }
948
949    /// Return the retained shared cuTENSOR contraction scratch, in bytes.
950    ///
951    /// All cached cuTENSOR contraction plans share one lazily grown workspace
952    /// per physical stream slot, so this is the sum over slots of the capacity
953    /// each slot currently holds. It is bounded by
954    /// [`CudaBackend::set_cutensor_workspace_max_retained_bytes`], reported
955    /// separately from the extension-cache byte statistics, and released by
956    /// clearing the extension cache or dropping the backend. This is retained
957    /// scratch, not total device memory: a workspace in use by a queued
958    /// contraction, a retiring allocation, and vendor-internal memory are all
959    /// excluded.
960    ///
961    /// Read this to size a retention cap: it is the high-water demand of the
962    /// workload shapes that have run so far.
963    ///
964    /// # Examples
965    ///
966    /// ```
967    /// use tenferro_gpu::cuda::{cuda_devices, gpu_available, CudaBackend};
968    ///
969    /// # fn main() -> Result<(), Box<dyn std::error::Error>> {
970    /// // `gpu_available` never panics without a CUDA driver, so this
971    /// // example also runs in CPU-only doctest environments.
972    /// if gpu_available() {
973    ///     let device = cuda_devices()?.remove(0);
974    ///     let backend = CudaBackend::new(device.id())?;
975    ///     let stats = backend.cutensor_workspace_stats()?;
976    ///     println!("{:?}", (stats.retained_entries, stats.retained_bytes));
977    /// }
978    /// # Ok(())
979    /// # }
980    /// ```
981    /// # Errors
982    ///
983    /// Returns [`crate::Error::RuntimeState`] if the cache mutex is poisoned.
984    pub fn cutensor_workspace_stats(&self) -> crate::Result<CutensorWorkspaceStats> {
985        gemm::cutensor_workspace_stats(self)
986    }
987
988    /// Return the device bytes retained by the shared cuTENSOR contraction
989    /// scratch. Equal to
990    /// [`CudaBackend::cutensor_workspace_stats`]`().retained_bytes`.
991    ///
992    /// # Examples
993    ///
994    /// ```
995    /// use tenferro_gpu::cuda::{cuda_devices, gpu_available, CudaBackend};
996    ///
997    /// # fn main() -> Result<(), Box<dyn std::error::Error>> {
998    /// // `gpu_available` never panics without a CUDA driver, so this
999    /// // example also runs in CPU-only doctest environments.
1000    /// if gpu_available() {
1001    ///     let device = cuda_devices()?.remove(0);
1002    ///     let backend = CudaBackend::new(device.id())?;
1003    ///     println!("{} bytes retained", backend.cutensor_workspace_bytes()?);
1004    /// }
1005    /// # Ok(())
1006    /// # }
1007    /// ```
1008    /// # Errors
1009    ///
1010    /// Returns [`crate::Error::RuntimeState`] if the cache mutex is poisoned.
1011    pub fn cutensor_workspace_bytes(&self) -> crate::Result<u64> {
1012        Ok(self.cutensor_workspace_stats()?.retained_bytes)
1013    }
1014
1015    /// Return the configured retention cap for shared cuTENSOR contraction
1016    /// scratch, in bytes.
1017    ///
1018    /// The default is 10 GiB. See
1019    /// [`CudaBackend::set_cutensor_workspace_max_retained_bytes`] for the
1020    /// contract; this value is not a device-memory reservation. The cap is
1021    /// plain backend state, so reading it cannot fail and never creates cache
1022    /// state.
1023    ///
1024    /// # Examples
1025    ///
1026    /// ```
1027    /// use tenferro_gpu::cuda::{cuda_devices, gpu_available, CudaBackend};
1028    ///
1029    /// # fn main() -> Result<(), Box<dyn std::error::Error>> {
1030    /// // `gpu_available` never panics without a CUDA driver, so this
1031    /// // example also runs in CPU-only doctest environments.
1032    /// if gpu_available() {
1033    ///     let device = cuda_devices()?.remove(0);
1034    ///     let backend = CudaBackend::new(device.id())?;
1035    ///     println!("cap {} bytes", backend.cutensor_workspace_max_retained_bytes());
1036    /// }
1037    /// # Ok(())
1038    /// # }
1039    /// ```
1040    pub fn cutensor_workspace_max_retained_bytes(&self) -> u64 {
1041        self.cutensor_workspace_limit()
1042    }
1043
1044    /// Configure the retention cap for shared cuTENSOR contraction scratch.
1045    ///
1046    /// The cap bounds how much scratch the backend keeps for reuse, summed over
1047    /// physical stream slots. It never refuses a contraction: a requirement that
1048    /// does not fit the remaining cap runs in a temporary workspace that is
1049    /// released afterwards, and shrinking the cap drops retained buffers
1050    /// without evicting any cached plan. Other slots are never evicted to make
1051    /// room.
1052    ///
1053    /// `0` disables retention entirely; it is not "unlimited". The default
1054    /// (10 GiB) is finite but is not a practical memory protection, and neither
1055    /// the cap nor the reported statistics bound total device memory.
1056    ///
1057    /// Setting a cap below the steady-state working set makes matching
1058    /// contractions allocate and retire their scratch on every call, which can
1059    /// increase workspace-retirement stream barrier fallbacks. To choose a
1060    /// value, run the workload and read the
1061    /// [`CudaBackend::cutensor_workspace_bytes`] high-water: retaining every
1062    /// slot's rounded high-water needs a cap of at least the sum of
1063    /// `next_power_of_two(max(request, 1 MiB))` over the stream slots.
1064    ///
1065    /// The setting is stored on the backend, so it survives
1066    /// `CudaBackend::clear_cuda_extension_cache` and extension-cache eviction,
1067    /// and is shared by clones of this backend.
1068    ///
1069    /// # Examples
1070    ///
1071    /// ```
1072    /// use tenferro_gpu::cuda::{cuda_devices, gpu_available, CudaBackend};
1073    ///
1074    /// # fn main() -> Result<(), Box<dyn std::error::Error>> {
1075    /// // `gpu_available` never panics without a CUDA driver, so this
1076    /// // example also runs in CPU-only doctest environments.
1077    /// if gpu_available() {
1078    ///     let device = cuda_devices()?.remove(0);
1079    ///     let backend = CudaBackend::new(device.id())?;
1080    ///     backend.set_cutensor_workspace_max_retained_bytes(4 << 30)?;
1081    /// }
1082    /// # Ok(())
1083    /// # }
1084    /// ```
1085    /// # Errors
1086    ///
1087    /// Returns [`crate::Error::RuntimeState`] if the plan-cache mutex is
1088    /// poisoned while releasing retained buffers.
1089    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    /// Current retention cap for shared cuTENSOR contraction scratch.
1097    fn cutensor_workspace_limit(&self) -> u64 {
1098        self.inner
1099            .cutensor_workspace_max_retained_bytes
1100            .load(Ordering::Relaxed)
1101    }
1102
1103    /// Return how many contractions ran in a temporary shared-scratch
1104    /// workspace because their requirement did not fit the retention cap.
1105    ///
1106    /// This is the direct signal that the cap is binding. A nonzero value means
1107    /// the matching contractions allocated and retired their scratch on every
1108    /// call instead of reusing a retained buffer; the high-water from
1109    /// [`CudaBackend::cutensor_workspace_bytes`] then under-reports the real
1110    /// requirement. Raise
1111    /// [`CudaBackend::set_cutensor_workspace_max_retained_bytes`] until this
1112    /// stops increasing, or accept the churn deliberately.
1113    ///
1114    /// The count is cumulative for the backend, shared by clones, and is not
1115    /// reset by `CudaBackend::clear_cuda_extension_cache`; diff two reads to
1116    /// measure an interval. Reading it cannot fail and never creates cache
1117    /// state.
1118    ///
1119    /// # Examples
1120    ///
1121    /// ```
1122    /// use tenferro_gpu::cuda::{cuda_devices, gpu_available, CudaBackend};
1123    ///
1124    /// # fn main() -> Result<(), Box<dyn std::error::Error>> {
1125    /// // `gpu_available` never panics without a CUDA driver, so this
1126    /// // example also runs in CPU-only doctest environments.
1127    /// if gpu_available() {
1128    ///     let device = cuda_devices()?.remove(0);
1129    ///     let backend = CudaBackend::new(device.id())?;
1130    ///     println!("{} temporary uses", backend.cutensor_workspace_temporary_uses());
1131    /// }
1132    /// # Ok(())
1133    /// # }
1134    /// ```
1135    pub fn cutensor_workspace_temporary_uses(&self) -> u64 {
1136        self.inner
1137            .cutensor_workspace_temporary_uses
1138            .load(Ordering::Relaxed)
1139    }
1140
1141    /// Record one cap-driven temporary workspace use.
1142    fn note_cutensor_temporary_workspace(&self) {
1143        self.inner
1144            .cutensor_workspace_temporary_uses
1145            .fetch_add(1, Ordering::Relaxed);
1146    }
1147
1148    /// Return deferred cuTENSOR workspace retirement counters.
1149    ///
1150    /// Retirements are deferred until the workspace's stream reaches the event
1151    /// recorded at retirement time. `in_flight` is the number of workspaces
1152    /// whose handle has not returned to the CubeCL pool yet.
1153    ///
1154    /// # Errors
1155    ///
1156    /// Returns [`crate::Error::RuntimeState`] if the retirement queue lock is
1157    /// poisoned.
1158    pub fn cutensor_workspace_retirement_stats(&self) -> crate::Result<WorkspaceRetirementStats> {
1159        gemm::cutensor_workspace_retirement_stats(self)
1160    }
1161
1162    /// Return the cuTENSOR contraction plan entry bound.
1163    /// # Errors
1164    ///
1165    /// Returns [`crate::Error::RuntimeState`] if the cache mutex is poisoned.
1166    pub fn cutensor_plan_cache_max_entries(&self) -> crate::Result<NonZeroUsize> {
1167        gemm::cutensor_plan_cache_max_entries(self)
1168    }
1169
1170    /// Configure the cuTENSOR contraction plan entry bound.
1171    ///
1172    /// The cache is initialized if it does not already exist so a setting made
1173    /// before the first CUDA `dot_general` call is preserved.
1174    /// # Errors
1175    ///
1176    /// Returns [`crate::Error::RuntimeState`] if the cache mutex is poisoned.
1177    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    /// Return cuTENSOR structural permutation plan cache stats.
1185    ///
1186    /// The returned entry count is the number of retained cuTENSOR permutation
1187    /// plans inside the CUDA backend's extension cache entry. Logical retained
1188    /// bytes include cached descriptor and plan state.
1189    /// # Errors
1190    ///
1191    /// Returns [`crate::Error::RuntimeState`] if the cache mutex is poisoned.
1192    pub fn cutensor_permutation_plan_cache_stats(&self) -> crate::Result<CacheStats> {
1193        permutation::cutensor_permutation_plan_cache_stats(self)
1194    }
1195
1196    /// Return the cuTENSOR structural permutation plan entry bound.
1197    /// # Errors
1198    ///
1199    /// Returns [`crate::Error::RuntimeState`] if the cache mutex is poisoned.
1200    pub fn cutensor_permutation_plan_cache_max_entries(&self) -> crate::Result<NonZeroUsize> {
1201        permutation::cutensor_permutation_plan_cache_max_entries(self)
1202    }
1203
1204    /// Configure the cuTENSOR structural permutation plan entry bound.
1205    ///
1206    /// The cache is initialized if it does not already exist so a setting made
1207    /// before the first CUDA structural permutation call is preserved.
1208    /// # Errors
1209    ///
1210    /// Returns [`crate::Error::RuntimeState`] if the cache mutex is poisoned.
1211    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    /// Materialize a strided `Bool` view through the `u8` storage it shares
1456    /// with the native permutation kernel, like [`Self::to_contiguous_view_typed`].
1457    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            // cuTENSOR 2.x rejects zero/negative-stride tensor descriptors. This
1497            // keeps existing CUDA view coverage for a layout the vendor
1498            // permutation path cannot represent; it is not a missing-library
1499            // fallback for cuTENSOR-supported descriptors.
1500            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        // Both operands keep their own strides and offsets: a region inside a
1555        // larger allocation is read and written in place instead of being
1556        // canonicalized into scratch first. Axis fusion collapses the affine
1557        // runs, so a compact sub-block still costs one flat pass.
1558        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        // A fused plan that reduces to one matrix transpose is exactly what the
1570        // tiled transpose kernel implements, and that kernel is coalesced on
1571        // both operands, so take it before the flat kernels. A flat pass over a
1572        // multi-axis permutation reads one contiguous run per source coordinate
1573        // and scatters one element run per destination coordinate, which is why
1574        // the 1 GiB class of copies stays on the generic kernel (issue #1891).
1575        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                // SAFETY: The source is a compact zero-offset CubeCL view on
1586                // this runtime. Allocation identity validation above proves
1587                // source and destination do not alias. The destination view
1588                // has validated reachable offsets and no internal overlap, and
1589                // the launch domain covers each source element and destination
1590                // logical coordinate exactly once.
1591                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            // SAFETY: Both array bindings cover their whole root allocation,
1613            // and `NativeStridedCopyPlan` proved every logical coordinate maps
1614            // inside the source and destination allocation spans, that the
1615            // destination is injective, and (with the allocation identity
1616            // check above) that the two allocations are distinct. The launch
1617            // domain visits each logical coordinate exactly once.
1618            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    /// Launch the tiled transpose kernel for a fused plan that it implements.
1637    ///
1638    /// Returns `Ok(false)` when the layout is outside that kernel's contract or
1639    /// the tiled configuration is disabled, so every other copy keeps its
1640    /// current kernel.
1641    ///
1642    /// # Errors
1643    ///
1644    /// Returns [`crate::Error::Validation`] when the metadata exceeds the
1645    /// kernel's launch limits.
1646    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        // A wide matrix needs a wider tile to stay inside the per-dimension
1664        // launch limit: an 1048576-element axis is 65536 blocks at the default
1665        // 16-wide tile, one past the limit. Widening the tile keeps the tiled
1666        // kernel available instead of falling back to a flat pass, and the
1667        // shared-memory budget bounds how far it can grow.
1668        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            // SAFETY: `NativeStridedCopyPlan` proved every logical coordinate
1706            // maps inside both allocation spans, that the destination is
1707            // injective, and that the two allocations are distinct. The
1708            // orientation check proves the source is row-major and the
1709            // destination column-major over the same matrix, which is the
1710            // kernel's indexing contract; its bounds guards cover edge tiles
1711            // and every unit reaches the shared-memory barrier.
1712            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    /// Build a tiled-transpose plan for a copy whose destination is a
1732    /// row-major-compact view at offset zero.
1733    ///
1734    /// Copying from a compact column-major source into a row-major compact
1735    /// destination of the same logical shape writes exactly the same physical
1736    /// bytes as materializing the transposed source into a compact
1737    /// column-major destination. Selecting that plan keeps the value-exact
1738    /// native path while replacing the uncoalesced access pattern of
1739    /// `contiguous_to_view_kernel` with the existing tiled transpose kernel.
1740    ///
1741    /// Returns `Ok(None)` whenever the layout is outside that narrow shape, so
1742    /// every other copy keeps its current kernel.
1743    ///
1744    /// # Errors
1745    ///
1746    /// Returns [`crate::Error::Validation`] when the transposed plan is rejected
1747    /// by bounds, stride-overflow, or destination-overlap validation; the
1748    /// remaining copy kernels would reject the same layouts.
1749    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        // The transpose kernel writes a compact column-major region starting at
1764        // the bound allocation, so the destination view must be exactly that
1765        // address range: row-major compact at offset zero.
1766        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    /// Copy through the cuTENSOR permutation executor when the layout supports
1815    /// it, otherwise through the exact native copy.
1816    ///
1817    /// Real dtypes use the vendor `alpha = 1` multiply, which is exact and
1818    /// markedly faster than the native kernel for a multi-axis permutation
1819    /// destination. Complex dtypes also reach this function: they are planned
1820    /// through their real view (`[...shape, 2]`, doubled leading strides,
1821    /// unit-stride trailing axis), so the vendor multiply is a *real* multiply
1822    /// and stays value-exact (issue #1891). The exact native tiled transpose is
1823    /// selected only for a transposing destination into a row-major compact
1824    /// view, where it is coalesced on both sides; every other complex layout
1825    /// keeps the vendor real-view plan.
1826    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        // cuTENSOR 2.x descriptors require positive strides on every operand,
1837        // so reversed and broadcast views stay on the native kernel. This is a
1838        // layout the vendor permutation path cannot represent, not a
1839        // missing-library fallback.
1840        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        // Complex operands are exact through cuTENSOR only via the real view,
1846        // whose unit-stride run is the 16-byte real/imaginary pair; that caps
1847        // an exact multi-axis permutation at a third of the achievable
1848        // bandwidth (issue #1891). When the copy is one tiled 2D transpose the
1849        // native kernel is exact as well and coalesced on both sides, so prefer
1850        // it and keep the vendor plan for every other layout.
1851        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                    // SAFETY: `convert_float_to_complex_raw` validated that
2127                    // `input` has `n` elements and `out` has `2 * n` scalar
2128                    // components. The kernel launches exactly `n` logical input
2129                    // positions and guards with `ABSOLUTE_POS < input.len()`.
2130                    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                    // SAFETY: `convert_float_to_complex_raw` validated that
2152                    // `input` has `n` elements and `out` has `2 * n` scalar
2153                    // components. The kernel launches exactly `n` logical input
2154                    // positions and guards with `ABSOLUTE_POS < input.len()`.
2155                    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                    // SAFETY: `convert_float_to_complex_raw` validated that
2177                    // `input` has `n` elements and `out` has `2 * n` scalar
2178                    // components. The kernel launches exactly `n` logical input
2179                    // positions and guards with `ABSOLUTE_POS < input.len()`.
2180                    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                    // SAFETY: `convert_float_to_complex_raw` validated that
2202                    // `input` has `n` elements and `out` has `2 * n` scalar
2203                    // components. The kernel launches exactly `n` logical input
2204                    // positions and guards with `ABSOLUTE_POS < input.len()`.
2205                    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    /// Generic float-to-complex conversion via raw interleaved kernel.
2219    ///
2220    /// The kernel writes `(re, 0, re, 0, ...)` into a raw float buffer that
2221    /// is then reinterpreted as complex.
2222    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        // SAFETY: The checked raw-array helpers prove that `input_arg` covers
2259        // exactly the dense input shape and `output_parts` covers the complete
2260        // real/imaginary scalar representation of the output allocation.
2261        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        // The first reduction reads the caller-owned input directly. Subsequent
2684        // axes consume the fresh keepdims result from the preceding launch.
2685        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, &current, 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(&current, 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            // INVARIANT: `concatenate_output_shape(inputs, axis)?` above checks
3324            // the total axis extent, so every partial offset stays bounded.
3325            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            // SAFETY: `scatter_launch_meta` validates the scatter/update
3533            // shapes and dimension-number mappings. `typed_tensor_binding`
3534            // validates input logical tensor buffers, while
3535            // `typed_tensor_array_arg_as` proves the atomic output view stays
3536            // within its backing allocation. The launch domain is
3537            // `scatter_update_len(meta)`, and the kernel maps each launched
3538            // update through the validated metadata before indexing.
3539            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        // num_complex::Complex<T> is repr(C) as { re: T, im: T }, so the
3636        // complex buffers can be viewed as real scalar parts for atomic add.
3637        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            // SAFETY: `scatter_launch_meta` validates the scatter/update
3644            // shapes and dimension-number mappings. `typed_tensor_binding`
3645            // validates logical tensor buffers, while `typed_tensor_array_arg_as`
3646            // proves complex real/imaginary part arrays stay within their
3647            // backing allocations. The launch domain is
3648            // `scatter_update_len(meta)` and the kernel indexes via the
3649            // validated metadata.
3650            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        // SAFETY: the flag allocation has two `F` elements, and the checked
3985        // reinterpretation above proves the atomic-u32 view fits that buffer.
3986        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        // SAFETY: the input binding was validated before allocation, the
3997        // launch domain is the checked input length, and the scalar flag view
3998        // was bounds-checked above.
3999        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        // SAFETY: one worker reads the atomically selected in-range index and
4013        // copies that single value into the second element of the same flag.
4014        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    // INVARIANT: `num_complex::Complex<T>` is `repr(C)` with interleaved `{ re, im }`
4187    // fields; the checked `2 * n_elements` length and binding validator prove this
4188    // real-component view covers exactly the resident complex allocation.
4189    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
4211/// The typed tensor a tag names, or a typed error if the payload does not carry it.
4212///
4213/// Dispatch on [`Tensor::dtype`] and this accessor are the pair that lets a GPU operation be written
4214/// against the tag rather than against every `Tensor` variant. The error is reachable only if a tag
4215/// and its payload ever disagree, which the tag itself excludes, so it is a typed refusal rather than
4216/// a panic.
4217fn 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    // Dispatch on the pair of tags and recover each typed tensor, which is what `as_typed` exists
4234    // for; the closure keeps the typed error inside the `Option` the caller expects.
4235    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
4350/// Operand accepted by a CUDA `_read` entry point.
4351///
4352/// The traced runtime prepares operands as `TensorRead`, which is either an
4353/// owned tensor or a borrowed view over another tensor's storage. Operations
4354/// without a native view kernel retain the explicit materialization fallback.
4355enum CudaReadInput<'a> {
4356    /// The caller owns the tensor for the duration of the call.
4357    Borrowed(&'a Tensor),
4358    /// A borrowed view materialized into backend storage for the call.
4359    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    /// Run the common same-shape elementwise cases directly into the caller's
4461    /// owned CUDA output. Views and broadcast/scalar cases retain the existing
4462    /// allocating fallback so their layout and ownership contracts are unchanged.
4463    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                    // SAFETY: native dispatch checks matching owned compact shapes/dtypes;
4494                    // validate_read_into_destination checks overlap, and the launch helper
4495                    // validates runtime residency and prepares the exclusive output write.
4496                    |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                    // SAFETY: native dispatch checks matching owned compact shapes/dtypes;
4512                    // validate_read_into_destination checks overlap, and the launch helper
4513                    // validates runtime residency and prepares the exclusive output write.
4514                    |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                    // SAFETY: launch_binary_views validates equal shapes, zero-offset
4595                    // compact layouts and residency, and allocates an independent output.
4596                    |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                    // SAFETY: launch_binary_views_parts validates equal shapes,
4624                    // zero-offset compact layouts and residency, and allocates an
4625                    // independent output; the parts kernel guards its index domain.
4626                    |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                            // SAFETY: launch_unary_view validates the shape, zero-offset
4725                            // compact layout and residency, and allocates a fresh output.
4726                            |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    /// Accept a read operand the runtime prepared, materializing a view only
4771    /// when the operation has no native borrowed implementation.
4772    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
4782/// The typed tensor behind a contiguous-read adapter's tensor, or its refusal.
4783fn 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
4792/// The typed tensor behind a copy read's tensor, or its refusal.
4793fn 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/// Operand of a fused CUDA kernel: an owned tensor or a compact borrowed view.
4825///
4826/// The eager einsum path prepares operands as borrowed views over already
4827/// allocated device storage, so fused entry points accept both forms, the way
4828/// the traced runtime already hands them.
4829// INVARIANT: the view variant carries the provider descriptor inline so a
4830// fused launch never allocates; the enum only lives for one kernel launch.
4831#[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
4860/// Dtype-erased borrowed view of the operands accepted by
4861/// [`CudaBackend::execute_broadcast_multiply`].
4862enum 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
4871/// Accept a view only when the fused kernel can index it directly.
4872///
4873/// The kernel reads each operand in compact column-major order, which is the
4874/// same requirement [`dispatch::typed_view_binding`] enforces; other view
4875/// forms keep the caller's materializing fallback.
4876fn 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    // `TypedTensor::into_parts` intentionally materializes host storage and
5144    // therefore cannot preserve a backend-owned root. Move the owner through
5145    // `TensorValue`/`AllocationGroup` instead so this metadata-only reshape
5146    // keeps the exact CubeCL allocation and performs no implicit download.
5147    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        // INVARIANT: the owner just published above was produced by this same
5166        // typed path, so its dtype tag always matches `T`.
5167        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            // INVARIANT: This boundary check intentionally mirrors CPU's
5194            // validator. CPU and GPU are independent backend leaves, and
5195            // sharing it via tenferro-tensor would require a new public
5196            // validation API.
5197            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;