Skip to main content

tenferro_gpu/cubecl/
gemm.rs

1use std::ffi::c_void;
2use std::num::NonZeroUsize;
3use std::sync::{Arc, Mutex};
4
5use cubecl::prelude::{CubeElement, CubePrimitive};
6use cubecl::stream_id::StreamId;
7use cubecl_cuda::CudaRuntime as CubeclCudaRuntime;
8use num_complex::{Complex32, Complex64};
9use num_traits::{One, Zero};
10
11use super::dispatch::{
12    alloc_output, cube_count_for_len, cube_dim_1d, cubecl_buffer, cubecl_view_buffer,
13    cubecl_view_mut_buffer, dtype_mismatch, ensure_resident_on_runtime,
14    ensure_view_mut_resident_on_runtime, ensure_view_resident_on_runtime, launch_nullary_into,
15    prepared_tensor_access, prepared_view_access, prepared_view_mut_access, CubeclPreparedAccess,
16};
17use super::error::{unsupported_dtype, unsupported_operation, workspace_size_overflow};
18use super::ffi::cutensor::{
19    CudaDataType, CutensorComputeDescriptor, CutensorCudaStream, CutensorHandle, CutensorOperator,
20    CutensorWorksizePreference, OperationDescriptor, Plan, PlanPreference, TensorDescriptor,
21};
22use super::interop::cuda_device_ptr_from_addr;
23use super::plan_cache::LruPlanCache;
24use super::{CudaBackend, CudaRuntime};
25use crate::config::DotGeneralConfig;
26use crate::kernels::structural;
27use crate::{col_major_strides, CubeclBuffer, Error, Tensor, TypedTensor};
28use tenferro_tensor::{
29    CacheStats, ContractionScalar, DType, DotGeneralAccumulation, TensorRead, TensorScalar,
30    TensorView, TensorViewMut, TensorWrite, TypedTensorView, TypedTensorViewMut,
31};
32
33const OP: &str = "dot_general";
34const CUDA_ALLOCATION_ALIGNMENT: u32 = 256;
35const DEFAULT_CUTENSOR_PLAN_CACHE_MAX_ENTRIES: usize = 64;
36type CutensorContractionPlanCache = LruPlanCache<CutensorContractionKey, CachedCutensorContraction>;
37type CutensorPlanCacheState = Arc<Mutex<CutensorContractionCacheState>>;
38
39/// cuTENSOR contraction plan cache plus the shared device scratch.
40///
41/// Every cached plan keeps only its cuTENSOR plan metadata and the scratch
42/// size it needs; one lazily grown workspace per physical stream slot is
43/// shared by all plans on that slot. The enclosing mutex serializes host
44/// enqueues, so all uses of a slot's scratch buffer are ordered on the same
45/// physical CUDA stream.
46struct CutensorContractionCacheState {
47    // INVARIANT: declared before `plans`, so teardown retires the shared
48    // device scratch before the plans and descriptors that referenced it are
49    // destroyed. Individual plan eviction leaves the shared scratch alone.
50    workspaces: Box<[Option<Workspace>]>,
51    plans: CutensorContractionPlanCache,
52}
53
54impl CutensorContractionCacheState {
55    fn new(max_entries: NonZeroUsize, stream_slots: usize) -> Self {
56        Self {
57            plans: CutensorContractionPlanCache::new(max_entries),
58            workspaces: (0..stream_slots).map(|_| None).collect(),
59        }
60    }
61
62    /// Bytes of the shared workspaces currently allocated (all stream slots).
63    fn workspace_bytes(&self) -> u64 {
64        retained_workspace_bytes(&self.workspaces)
65    }
66
67    /// Retained shared buffers and their total capacity.
68    ///
69    /// INVARIANT: a stored workspace always has nonzero capacity, because
70    /// `plan_workspace` answers `Reuse` for a zero request instead of storing a
71    /// zero-size placeholder. `retained_entries` therefore counts exactly the
72    /// slots holding a device buffer.
73    fn workspace_stats(&self) -> CutensorWorkspaceStats {
74        CutensorWorkspaceStats {
75            retained_entries: self.workspaces.iter().flatten().count(),
76            retained_bytes: self.workspace_bytes(),
77        }
78    }
79
80    /// Drop every retained shared buffer. Cached plans and descriptors are
81    /// untouched; each dropped buffer is retired through the event queue.
82    fn release_workspaces(&mut self) {
83        for workspace in self.workspaces.iter_mut() {
84            *workspace = None;
85        }
86    }
87}
88
89fn retained_workspace_bytes(workspaces: &[Option<Workspace>]) -> u64 {
90    // INVARIANT: the slice holds one entry per physical stream slot, a count
91    // fixed at runtime construction, so this fold has a constant configured
92    // ceiling.
93    workspaces.iter().flatten().fold(0_u64, |total, workspace| {
94        total.saturating_add(workspace.size)
95    })
96}
97
98/// Retained shared cuTENSOR contraction scratch, as reported to callers.
99///
100/// This is scratch the backend keeps for reuse, not total device memory: a
101/// workspace in use by a queued contraction, a retiring allocation, the
102/// allocator arena, and vendor-internal allocations are all excluded.
103///
104/// # Examples
105///
106/// ```
107/// use tenferro_gpu::cuda::CutensorWorkspaceStats;
108///
109/// let stats = CutensorWorkspaceStats {
110///     retained_entries: 1,
111///     retained_bytes: 4 << 20,
112/// };
113/// assert_eq!(stats.retained_bytes, 4 << 20);
114/// ```
115#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
116pub struct CutensorWorkspaceStats {
117    /// Physical stream slots currently holding a shared workspace buffer.
118    pub retained_entries: usize,
119    /// Total retained capacity of those buffers, in bytes.
120    pub retained_bytes: u64,
121}
122
123/// What one contraction should do with the shared scratch for its stream slot.
124#[derive(Clone, Copy, Debug, PartialEq, Eq)]
125pub(super) enum WorkspacePlan {
126    /// The slot's current buffer is large enough; nothing to allocate.
127    Reuse,
128    /// Replace the slot's buffer with this capacity and retain it.
129    Retain(u64),
130    /// Execute with a temporary workspace of this exact size and leave the
131    /// slot's current buffer unchanged.
132    Temporary(u64),
133}
134
135/// Choose the shared-scratch action for one contraction.
136///
137/// `requested` is the cuTENSOR workspace estimate, `current_capacity` the
138/// slot's retained capacity, `retained_total` the capacity retained by every
139/// slot, and `limit` the configured backend-wide retention cap.
140///
141/// A contraction is never refused because of the cap: when the requested size
142/// does not fit the remaining headroom, it runs in an exact-size temporary
143/// workspace instead of a retained one. Other slots are never evicted to make
144/// room, and `limit` is not a bound on what a single contraction may use.
145pub(super) fn plan_workspace(
146    requested: u64,
147    current_capacity: u64,
148    retained_total: u64,
149    limit: u64,
150) -> WorkspacePlan {
151    if current_capacity >= requested {
152        return WorkspacePlan::Reuse;
153    }
154    // Headroom after keeping every other slot's buffer. Saturating arithmetic
155    // keeps this well-defined for any `u64` limit, including values near
156    // `u64::MAX`.
157    let headroom = limit.saturating_sub(retained_total.saturating_sub(current_capacity));
158    match shared_workspace_capacity(requested) {
159        Some(rounded) if rounded <= headroom => WorkspacePlan::Retain(rounded),
160        // Rounding overflow, or the rounded capacity does not fit: retain the
161        // exact request when that fits.
162        _ if requested <= headroom => WorkspacePlan::Retain(requested),
163        _ => WorkspacePlan::Temporary(requested),
164    }
165}
166
167trait CutensorScalar: CubeElement + TensorScalar + CubePrimitive + Clone + One + Zero {
168    const DATA_TYPE: CudaDataType;
169    const DTYPE: DType;
170    const IS_COMPLEX: bool;
171
172    fn compute_descriptor(handle: &CutensorHandle) -> CutensorComputeDescriptor;
173
174    /// Unwrap the matching dtype variant from dtype-erased tensors.
175    fn unwrap_tensor(tensor: &Tensor) -> Option<&TypedTensor<Self>>;
176
177    /// Unwrap the matching dtype variant from dtype-erased borrowed views.
178    fn unwrap_view<'a, 'b>(view: &'a TensorView<'b>) -> Option<&'a TypedTensorView<'b, Self>>;
179    fn unwrap_view_mut<'a, 'b>(
180        view: &'a mut TensorViewMut<'b>,
181    ) -> Option<&'a mut TypedTensorViewMut<'b, Self>>;
182    fn unwrap_tensor_mut(tensor: &mut Tensor) -> Option<&mut TypedTensor<Self>>;
183
184    /// Launch the in-place scale kernel for this scalar type.
185    fn launch_scale_in_place(
186        client: &cubecl::prelude::ComputeClient<CubeclCudaRuntime>,
187        count: cubecl::prelude::CubeCount,
188        dim: cubecl::prelude::CubeDim,
189        out: cubecl::prelude::ArrayArg<CubeclCudaRuntime>,
190        factor: cubecl::prelude::ArrayArg<CubeclCudaRuntime>,
191    );
192}
193
194/// Implement the dtype-erased variant accessors for one scalar type.
195macro_rules! cutensor_variant_accessors {
196    ($variant:ident) => {
197        fn unwrap_view<'a, 'b>(view: &'a TensorView<'b>) -> Option<&'a TypedTensorView<'b, Self>> {
198            match view {
199                TensorView::$variant(view) => Some(view),
200                _ => None,
201            }
202        }
203
204        fn unwrap_view_mut<'a, 'b>(
205            view: &'a mut TensorViewMut<'b>,
206        ) -> Option<&'a mut TypedTensorViewMut<'b, Self>> {
207            match view {
208                TensorViewMut::$variant(view) => Some(view),
209                _ => None,
210            }
211        }
212
213        fn unwrap_tensor_mut(tensor: &mut Tensor) -> Option<&mut TypedTensor<Self>> {
214            tensor.as_typed_mut::<Self>()
215        }
216    };
217}
218
219impl CutensorScalar for f32 {
220    cutensor_variant_accessors!(F32);
221
222    const DATA_TYPE: CudaDataType = CudaDataType::R32F;
223    const DTYPE: DType = DType::F32;
224    const IS_COMPLEX: bool = false;
225
226    fn compute_descriptor(handle: &CutensorHandle) -> CutensorComputeDescriptor {
227        handle.compute_desc_32f()
228    }
229    fn unwrap_tensor(tensor: &Tensor) -> Option<&TypedTensor<Self>> {
230        tensor.as_typed::<Self>()
231    }
232
233    fn launch_scale_in_place(
234        client: &cubecl::prelude::ComputeClient<CubeclCudaRuntime>,
235        count: cubecl::prelude::CubeCount,
236        dim: cubecl::prelude::CubeDim,
237        out: cubecl::prelude::ArrayArg<CubeclCudaRuntime>,
238        factor: cubecl::prelude::ArrayArg<CubeclCudaRuntime>,
239    ) {
240        // SAFETY: caller validated residency, lengths, and launch domain.
241        unsafe {
242            structural::scale_in_place_float_kernel::launch_unchecked::<f32, CubeclCudaRuntime>(
243                client, count, dim, out, factor,
244            );
245        }
246    }
247}
248
249impl CutensorScalar for f64 {
250    cutensor_variant_accessors!(F64);
251
252    const DATA_TYPE: CudaDataType = CudaDataType::R64F;
253    const DTYPE: DType = DType::F64;
254    const IS_COMPLEX: bool = false;
255
256    fn compute_descriptor(handle: &CutensorHandle) -> CutensorComputeDescriptor {
257        handle.compute_desc_64f()
258    }
259    fn unwrap_tensor(tensor: &Tensor) -> Option<&TypedTensor<Self>> {
260        tensor.as_typed::<Self>()
261    }
262
263    fn launch_scale_in_place(
264        client: &cubecl::prelude::ComputeClient<CubeclCudaRuntime>,
265        count: cubecl::prelude::CubeCount,
266        dim: cubecl::prelude::CubeDim,
267        out: cubecl::prelude::ArrayArg<CubeclCudaRuntime>,
268        factor: cubecl::prelude::ArrayArg<CubeclCudaRuntime>,
269    ) {
270        // SAFETY: caller validated residency, lengths, and launch domain.
271        unsafe {
272            structural::scale_in_place_float_kernel::launch_unchecked::<f64, CubeclCudaRuntime>(
273                client, count, dim, out, factor,
274            );
275        }
276    }
277}
278
279impl CutensorScalar for Complex32 {
280    cutensor_variant_accessors!(C32);
281
282    const DATA_TYPE: CudaDataType = CudaDataType::C32F;
283    const DTYPE: DType = DType::C32;
284    const IS_COMPLEX: bool = true;
285
286    fn compute_descriptor(handle: &CutensorHandle) -> CutensorComputeDescriptor {
287        handle.compute_desc_32f()
288    }
289    fn unwrap_tensor(tensor: &Tensor) -> Option<&TypedTensor<Self>> {
290        tensor.as_typed::<Self>()
291    }
292
293    fn launch_scale_in_place(
294        client: &cubecl::prelude::ComputeClient<CubeclCudaRuntime>,
295        count: cubecl::prelude::CubeCount,
296        dim: cubecl::prelude::CubeDim,
297        out: cubecl::prelude::ArrayArg<CubeclCudaRuntime>,
298        factor: cubecl::prelude::ArrayArg<CubeclCudaRuntime>,
299    ) {
300        // SAFETY: caller validated residency, lengths, and launch domain.
301        unsafe {
302            structural::scale_in_place_complex_kernel::launch_unchecked::<
303                Complex32,
304                CubeclCudaRuntime,
305            >(client, count, dim, out, factor);
306        }
307    }
308}
309
310impl CutensorScalar for Complex64 {
311    cutensor_variant_accessors!(C64);
312
313    const DATA_TYPE: CudaDataType = CudaDataType::C64F;
314    const DTYPE: DType = DType::C64;
315    const IS_COMPLEX: bool = true;
316
317    fn compute_descriptor(handle: &CutensorHandle) -> CutensorComputeDescriptor {
318        handle.compute_desc_64f()
319    }
320    fn unwrap_tensor(tensor: &Tensor) -> Option<&TypedTensor<Self>> {
321        tensor.as_typed::<Self>()
322    }
323
324    fn launch_scale_in_place(
325        client: &cubecl::prelude::ComputeClient<CubeclCudaRuntime>,
326        count: cubecl::prelude::CubeCount,
327        dim: cubecl::prelude::CubeDim,
328        out: cubecl::prelude::ArrayArg<CubeclCudaRuntime>,
329        factor: cubecl::prelude::ArrayArg<CubeclCudaRuntime>,
330    ) {
331        // SAFETY: caller validated residency, lengths, and launch domain.
332        unsafe {
333            structural::scale_in_place_complex_kernel::launch_unchecked::<
334                Complex64,
335                CubeclCudaRuntime,
336            >(client, count, dim, out, factor);
337        }
338    }
339}
340
341struct DotGeneralLayout {
342    lhs_modes: Vec<i32>,
343    rhs_modes: Vec<i32>,
344    output_modes: Vec<i32>,
345    output_shape: Vec<usize>,
346    lhs_extents: Vec<i64>,
347    rhs_extents: Vec<i64>,
348    output_extents: Vec<i64>,
349    lhs_strides: Vec<i64>,
350    rhs_strides: Vec<i64>,
351    output_strides: Vec<i64>,
352    contracting_elements: usize,
353}
354
355struct Workspace {
356    // CubeCL owns the allocation. `Drop` retires the stream on which cuTENSOR
357    // used it before releasing this handle.
358    _handle: Option<cubecl_runtime::server::Handle>,
359    ptr: *mut c_void,
360    size: u64,
361    runtime: Option<CudaRuntime>,
362    stream: u64,
363}
364
365impl Workspace {
366    fn none() -> Self {
367        Self {
368            _handle: None,
369            ptr: std::ptr::null_mut(),
370            size: 0,
371            runtime: None,
372            stream: 0,
373        }
374    }
375}
376
377impl Drop for Workspace {
378    fn drop(&mut self) {
379        let (Some(runtime), Some(handle)) = (self.runtime.as_ref(), self._handle.take()) else {
380            return;
381        };
382        // Defer the handle release until this workspace's stream reaches the
383        // event recorded now, instead of draining the pipeline here. The event
384        // is the same completion witness the barrier provided.
385        let mut retirements = runtime
386            .workspace_retirements()
387            .lock()
388            .unwrap_or_else(|error| error.into_inner());
389        retirements.retire(runtime.state(), self.stream, handle);
390    }
391}
392
393// SAFETY: `Workspace` owns a CubeCL server handle that keeps the device
394// allocation alive. The raw pointer is only submitted back to cuTENSOR while
395// the cached contraction mutex is held.
396unsafe impl Send for Workspace {}
397
398#[derive(Clone, Debug, PartialEq, Eq, Hash)]
399struct CutensorOperandLayoutKey {
400    extents: Vec<i64>,
401    strides: Vec<i64>,
402    modes: Vec<i32>,
403}
404
405impl CutensorOperandLayoutKey {
406    fn new(extents: &[i64], strides: &[i64], modes: &[i32]) -> Self {
407        Self {
408            extents: extents.to_vec(),
409            strides: strides.to_vec(),
410            modes: modes.to_vec(),
411        }
412    }
413
414    fn retained_bytes(&self) -> usize {
415        std::mem::size_of::<Self>()
416            .saturating_add(self.extents.capacity() * std::mem::size_of::<i64>())
417            .saturating_add(self.strides.capacity() * std::mem::size_of::<i64>())
418            .saturating_add(self.modes.capacity() * std::mem::size_of::<i32>())
419    }
420}
421
422#[derive(Clone, Debug, PartialEq, Eq, Hash)]
423struct CutensorContractionKey {
424    dtype: DType,
425    lhs: CutensorOperandLayoutKey,
426    rhs: CutensorOperandLayoutKey,
427    output: CutensorOperandLayoutKey,
428    lhs_alignment_requirement: u32,
429    rhs_alignment_requirement: u32,
430    output_alignment_requirement: u32,
431    lhs_op: CutensorOperator,
432    rhs_op: CutensorOperator,
433    workspace_preference: CutensorWorksizePreference,
434}
435
436impl CutensorContractionKey {
437    fn from_spec<T: CutensorScalar>(spec: &CutensorContractionSpec<'_>) -> Self {
438        Self {
439            dtype: T::DTYPE,
440            lhs: CutensorOperandLayoutKey::new(
441                &spec.layout.lhs_extents,
442                spec.lhs_strides,
443                &spec.layout.lhs_modes,
444            ),
445            rhs: CutensorOperandLayoutKey::new(
446                &spec.layout.rhs_extents,
447                spec.rhs_strides,
448                &spec.layout.rhs_modes,
449            ),
450            output: CutensorOperandLayoutKey::new(
451                &spec.layout.output_extents,
452                spec.output_strides,
453                &spec.layout.output_modes,
454            ),
455            lhs_alignment_requirement: spec.lhs_alignment_requirement,
456            rhs_alignment_requirement: spec.rhs_alignment_requirement,
457            output_alignment_requirement: spec.output_alignment_requirement,
458            lhs_op: cutensor_conj_op::<T>(spec.lhs_conj),
459            rhs_op: cutensor_conj_op::<T>(spec.rhs_conj),
460            workspace_preference: spec.workspace_preference,
461        }
462    }
463
464    fn retained_bytes(&self) -> usize {
465        std::mem::size_of::<Self>()
466            .saturating_add(self.lhs.retained_bytes())
467            .saturating_add(self.rhs.retained_bytes())
468            .saturating_add(self.output.retained_bytes())
469    }
470}
471
472struct CutensorContractionSpec<'a> {
473    layout: &'a DotGeneralLayout,
474    lhs_strides: &'a [i64],
475    rhs_strides: &'a [i64],
476    output_strides: &'a [i64],
477    lhs_alignment_requirement: u32,
478    rhs_alignment_requirement: u32,
479    output_alignment_requirement: u32,
480    lhs_conj: bool,
481    rhs_conj: bool,
482    workspace_preference: CutensorWorksizePreference,
483}
484
485struct CachedCutensorContraction {
486    // The scratch buffer is shared per stream slot in
487    // `CutensorContractionCacheState`; a cached plan only records how large a
488    // request its plan needs.
489    workspace_size: u64,
490    // Drop the cuTENSOR plan before the descriptor objects it was built from.
491    plan: Plan,
492    _plan_preference: PlanPreference,
493    _operation_descriptor: OperationDescriptor,
494    _output_descriptor: TensorDescriptor,
495    _rhs_descriptor: TensorDescriptor,
496    _lhs_descriptor: TensorDescriptor,
497}
498
499// SAFETY: cached cuTENSOR state is tied to one `CudaBackend` and is only used
500// while holding the enclosing plan-cache mutex. The opaque cuTENSOR handles are
501// created and destroyed through the same loaded cuTENSOR library.
502unsafe impl Send for CachedCutensorContraction {}
503
504impl CachedCutensorContraction {
505    fn new<T>(cutensor: &CutensorHandle, spec: &CutensorContractionSpec<'_>) -> crate::Result<Self>
506    where
507        T: CutensorScalar,
508    {
509        let desc_a = TensorDescriptor::new(
510            cutensor,
511            &spec.layout.lhs_extents,
512            spec.lhs_strides,
513            T::DATA_TYPE,
514            spec.lhs_alignment_requirement,
515            OP,
516        )?;
517        let desc_b = TensorDescriptor::new(
518            cutensor,
519            &spec.layout.rhs_extents,
520            spec.rhs_strides,
521            T::DATA_TYPE,
522            spec.rhs_alignment_requirement,
523            OP,
524        )?;
525        let desc_out = TensorDescriptor::new(
526            cutensor,
527            &spec.layout.output_extents,
528            spec.output_strides,
529            T::DATA_TYPE,
530            spec.output_alignment_requirement,
531            OP,
532        )?;
533        let op_desc = OperationDescriptor::new_contraction_with_ops(
534            cutensor,
535            &desc_a,
536            &spec.layout.lhs_modes,
537            cutensor_conj_op::<T>(spec.lhs_conj),
538            &desc_b,
539            &spec.layout.rhs_modes,
540            cutensor_conj_op::<T>(spec.rhs_conj),
541            &desc_out,
542            &spec.layout.output_modes,
543            &desc_out,
544            &spec.layout.output_modes,
545            T::compute_descriptor(cutensor),
546            OP,
547        )?;
548        let pref = PlanPreference::new_default(cutensor, OP)?;
549        let workspace_size =
550            cutensor.estimate_workspace_size(&op_desc, &pref, spec.workspace_preference, OP)?;
551        let plan = Plan::new(cutensor, &op_desc, &pref, workspace_size, OP)?;
552        Ok(Self {
553            workspace_size,
554            plan,
555            _plan_preference: pref,
556            _operation_descriptor: op_desc,
557            _output_descriptor: desc_out,
558            _rhs_descriptor: desc_b,
559            _lhs_descriptor: desc_a,
560        })
561    }
562
563    fn retained_bytes(&self) -> usize {
564        std::mem::size_of::<Self>()
565    }
566}
567
568/// Hash `spec` into the plan-cache key hash without materializing an owned
569/// key. Field order mirrors [`CutensorContractionKey::from_spec`]; the stored
570/// key is compared with [`key_matches_spec`] on lookup, so a 64-bit collision
571/// degrades to a plan rebuild instead of a wrong plan.
572fn spec_hash<T: CutensorScalar>(spec: &CutensorContractionSpec<'_>) -> u64 {
573    use std::hash::{Hash, Hasher};
574    let mut hasher = std::collections::hash_map::DefaultHasher::new();
575    T::DTYPE.hash(&mut hasher);
576    spec.layout.lhs_extents.hash(&mut hasher);
577    spec.lhs_strides.hash(&mut hasher);
578    spec.layout.lhs_modes.hash(&mut hasher);
579    spec.layout.rhs_extents.hash(&mut hasher);
580    spec.rhs_strides.hash(&mut hasher);
581    spec.layout.rhs_modes.hash(&mut hasher);
582    spec.layout.output_extents.hash(&mut hasher);
583    spec.output_strides.hash(&mut hasher);
584    spec.layout.output_modes.hash(&mut hasher);
585    spec.lhs_alignment_requirement.hash(&mut hasher);
586    spec.rhs_alignment_requirement.hash(&mut hasher);
587    spec.output_alignment_requirement.hash(&mut hasher);
588    cutensor_conj_op::<T>(spec.lhs_conj).hash(&mut hasher);
589    cutensor_conj_op::<T>(spec.rhs_conj).hash(&mut hasher);
590    spec.workspace_preference.hash(&mut hasher);
591    hasher.finish()
592}
593
594/// Verify a stored materialized key against a borrowed spec.
595fn key_matches_spec<T: CutensorScalar>(
596    key: &CutensorContractionKey,
597    spec: &CutensorContractionSpec<'_>,
598) -> bool {
599    key.dtype == T::DTYPE
600        && key.lhs.extents == spec.layout.lhs_extents
601        && key.lhs.strides == spec.lhs_strides
602        && key.lhs.modes == spec.layout.lhs_modes
603        && key.rhs.extents == spec.layout.rhs_extents
604        && key.rhs.strides == spec.rhs_strides
605        && key.rhs.modes == spec.layout.rhs_modes
606        && key.output.extents == spec.layout.output_extents
607        && key.output.strides == spec.output_strides
608        && key.output.modes == spec.layout.output_modes
609        && key.lhs_alignment_requirement == spec.lhs_alignment_requirement
610        && key.rhs_alignment_requirement == spec.rhs_alignment_requirement
611        && key.output_alignment_requirement == spec.output_alignment_requirement
612        && key.lhs_op == cutensor_conj_op::<T>(spec.lhs_conj)
613        && key.rhs_op == cutensor_conj_op::<T>(spec.rhs_conj)
614        && key.workspace_preference == spec.workspace_preference
615}
616
617/// The typed operands behind a same-dtype pair, or the refusal a mismatched pair reports.
618fn gemm_pair_operands<'a, T: TensorScalar>(
619    op: &'static str,
620    lhs: &'a Tensor,
621    rhs: &'a Tensor,
622) -> crate::Result<(&'a TypedTensor<T>, &'a TypedTensor<T>)> {
623    let lhs_t = lhs
624        .as_typed::<T>()
625        .ok_or_else(|| dtype_mismatch(op, lhs, rhs))?;
626    let rhs_t = rhs
627        .as_typed::<T>()
628        .ok_or_else(|| dtype_mismatch(op, lhs, rhs))?;
629    Ok((lhs_t, rhs_t))
630}
631
632pub(super) fn dot_general_with_conj(
633    backend: &CudaBackend,
634    lhs: &Tensor,
635    rhs: &Tensor,
636    config: &DotGeneralConfig,
637    lhs_conj: bool,
638    rhs_conj: bool,
639) -> crate::Result<Tensor> {
640    match (lhs.dtype(), rhs.dtype()) {
641        (DType::F32, DType::F32) => {
642            let (lhs, rhs) = gemm_pair_operands::<f32>(OP, lhs, rhs)?;
643            dot_general_typed_with_conj(backend, lhs, rhs, config, lhs_conj, rhs_conj)
644                .map(Tensor::from_typed::<f32>)
645        }
646        (DType::F64, DType::F64) => {
647            let (lhs, rhs) = gemm_pair_operands::<f64>(OP, lhs, rhs)?;
648            dot_general_typed_with_conj(backend, lhs, rhs, config, lhs_conj, rhs_conj)
649                .map(Tensor::from_typed::<f64>)
650        }
651        (DType::C32, DType::C32) => {
652            let (lhs, rhs) = gemm_pair_operands::<Complex32>(OP, lhs, rhs)?;
653            dot_general_typed_with_conj(backend, lhs, rhs, config, lhs_conj, rhs_conj)
654                .map(Tensor::from_typed::<Complex32>)
655        }
656        (DType::C64, DType::C64) => {
657            let (lhs, rhs) = gemm_pair_operands::<Complex64>(OP, lhs, rhs)?;
658            dot_general_typed_with_conj(backend, lhs, rhs, config, lhs_conj, rhs_conj)
659                .map(Tensor::from_typed::<Complex64>)
660        }
661        _ => Err(dtype_mismatch(OP, lhs, rhs)),
662    }
663}
664
665/// Local extraction of typed accumulation coefficients; dtype mismatches
666/// between the coefficient and the operand dtype are explicit errors.
667trait FromContractionScalar: Sized {
668    fn from_contraction_scalar(value: ContractionScalar) -> crate::Result<Self>;
669}
670
671macro_rules! impl_from_contraction_scalar {
672    ($ty:ty, $variant:ident) => {
673        impl FromContractionScalar for $ty {
674            fn from_contraction_scalar(value: ContractionScalar) -> crate::Result<Self> {
675                match value {
676                    ContractionScalar::$variant(value) => Ok(value),
677                    other => Err(Error::dtype_mismatch(
678                        OP,
679                        <$ty as tenferro_tensor::TensorScalar>::dtype(),
680                        other.dtype(),
681                    )),
682                }
683            }
684        }
685    };
686}
687
688impl_from_contraction_scalar!(f32, F32);
689impl_from_contraction_scalar!(f64, F64);
690impl_from_contraction_scalar!(Complex32, C32);
691impl_from_contraction_scalar!(Complex64, C64);
692
693/// Read-slot operand for the accumulate path: an owned compact tensor or a
694/// borrowed strided view over a device buffer.
695enum ReadOperand<'a, 'b, T> {
696    Owned(&'a TypedTensor<T>),
697    View(&'a TypedTensorView<'b, T>),
698}
699
700impl<T: 'static> ReadOperand<'_, '_, T> {
701    fn shape(&self) -> &[usize] {
702        match self {
703            Self::Owned(tensor) => tensor.shape(),
704            Self::View(view) => view.shape(),
705        }
706    }
707
708    fn handle(&self) -> crate::Result<&cubecl_runtime::server::Handle> {
709        match self {
710            Self::Owned(tensor) => Ok(cubecl_buffer(tensor, OP)?.handle()),
711            Self::View(view) => Ok(cubecl_view_buffer(view, OP)?.handle()),
712        }
713    }
714}
715
716fn read_operand_alignment_requirement<T: CutensorScalar>(operand: &ReadOperand<'_, '_, T>) -> u32 {
717    match operand {
718        ReadOperand::Owned(_) => CUDA_ALLOCATION_ALIGNMENT,
719        ReadOperand::View(_) => view_descriptor_alignment_requirement::<T>(),
720    }
721}
722
723/// Write-slot operand for the accumulate path.
724enum WriteOperand<'a, 'b, T> {
725    Owned(&'a mut TypedTensor<T>),
726    View(&'a mut TypedTensorViewMut<'b, T>),
727}
728
729impl<T: 'static> WriteOperand<'_, '_, T> {
730    fn shape(&self) -> &[usize] {
731        match self {
732            Self::Owned(tensor) => tensor.shape(),
733            Self::View(view) => view.shape(),
734        }
735    }
736
737    fn n_elements(&self) -> usize {
738        match self {
739            Self::Owned(tensor) => tensor.n_elements(),
740            Self::View(view) => view.n_elements(),
741        }
742    }
743
744    fn handle(&self) -> crate::Result<&cubecl_runtime::server::Handle> {
745        match self {
746            Self::Owned(tensor) => Ok(cubecl_buffer(tensor, OP)?.handle()),
747            Self::View(view) => Ok(cubecl_view_mut_buffer(view, OP)?.handle()),
748        }
749    }
750}
751
752fn cross_stream_handles<'a>(
753    rt: &CudaRuntime,
754    handles: impl IntoIterator<Item = &'a cubecl_runtime::server::Handle>,
755) -> Vec<cubecl_runtime::server::Handle> {
756    handles
757        .into_iter()
758        .filter(|handle| !rt.is_current_stream_slot(handle))
759        .cloned()
760        .collect()
761}
762
763fn write_operand_alignment_requirement<T: CutensorScalar>(
764    operand: &WriteOperand<'_, '_, T>,
765) -> u32 {
766    match operand {
767        WriteOperand::Owned(_) => CUDA_ALLOCATION_ALIGNMENT,
768        WriteOperand::View(_) => view_descriptor_alignment_requirement::<T>(),
769    }
770}
771
772fn view_descriptor_alignment_requirement<T: CutensorScalar>() -> u32 {
773    u32::try_from(std::mem::size_of::<T>()).unwrap_or(CUDA_ALLOCATION_ALIGNMENT)
774}
775
776/// Device pointer plus cuTENSOR descriptor metadata for one operand.
777/// Compact owned operands borrow the layout's precomputed strides; strided
778/// views own their converted strides.
779struct ResolvedOperand<'a> {
780    ptr: *mut c_void,
781    strides: std::borrow::Cow<'a, [i64]>,
782    alignment: u32,
783}
784
785fn read_operand<'a, 'b, T: CutensorScalar>(
786    read: &'a TensorRead<'b>,
787) -> Option<ReadOperand<'a, 'b, T>> {
788    match read {
789        TensorRead::Tensor(tensor) => T::unwrap_tensor(tensor).map(ReadOperand::Owned),
790        TensorRead::View(view) => T::unwrap_view(view).map(ReadOperand::View),
791    }
792}
793
794fn write_operand<'a, 'b, T: CutensorScalar>(
795    write: &'a mut TensorWrite<'b>,
796) -> Option<WriteOperand<'a, 'b, T>> {
797    match write {
798        TensorWrite::Tensor(tensor) => T::unwrap_tensor_mut(tensor).map(WriteOperand::Owned),
799        TensorWrite::View(view) => T::unwrap_view_mut(view).map(WriteOperand::View),
800    }
801}
802
803/// CUDA-native accumulate-form contraction:
804/// `out = alpha * op(lhs) * op(rhs) + beta * out` executed by a single
805/// cuTENSOR contraction with `C = D = out` (no temporary result tensor).
806///
807/// Stage-2 scope (tensor4all/tenferro-rs#1287): every slot accepts either a
808/// compact GPU-resident owned tensor or a borrowed strided view over a
809/// GPU-resident device buffer. Host-backed views, negative view strides, and
810/// out-of-bounds regions are explicit errors — no hidden host transfer and no
811/// silent fallback.
812/// Allocate the dot-general output and contract straight from the reads.
813///
814/// The `TensorDot` default implementations run any strided read through
815/// `to_contiguous_read` before contracting, which turns every view operand into
816/// a materialized copy. The device descriptors already carry the operand
817/// strides (see [`dot_general_read_into_accum`]), so view operands are consumed
818/// in place here instead.
819pub(super) fn dot_general_read_allocating(
820    backend: &mut CudaBackend,
821    lhs: TensorRead<'_>,
822    rhs: TensorRead<'_>,
823    config: &DotGeneralConfig,
824    lhs_conj: bool,
825    rhs_conj: bool,
826) -> crate::Result<Tensor> {
827    if let (Some(lhs_owned), Some(rhs_owned)) = (lhs.as_tensor(), rhs.as_tensor()) {
828        return dot_general_with_conj(backend, lhs_owned, rhs_owned, config, lhs_conj, rhs_conj);
829    }
830    let dtype = lhs.dtype();
831    let shape =
832        tenferro_tensor::backend::dot_general_output_shape(lhs.shape(), rhs.shape(), config, OP)?;
833    let mut out = match dtype {
834        DType::F32 => Tensor::from_typed::<f32>(alloc_output::<f32>(backend.runtime(), &shape)?),
835        DType::F64 => Tensor::from_typed::<f64>(alloc_output::<f64>(backend.runtime(), &shape)?),
836        DType::C32 => {
837            Tensor::from_typed::<Complex32>(alloc_output::<Complex32>(backend.runtime(), &shape)?)
838        }
839        DType::C64 => {
840            Tensor::from_typed::<Complex64>(alloc_output::<Complex64>(backend.runtime(), &shape)?)
841        }
842        dtype => return Err(unsupported_dtype(OP, dtype)),
843    };
844    let accumulation = DotGeneralAccumulation {
845        lhs_conj,
846        rhs_conj,
847        ..DotGeneralAccumulation::overwrite(dtype)?
848    };
849    {
850        let mut out_write = TensorWrite::from_tensor(&mut out);
851        dot_general_read_into_accum(backend, &lhs, &rhs, config, accumulation, &mut out_write)?;
852    }
853    Ok(out)
854}
855
856pub(super) fn dot_general_read_into_accum(
857    backend: &CudaBackend,
858    lhs: &TensorRead<'_>,
859    rhs: &TensorRead<'_>,
860    config: &DotGeneralConfig,
861    accumulation: DotGeneralAccumulation,
862    out: &mut TensorWrite<'_>,
863) -> crate::Result<()> {
864    match lhs.dtype() {
865        DType::F32 => accum_erased::<f32>(backend, lhs, rhs, config, accumulation, out),
866        DType::F64 => accum_erased::<f64>(backend, lhs, rhs, config, accumulation, out),
867        DType::C32 => accum_erased::<Complex32>(backend, lhs, rhs, config, accumulation, out),
868        DType::C64 => accum_erased::<Complex64>(backend, lhs, rhs, config, accumulation, out),
869        dtype => Err(unsupported_dtype(OP, dtype)),
870    }
871}
872
873fn accum_erased<T>(
874    backend: &CudaBackend,
875    lhs: &TensorRead<'_>,
876    rhs: &TensorRead<'_>,
877    config: &DotGeneralConfig,
878    accumulation: DotGeneralAccumulation,
879    out: &mut TensorWrite<'_>,
880) -> crate::Result<()>
881where
882    T: CutensorScalar + FromContractionScalar + PartialEq + tenferro_tensor::TensorScalar,
883{
884    let (lhs_dtype, rhs_dtype, out_dtype) = (lhs.dtype(), rhs.dtype(), out.dtype());
885    let (Some(lhs), Some(rhs), Some(out)) = (
886        read_operand::<T>(lhs),
887        read_operand::<T>(rhs),
888        write_operand::<T>(out),
889    ) else {
890        let (expected, actual) = if lhs_dtype != rhs_dtype {
891            (lhs_dtype, rhs_dtype)
892        } else {
893            (lhs_dtype, out_dtype)
894        };
895        return Err(Error::dtype_mismatch(OP, expected, actual));
896    };
897    dot_general_typed_into_accum(
898        backend,
899        lhs,
900        rhs,
901        config,
902        accumulation.lhs_conj,
903        accumulation.rhs_conj,
904        T::from_contraction_scalar(accumulation.alpha)?,
905        T::from_contraction_scalar(accumulation.beta)?,
906        out,
907    )
908}
909
910#[allow(clippy::too_many_arguments)]
911fn dot_general_typed_into_accum<T>(
912    backend: &CudaBackend,
913    lhs: ReadOperand<'_, '_, T>,
914    rhs: ReadOperand<'_, '_, T>,
915    config: &DotGeneralConfig,
916    lhs_conj: bool,
917    rhs_conj: bool,
918    alpha: T,
919    beta: T,
920    mut out: WriteOperand<'_, '_, T>,
921) -> crate::Result<()>
922where
923    T: CutensorScalar + PartialEq + tenferro_tensor::TensorScalar,
924{
925    backend.runtime().set_current_cuda_context(OP)?;
926    validate_dot_general(lhs.shape(), rhs.shape(), config)?;
927    let layout = build_layout(lhs.shape(), rhs.shape(), config)?;
928    if out.shape() != layout.output_shape.as_slice() {
929        return Err(Error::shape_mismatch(
930            OP,
931            out.shape().to_vec(),
932            layout.output_shape.clone(),
933        ));
934    }
935    let cross_stream_handles = cross_stream_handles(
936        backend.runtime(),
937        [lhs.handle()?, rhs.handle()?, out.handle()?],
938    );
939    // Residency, buffer-family, stride-sign, and bounds validation for all
940    // three slots happens here, before any degenerate-case early return.
941    let lhs_res = resolve_read_operand(backend.runtime(), &lhs, &layout.lhs_strides)?;
942    let rhs_res = resolve_read_operand(backend.runtime(), &rhs, &layout.rhs_strides)?;
943    let out_res = resolve_write_operand(backend.runtime(), &mut out, &layout.output_strides)?;
944    if out.n_elements() == 0 {
945        return Ok(());
946    }
947    if layout.contracting_elements == 0 {
948        // The contraction sum is empty: out = beta * out.
949        return match out {
950            WriteOperand::Owned(tensor) => scale_in_place(backend.runtime(), tensor, beta),
951            WriteOperand::View(_) => {
952                if beta == T::one() {
953                    Ok(())
954                } else {
955                    // No strided in-place scale kernel exists yet; an explicit
956                    // error is required instead of a silent wrong result.
957                    Err(unsupported_operation(
958                        OP,
959                        "zero-sized contraction with beta != 1 is not supported for borrowed view outputs",
960                    ))
961                }
962            }
963        };
964    }
965
966    let stream = raw_stream(backend.runtime())?;
967    let spec = CutensorContractionSpec {
968        layout: &layout,
969        lhs_strides: &lhs_res.strides,
970        rhs_strides: &rhs_res.strides,
971        output_strides: &out_res.strides,
972        lhs_alignment_requirement: read_operand_alignment_requirement(&lhs),
973        rhs_alignment_requirement: read_operand_alignment_requirement(&rhs),
974        output_alignment_requirement: write_operand_alignment_requirement(&out),
975        lhs_conj,
976        rhs_conj,
977        workspace_preference: CutensorWorksizePreference::Default,
978    };
979    validate_descriptor_alignment(lhs_res.alignment, spec.lhs_alignment_requirement, "lhs")?;
980    validate_descriptor_alignment(rhs_res.alignment, spec.rhs_alignment_requirement, "rhs")?;
981    validate_descriptor_alignment(out_res.alignment, spec.output_alignment_requirement, "out")?;
982    // C = D = out: cuTENSOR reads the destination as the accumulator (skipped
983    // by cuTENSOR itself when beta == 0) and writes the result in place.
984    // Overlap between the out region and the lhs/rhs regions is the caller's
985    // responsibility, as with BLAS-style in-place update APIs.
986    cached_cutensor_contraction::<T, _>(
987        backend,
988        &spec,
989        cross_stream_handles,
990        |cutensor, plan, workspace| unsafe {
991            cutensor.contract(
992                plan,
993                &alpha as *const T as *const c_void,
994                lhs_res.ptr as *const c_void,
995                rhs_res.ptr as *const c_void,
996                &beta as *const T as *const c_void,
997                out_res.ptr as *const c_void,
998                out_res.ptr,
999                workspace.ptr,
1000                workspace.size,
1001                stream,
1002                OP,
1003            )
1004        },
1005    )
1006}
1007
1008fn resolve_read_operand<'a, T>(
1009    rt: &CudaRuntime,
1010    operand: &ReadOperand<'_, '_, T>,
1011    compact_strides: &'a [i64],
1012) -> crate::Result<ResolvedOperand<'a>>
1013where
1014    T: CutensorScalar + 'static,
1015{
1016    match operand {
1017        ReadOperand::Owned(tensor) => Ok(ResolvedOperand {
1018            ptr: typed_device_ptr(rt, tensor, OP)?,
1019            strides: std::borrow::Cow::Borrowed(compact_strides),
1020            alignment: CUDA_ALLOCATION_ALIGNMENT,
1021        }),
1022        ReadOperand::View(view) => {
1023            ensure_view_resident_on_runtime(rt, view, OP)?;
1024            let prepared = prepared_view_access(view, OP)?;
1025            // A read shares the root buffer's memoized address, like an
1026            // owned operand, instead of a blocking round trip per view (#1925).
1027            let base = memoized_device_addr(rt, cubecl_view_buffer(view, OP)?, prepared, OP)?;
1028            resolve_prepared_device_region::<T>(base, view.strides(), view.offset())
1029        }
1030    }
1031}
1032
1033fn resolve_write_operand<'a, T>(
1034    rt: &CudaRuntime,
1035    operand: &mut WriteOperand<'_, '_, T>,
1036    compact_strides: &'a [i64],
1037) -> crate::Result<ResolvedOperand<'a>>
1038where
1039    T: CutensorScalar + 'static,
1040{
1041    match operand {
1042        WriteOperand::Owned(tensor) => Ok(ResolvedOperand {
1043            ptr: write_device_ptr(rt, tensor, OP)?,
1044            strides: std::borrow::Cow::Borrowed(compact_strides),
1045            alignment: CUDA_ALLOCATION_ALIGNMENT,
1046        }),
1047        WriteOperand::View(view) => {
1048            ensure_view_mut_resident_on_runtime(rt, view, OP)?;
1049            let prepared = prepared_view_mut_access(view, OP)?;
1050            // A borrowed destination keeps the blocking round trip, which
1051            // orders the vendor write after queued CubeCL work on the stream.
1052            let base = rt
1053                .client()
1054                .get_resource(prepared.into_handle())
1055                .map_err(|err| Error::backend_source(OP, err))?
1056                .resource()
1057                .ptr;
1058            resolve_prepared_device_region::<T>(base, view.strides(), view.offset())
1059        }
1060    }
1061}
1062
1063/// Resolve a strided view region over a device buffer into an effective
1064/// cuTENSOR operand: `ptr = base + offset * size_of::<T>()`, the view's own
1065/// element strides, and the alignment actually guaranteed by the effective
1066/// byte address.
1067fn resolve_prepared_device_region<T: CutensorScalar + 'static>(
1068    base_addr: u64,
1069    strides: &[isize],
1070    offset: isize,
1071) -> crate::Result<ResolvedOperand<'static>> {
1072    let mut strides_i64 = Vec::with_capacity(strides.len());
1073    for &stride in strides {
1074        if stride < 0 {
1075            return Err(Error::invalid_argument(
1076                OP,
1077                "layout",
1078                format!(
1079                    "cuTENSOR dot-general accumulation requires nonnegative view strides, got {strides:?}; canonicalize the view on device first"
1080                ),
1081            ));
1082        }
1083        strides_i64.push(stride as i64);
1084    }
1085    let offset = usize::try_from(offset)
1086        .map_err(|_| Error::invalid_argument(OP, "layout", "view offset must be nonnegative"))?;
1087    let offset_bytes = offset
1088        .checked_mul(std::mem::size_of::<T>())
1089        .ok_or_else(|| Error::invalid_argument(OP, "layout", "view byte offset overflows"))?;
1090    let addr = base_addr
1091        .checked_add(offset_bytes as u64)
1092        .ok_or_else(|| Error::invalid_argument(OP, "layout", "view device address overflows"))?;
1093    // INVARIANT: CubeCL root allocations are at least 256-byte aligned, and
1094    // the checked element offset preserves alignment to the scalar size. A
1095    // strided view therefore guarantees the scalar-sized cuTENSOR requirement,
1096    // even when the view starts inside the root allocation.
1097    Ok(ResolvedOperand {
1098        ptr: cuda_device_ptr_from_addr(addr, OP)?,
1099        strides: std::borrow::Cow::Owned(strides_i64),
1100        alignment: view_descriptor_alignment_requirement::<T>(),
1101    })
1102}
1103
1104/// Device-side `out *= beta` for the degenerate zero-contraction case. The
1105/// factor is materialized as an explicit one-element device constant; user
1106/// operand tensors are never transferred.
1107fn scale_in_place<T>(rt: &CudaRuntime, out: &mut TypedTensor<T>, beta: T) -> crate::Result<()>
1108where
1109    T: CutensorScalar + PartialEq + tenferro_tensor::TensorScalar,
1110{
1111    if beta == T::one() {
1112        return Ok(());
1113    }
1114    if beta == T::zero() {
1115        return launch_nullary_into(
1116            rt,
1117            out,
1118            OP,
1119            cube_count_for_len(out.n_elements())?,
1120            cube_dim_1d(),
1121            |client, count, dim, out| unsafe {
1122                structural::fill_zero_kernel::launch_unchecked::<T, CubeclCudaRuntime>(
1123                    client, count, dim, out,
1124                );
1125            },
1126        );
1127    }
1128    super::interop::scale_typed_tensor_for_op(rt, out, beta, OP, T::launch_scale_in_place)
1129}
1130
1131fn dot_general_typed_with_conj<T>(
1132    backend: &CudaBackend,
1133    lhs: &TypedTensor<T>,
1134    rhs: &TypedTensor<T>,
1135    config: &DotGeneralConfig,
1136    lhs_conj: bool,
1137    rhs_conj: bool,
1138) -> crate::Result<TypedTensor<T>>
1139where
1140    T: CutensorScalar,
1141{
1142    backend.runtime().set_current_cuda_context(OP)?;
1143    validate_dot_general(lhs.shape(), rhs.shape(), config)?;
1144    let layout = build_layout(lhs.shape(), rhs.shape(), config)?;
1145    let output = alloc_output::<T>(backend.runtime(), &layout.output_shape)?;
1146    if output.n_elements() == 0 {
1147        return Ok(output);
1148    }
1149    if layout.contracting_elements == 0 {
1150        // The contraction sum is empty: fill the already-allocated output with
1151        // zeros instead of allocating a second output tensor.
1152        launch_nullary_into(
1153            backend.runtime(),
1154            &output,
1155            OP,
1156            cube_count_for_len(output.n_elements())?,
1157            cube_dim_1d(),
1158            |client, count, dim, out| unsafe {
1159                structural::fill_zero_kernel::launch_unchecked::<T, CubeclCudaRuntime>(
1160                    client, count, dim, out,
1161                );
1162            },
1163        )?;
1164        return Ok(output);
1165    }
1166
1167    let lhs_ptr = typed_device_ptr(backend.runtime(), lhs, OP)?;
1168    let rhs_ptr = typed_device_ptr(backend.runtime(), rhs, OP)?;
1169    let output_ptr = typed_device_ptr(backend.runtime(), &output, OP)?;
1170
1171    let alpha = T::one();
1172    let beta = T::zero();
1173    let stream = raw_stream(backend.runtime())?;
1174    let spec = CutensorContractionSpec {
1175        layout: &layout,
1176        lhs_strides: &layout.lhs_strides,
1177        rhs_strides: &layout.rhs_strides,
1178        output_strides: &layout.output_strides,
1179        lhs_alignment_requirement: CUDA_ALLOCATION_ALIGNMENT,
1180        rhs_alignment_requirement: CUDA_ALLOCATION_ALIGNMENT,
1181        output_alignment_requirement: CUDA_ALLOCATION_ALIGNMENT,
1182        lhs_conj,
1183        rhs_conj,
1184        workspace_preference: CutensorWorksizePreference::Default,
1185    };
1186    // C = D = output: cuTENSOR never reads the accumulator slot when
1187    // beta == 0, so the freshly allocated output serves as both C and D and
1188    // no separate accumulator tensor is needed.
1189    let cross_stream_handles = cross_stream_handles(
1190        backend.runtime(),
1191        [
1192            cubecl_buffer(lhs, OP)?.handle(),
1193            cubecl_buffer(rhs, OP)?.handle(),
1194            cubecl_buffer(&output, OP)?.handle(),
1195        ],
1196    );
1197    cached_cutensor_contraction::<T, _>(
1198        backend,
1199        &spec,
1200        cross_stream_handles,
1201        |cutensor, plan, workspace| unsafe {
1202            cutensor.contract(
1203                plan,
1204                &alpha as *const T as *const c_void,
1205                lhs_ptr as *const c_void,
1206                rhs_ptr as *const c_void,
1207                &beta as *const T as *const c_void,
1208                output_ptr as *const c_void,
1209                output_ptr,
1210                workspace.ptr,
1211                workspace.size,
1212                stream,
1213                OP,
1214            )
1215        },
1216    )?;
1217
1218    Ok(output)
1219}
1220
1221fn cutensor_conj_op<T: CutensorScalar>(conj: bool) -> CutensorOperator {
1222    if conj && T::IS_COMPLEX {
1223        CutensorOperator::Conj
1224    } else {
1225        CutensorOperator::Identity
1226    }
1227}
1228
1229fn default_cutensor_plan_cache_max_entries() -> NonZeroUsize {
1230    NonZeroUsize::new(DEFAULT_CUTENSOR_PLAN_CACHE_MAX_ENTRIES).unwrap_or(NonZeroUsize::MIN)
1231}
1232
1233fn new_cutensor_plan_cache_state(
1234    max_entries: NonZeroUsize,
1235    stream_slots: usize,
1236) -> CutensorPlanCacheState {
1237    Arc::new(Mutex::new(CutensorContractionCacheState::new(
1238        max_entries,
1239        stream_slots,
1240    )))
1241}
1242
1243fn get_or_init_cutensor_plan_cache(backend: &CudaBackend) -> crate::Result<CutensorPlanCacheState> {
1244    let guard = backend
1245        .cuda_extension_cache()
1246        .get_or_try_init::<CutensorPlanCacheState>(|| {
1247            Ok(new_cutensor_plan_cache_state(
1248                default_cutensor_plan_cache_max_entries(),
1249                backend.runtime().stream_slot_count(),
1250            ))
1251        })?;
1252    Ok(Arc::clone(&guard))
1253}
1254
1255fn lock_cutensor_plan_cache(
1256    cache: &CutensorPlanCacheState,
1257) -> crate::Result<std::sync::MutexGuard<'_, CutensorContractionCacheState>> {
1258    cache
1259        .lock()
1260        .map_err(|_| Error::runtime_state("cutensor_plan_cache", "plan cache lock poisoned"))
1261}
1262
1263pub(super) fn cutensor_plan_cache_stats(backend: &CudaBackend) -> crate::Result<CacheStats> {
1264    let Some(plan_cache) = backend
1265        .cuda_extension_cache()
1266        .get_cloned::<CutensorPlanCacheState>()?
1267    else {
1268        return Ok(CacheStats::empty());
1269    };
1270    let plan_cache = lock_cutensor_plan_cache(&plan_cache)?;
1271    Ok(plan_cache.plans.stats())
1272}
1273
1274/// Retained shared scratch and its occupancy, without creating cache state.
1275pub(super) fn cutensor_workspace_stats(
1276    backend: &CudaBackend,
1277) -> crate::Result<CutensorWorkspaceStats> {
1278    let Some(plan_cache) = backend
1279        .cuda_extension_cache()
1280        .get_cloned::<CutensorPlanCacheState>()?
1281    else {
1282        return Ok(CutensorWorkspaceStats::default());
1283    };
1284    let plan_cache = lock_cutensor_plan_cache(&plan_cache)?;
1285    Ok(plan_cache.workspace_stats())
1286}
1287
1288/// Apply a new retention cap, releasing retained buffers that no longer fit.
1289///
1290/// Cached plans are never evicted by a cap change. Dropping a retained buffer
1291/// retires it through the event queue, so this does not force a stream barrier
1292/// and does not reclaim device memory synchronously.
1293pub(super) fn set_cutensor_workspace_max_retained_bytes(
1294    backend: &CudaBackend,
1295    limit: u64,
1296) -> crate::Result<()> {
1297    let Some(plan_cache) = backend
1298        .cuda_extension_cache()
1299        .get_cloned::<CutensorPlanCacheState>()?
1300    else {
1301        return Ok(());
1302    };
1303    let mut plan_cache = lock_cutensor_plan_cache(&plan_cache)?;
1304    if plan_cache.workspace_bytes() > limit {
1305        plan_cache.release_workspaces();
1306    }
1307    Ok(())
1308}
1309
1310/// Deferred workspace retirement counters for tests and diagnostics.
1311pub(crate) fn cutensor_workspace_retirement_stats(
1312    backend: &CudaBackend,
1313) -> crate::Result<super::workspace_retirement::WorkspaceRetirementStats> {
1314    Ok(backend
1315        .runtime()
1316        .workspace_retirements()
1317        .lock()
1318        .map_err(|_| {
1319            crate::Error::runtime_state(
1320                "cutensor_workspace_retirement",
1321                "retirement queue lock poisoned",
1322            )
1323        })?
1324        .stats())
1325}
1326
1327pub(super) fn cutensor_plan_cache_max_entries(
1328    backend: &CudaBackend,
1329) -> crate::Result<NonZeroUsize> {
1330    let Some(plan_cache) = backend
1331        .cuda_extension_cache()
1332        .get_cloned::<CutensorPlanCacheState>()?
1333    else {
1334        return Ok(default_cutensor_plan_cache_max_entries());
1335    };
1336    let plan_cache = lock_cutensor_plan_cache(&plan_cache)?;
1337    Ok(plan_cache.plans.max_entries())
1338}
1339
1340pub(super) fn set_cutensor_plan_cache_max_entries(
1341    backend: &CudaBackend,
1342    max_entries: NonZeroUsize,
1343) -> crate::Result<()> {
1344    let plan_cache = get_or_init_cutensor_plan_cache(backend)?;
1345    let mut plan_cache = lock_cutensor_plan_cache(&plan_cache)?;
1346    plan_cache.plans.set_max_entries(max_entries);
1347    let retained_bytes = plan_cache.plans.retained_bytes();
1348    backend
1349        .cuda_extension_cache()
1350        .update_retained_bytes::<CutensorPlanCacheState>(retained_bytes)
1351}
1352
1353fn cached_cutensor_contraction<T, R>(
1354    backend: &CudaBackend,
1355    spec: &CutensorContractionSpec<'_>,
1356    cross_stream_handles: Vec<cubecl_runtime::server::Handle>,
1357    execute: impl FnOnce(&CutensorHandle, &Plan, &Workspace) -> crate::Result<R>,
1358) -> crate::Result<R>
1359where
1360    T: CutensorScalar,
1361{
1362    let cutensor = backend.cutensor_handle()?;
1363    let hash = spec_hash::<T>(spec);
1364    let plan_cache = get_or_init_cutensor_plan_cache(backend)?;
1365    let mut plan_cache = lock_cutensor_plan_cache(&plan_cache)?;
1366    let entries_changed = plan_cache.plans.ensure(
1367        hash,
1368        |key| key_matches_spec::<T>(key, spec),
1369        || {
1370            let cached = CachedCutensorContraction::new::<T>(cutensor, spec)?;
1371            let key = CutensorContractionKey::from_spec::<T>(spec);
1372            let retained_bytes = key.retained_bytes().saturating_add(cached.retained_bytes());
1373            Ok((key, cached, retained_bytes))
1374        },
1375    )?;
1376    if entries_changed {
1377        // Retained bytes only move on insert/evict, so cache hits skip the
1378        // extension-cache accounting write entirely. The shared workspaces
1379        // are deliberately outside this budget: they have their own retention
1380        // cap, and counting them here would evict the whole typed entry
1381        // (including every plan) whenever a shape needs more scratch.
1382        let retained_bytes = plan_cache.plans.retained_bytes();
1383        backend
1384            .cuda_extension_cache()
1385            .update_retained_bytes::<CutensorPlanCacheState>(retained_bytes)?;
1386    }
1387    let state = &mut *plan_cache;
1388    let slot = backend.runtime().stream_slot();
1389    let required = state
1390        .plans
1391        .get(hash, |key| key_matches_spec::<T>(key, spec))
1392        .map(|cached| cached.workspace_size)
1393        .ok_or_else(|| {
1394            Error::runtime_state(
1395                "cutensor_plan_cache",
1396                "cached cuTENSOR contraction was evicted before use",
1397            )
1398        })?;
1399    let current_capacity = state.workspaces[slot]
1400        .as_ref()
1401        .map_or(0, |workspace| workspace.size);
1402    let decision = plan_workspace(
1403        required,
1404        current_capacity,
1405        retained_workspace_bytes(&state.workspaces),
1406        backend.cutensor_workspace_limit(),
1407    );
1408    match decision {
1409        WorkspacePlan::Reuse => {}
1410        WorkspacePlan::Retain(capacity) => {
1411            // Allocate the replacement before dropping the current buffer: a
1412            // failed allocation must not cost this slot its usable scratch.
1413            let replacement = alloc_workspace(backend.runtime(), capacity)?;
1414            state.workspaces[slot] = Some(replacement);
1415        }
1416        WorkspacePlan::Temporary(capacity) => {
1417            // The request does not fit the retention cap. Run it in a
1418            // temporary buffer that is retired through the event queue after
1419            // the call; the slot keeps whatever buffer it already had. Record
1420            // the use so a caller can tell that the cap is binding.
1421            backend.note_cutensor_temporary_workspace();
1422            let temporary = alloc_workspace(backend.runtime(), capacity)?;
1423            let cached = state
1424                .plans
1425                .get(hash, |key| key_matches_spec::<T>(key, spec))
1426                .ok_or_else(|| {
1427                    Error::runtime_state(
1428                        "cutensor_plan_cache",
1429                        "cached cuTENSOR contraction was evicted before use",
1430                    )
1431                })?;
1432            let execute_result = execute(cutensor, &cached.plan, &temporary);
1433            return backend.runtime().finish_vendor_enqueue(
1434                OP,
1435                cross_stream_handles,
1436                execute_result,
1437            );
1438        }
1439    }
1440    let cached = state
1441        .plans
1442        .get(hash, |key| key_matches_spec::<T>(key, spec))
1443        .ok_or_else(|| {
1444            Error::runtime_state(
1445                "cutensor_plan_cache",
1446                "cached cuTENSOR contraction was evicted before use",
1447            )
1448        })?;
1449    // A zero request stores no buffer, so the slot may still be empty here;
1450    // any nonzero requirement above guarantees a retained buffer.
1451    let empty = Workspace::none();
1452    let workspace = state.workspaces[slot].as_ref().unwrap_or(&empty);
1453    debug_assert!(workspace.size >= cached.workspace_size);
1454    // INVARIANT: the cache mutex serializes host enqueues, and all uses of a
1455    // slot's scratch buffer execute in order on the same physical CUDA stream.
1456    let execute_result = execute(cutensor, &cached.plan, workspace);
1457    backend
1458        .runtime()
1459        .finish_vendor_enqueue(OP, cross_stream_handles, execute_result)
1460}
1461
1462/// Zero requests allocate nothing; nonzero requests grow geometrically with a
1463/// 1 MiB floor. `None` when the rounded capacity is not representable, which
1464/// routes the request to the exact-size path instead of failing.
1465pub(super) fn shared_workspace_capacity(requested: u64) -> Option<u64> {
1466    const MIN_CAPACITY: u64 = 1 << 20;
1467    if requested == 0 {
1468        return Some(0);
1469    }
1470    requested.max(MIN_CAPACITY).checked_next_power_of_two()
1471}
1472
1473fn validate_descriptor_alignment(
1474    actual_alignment: u32,
1475    alignment_requirement: u32,
1476    slot: &'static str,
1477) -> crate::Result<()> {
1478    if actual_alignment >= alignment_requirement {
1479        return Ok(());
1480    }
1481    Err(Error::invalid_argument(
1482        OP,
1483        "alignment",
1484        format!(
1485            "{slot} device pointer alignment {actual_alignment} is smaller than the cuTENSOR \
1486             descriptor requirement {alignment_requirement}"
1487        ),
1488    ))
1489}
1490
1491fn raw_stream(rt: &CudaRuntime) -> crate::Result<CutensorCudaStream> {
1492    Ok(rt.raw_cuda_stream()? as usize as CutensorCudaStream)
1493}
1494
1495fn alloc_workspace(rt: &CudaRuntime, workspace_size: u64) -> crate::Result<Workspace> {
1496    if workspace_size == 0 {
1497        return Ok(Workspace::none());
1498    }
1499    // Memory pressure only appears where a workspace is allocated, and eviction
1500    // (which queues a retirement) is what leads here: release completed
1501    // retirements before asking CubeCL for a new block.
1502    rt.workspace_retirements()
1503        .lock()
1504        .unwrap_or_else(|error| error.into_inner())
1505        .drain(rt.state());
1506    let workspace_len =
1507        usize::try_from(workspace_size).map_err(|_| workspace_size_overflow(OP, workspace_size))?;
1508    let handle = rt.client().empty(workspace_len);
1509    let resource = rt
1510        .client()
1511        .get_resource(handle.clone())
1512        .map_err(|err| crate::Error::backend_source(OP, err))?;
1513    let stream = rt.raw_cuda_stream()?;
1514    Ok(Workspace {
1515        _handle: Some(handle),
1516        ptr: cuda_device_ptr_from_addr(resource.resource().ptr, OP)?,
1517        size: workspace_size,
1518        runtime: Some(rt.clone()),
1519        stream,
1520    })
1521}
1522
1523pub(super) fn typed_device_ptr<T: TensorScalar + 'static>(
1524    rt: &CudaRuntime,
1525    tensor: &TypedTensor<T>,
1526    op: &'static str,
1527) -> crate::Result<*mut c_void> {
1528    ensure_resident_on_runtime(rt, tensor, op)?;
1529    let prepared = prepared_tensor_access(tensor, op)?;
1530    let buffer = cubecl_buffer(tensor, op)?;
1531    let addr = memoized_device_addr(rt, buffer, prepared, op)?;
1532    // The residency check above ties this raw FFI pointer to the caller's runtime/device.
1533    cuda_device_ptr_from_addr(addr, op)
1534}
1535
1536/// The device base address of an owned *destination* for a raw vendor call
1537/// that writes it.
1538///
1539/// A queued CubeCL kernel only drops a buffer's memoized address when it
1540/// *writes* the buffer (#1868). A queued *read* leaves the memo valid, so a
1541/// raw vendor write resolved through the memo could be issued ahead of a read
1542/// that precedes it in program order — a write-after-read window (#1949).
1543/// Taking the blocking `get_resource` round trip here pushes every kernel
1544/// queued on this stream onto the CUstream first, so the vendor write is
1545/// ordered after them. Kernels queued on another stream are not covered by
1546/// this barrier; cross-stream reader ordering is tracked separately.
1547///
1548/// The read-only fast path is unchanged: [`typed_device_ptr`] and
1549/// [`memoized_device_addr`] still consult the memo, which is only ever
1550/// resolved from a read or from this write path after the round trip. The
1551/// resolved address is memoized so a later raw read can reuse it (reads need
1552/// no ordering against queued reads, and any intervening queued write
1553/// invalidates it).
1554pub(super) fn write_device_ptr<T: TensorScalar + 'static>(
1555    rt: &CudaRuntime,
1556    tensor: &TypedTensor<T>,
1557    op: &'static str,
1558) -> crate::Result<*mut c_void> {
1559    ensure_resident_on_runtime(rt, tensor, op)?;
1560    let prepared = prepared_tensor_access(tensor, op)?;
1561    let resource = rt
1562        .client()
1563        .get_resource(prepared.into_handle())
1564        .map_err(|err| crate::Error::backend_source(op, err))?;
1565    let addr = resource.resource().ptr;
1566    // See `CubeclBuffer::device_addr` for the address-stability invariant.
1567    cubecl_buffer(tensor, op)?.memoize_device_addr(addr);
1568    // The residency check above ties this raw FFI pointer to the caller's runtime/device.
1569    cuda_device_ptr_from_addr(addr, op)
1570}
1571
1572/// The device base address of `buffer` for a raw vendor call, without a
1573/// blocking server round trip when it is already known.
1574///
1575/// Fast path: reuse the buffer's memoized device address when executing on
1576/// the stream that created the allocation. In that case `get_resource` is
1577/// only a pointer lookup — CubeCL's cross-stream alignment pass skips
1578/// bindings whose creation stream equals the current stream — so no
1579/// synchronization is lost. Any other stream takes the full `get_resource`
1580/// round trip, preserving CubeCL's cross-stream alignment. A queued CubeCL
1581/// write drops the memoized address (#1868), so a raw call never reads ahead
1582/// of it. Borrowed views share their root buffer's memo (#1925).
1583///
1584/// This memoized path is for read-only operands and for fresh internal
1585/// outputs, whose memo is empty. An owned *caller-provided destination* is
1586/// resolved by [`write_device_ptr`], which always takes the round trip so a
1587/// vendor write cannot overtake a queued CubeCL read (#1949).
1588pub(super) fn memoized_device_addr(
1589    rt: &CudaRuntime,
1590    buffer: &CubeclBuffer,
1591    prepared: CubeclPreparedAccess,
1592    op: &'static str,
1593) -> crate::Result<u64> {
1594    let same_stream = StreamId::current() == buffer.handle().stream;
1595    if same_stream {
1596        if let Some(addr) = buffer.cached_device_addr() {
1597            return Ok(addr);
1598        }
1599    }
1600    let resource = rt
1601        .client()
1602        .get_resource(prepared.into_handle())
1603        .map_err(|err| crate::Error::backend_source(op, err))?;
1604    let addr = resource.resource().ptr;
1605    // See `CubeclBuffer::device_addr` for the address-stability invariant.
1606    buffer.memoize_device_addr(addr);
1607    Ok(addr)
1608}
1609
1610fn build_layout(
1611    lhs_shape: &[usize],
1612    rhs_shape: &[usize],
1613    config: &DotGeneralConfig,
1614) -> crate::Result<DotGeneralLayout> {
1615    let lhs_free = free_axes(
1616        lhs_shape.len(),
1617        &config.lhs_contracting_dims,
1618        &config.lhs_batch_dims,
1619    );
1620    let rhs_free = free_axes(
1621        rhs_shape.len(),
1622        &config.rhs_contracting_dims,
1623        &config.rhs_batch_dims,
1624    );
1625
1626    let mut lhs_modes = vec![-1i32; lhs_shape.len()];
1627    let mut rhs_modes = vec![-1i32; rhs_shape.len()];
1628    let mut output_modes =
1629        Vec::with_capacity(lhs_free.len() + rhs_free.len() + config.lhs_batch_dims.len());
1630    let mut output_shape = Vec::with_capacity(output_modes.capacity());
1631    let mut batch_modes = Vec::with_capacity(config.lhs_batch_dims.len());
1632    let mut batch_shape = Vec::with_capacity(config.lhs_batch_dims.len());
1633    let mut next_mode = 0i32;
1634    let mut contracting_elements = 1usize;
1635
1636    for (&lhs_axis, &rhs_axis) in config
1637        .lhs_contracting_dims
1638        .iter()
1639        .zip(&config.rhs_contracting_dims)
1640    {
1641        let mode = next_mode;
1642        next_mode += 1;
1643        lhs_modes[lhs_axis] = mode;
1644        rhs_modes[rhs_axis] = mode;
1645        contracting_elements = contracting_elements
1646            .checked_mul(lhs_shape[lhs_axis])
1647            .ok_or_else(|| {
1648                Error::invalid_argument(
1649                    OP,
1650                    "shape",
1651                    format!(
1652                        "contracting dimension product overflows usize for lhs shape {lhs_shape:?}"
1653                    ),
1654                )
1655            })?;
1656    }
1657
1658    for (&lhs_axis, &rhs_axis) in config.lhs_batch_dims.iter().zip(&config.rhs_batch_dims) {
1659        let mode = next_mode;
1660        next_mode += 1;
1661        lhs_modes[lhs_axis] = mode;
1662        rhs_modes[rhs_axis] = mode;
1663        batch_modes.push(mode);
1664        batch_shape.push(lhs_shape[lhs_axis]);
1665    }
1666
1667    for &lhs_axis in &lhs_free {
1668        let mode = next_mode;
1669        next_mode += 1;
1670        lhs_modes[lhs_axis] = mode;
1671        output_modes.push(mode);
1672        output_shape.push(lhs_shape[lhs_axis]);
1673    }
1674
1675    for &rhs_axis in &rhs_free {
1676        let mode = next_mode;
1677        next_mode += 1;
1678        rhs_modes[rhs_axis] = mode;
1679        output_modes.push(mode);
1680        output_shape.push(rhs_shape[rhs_axis]);
1681    }
1682
1683    output_modes.extend_from_slice(&batch_modes);
1684    output_shape.extend_from_slice(&batch_shape);
1685
1686    let lhs_extents = dims_to_i64(lhs_shape)?;
1687    let rhs_extents = dims_to_i64(rhs_shape)?;
1688    let output_extents = dims_to_i64(&output_shape)?;
1689    let lhs_strides = strides_to_i64(&col_major_strides(lhs_shape)?)?;
1690    let rhs_strides = strides_to_i64(&col_major_strides(rhs_shape)?)?;
1691    let output_strides = strides_to_i64(&col_major_strides(&output_shape)?)?;
1692
1693    Ok(DotGeneralLayout {
1694        lhs_modes,
1695        rhs_modes,
1696        output_modes,
1697        output_shape,
1698        lhs_extents,
1699        rhs_extents,
1700        output_extents,
1701        lhs_strides,
1702        rhs_strides,
1703        output_strides,
1704        contracting_elements,
1705    })
1706}
1707
1708fn dims_to_i64(dims: &[usize]) -> crate::Result<Vec<i64>> {
1709    dims.iter()
1710        .map(|&dim| {
1711            i64::try_from(dim).map_err(|_| {
1712                Error::invalid_argument(
1713                    OP,
1714                    "shape",
1715                    format!("extent {dim} exceeds cuTENSOR i64 limit"),
1716                )
1717            })
1718        })
1719        .collect()
1720}
1721
1722fn strides_to_i64(strides: &[isize]) -> crate::Result<Vec<i64>> {
1723    strides
1724        .iter()
1725        .map(|&stride| {
1726            i64::try_from(stride).map_err(|_| {
1727                Error::invalid_argument(
1728                    OP,
1729                    "stride",
1730                    format!("stride {stride} exceeds cuTENSOR i64 limit"),
1731                )
1732            })
1733        })
1734        .collect()
1735}
1736
1737fn free_axes(rank: usize, contracting: &[usize], batch: &[usize]) -> Vec<usize> {
1738    (0..rank)
1739        .filter(|axis| !contracting.contains(axis) && !batch.contains(axis))
1740        .collect()
1741}
1742
1743fn validate_axis_list(
1744    op: &'static str,
1745    role: &'static str,
1746    axes: &[usize],
1747    rank: usize,
1748) -> crate::Result<()> {
1749    let mut seen = vec![false; rank];
1750    for &axis in axes {
1751        if axis >= rank {
1752            return Err(Error::axis_out_of_bounds(op, axis, rank));
1753        }
1754        if seen[axis] {
1755            return Err(Error::duplicate_axis(op, axis, role));
1756        }
1757        seen[axis] = true;
1758    }
1759    Ok(())
1760}
1761
1762fn validate_role_disjoint(
1763    op: &'static str,
1764    first_role: &'static str,
1765    first_axes: &[usize],
1766    second_role: &'static str,
1767    second_axes: &[usize],
1768) -> crate::Result<()> {
1769    for &axis in first_axes {
1770        if second_axes.contains(&axis) {
1771            return Err(Error::validation(
1772                op,
1773                tenferro_tensor::ValidationError::AxisRoleConflict {
1774                    axis,
1775                    first_role,
1776                    second_role,
1777                },
1778            ));
1779        }
1780    }
1781    Ok(())
1782}
1783
1784fn validate_dot_general(
1785    lhs_shape: &[usize],
1786    rhs_shape: &[usize],
1787    config: &DotGeneralConfig,
1788) -> crate::Result<()> {
1789    if config.lhs_contracting_dims.len() != config.rhs_contracting_dims.len() {
1790        return Err(Error::invalid_argument(
1791            OP,
1792            "contracting_dims",
1793            "lhs/rhs contracting dim counts differ",
1794        ));
1795    }
1796    if config.lhs_batch_dims.len() != config.rhs_batch_dims.len() {
1797        return Err(Error::invalid_argument(
1798            OP,
1799            "batch_dims",
1800            "lhs/rhs batch dim counts differ",
1801        ));
1802    }
1803
1804    let lhs_rank = lhs_shape.len();
1805    let rhs_rank = rhs_shape.len();
1806
1807    validate_axis_list(
1808        OP,
1809        "lhs_contracting",
1810        &config.lhs_contracting_dims,
1811        lhs_rank,
1812    )?;
1813    validate_axis_list(
1814        OP,
1815        "rhs_contracting",
1816        &config.rhs_contracting_dims,
1817        rhs_rank,
1818    )?;
1819    validate_axis_list(OP, "lhs_batch", &config.lhs_batch_dims, lhs_rank)?;
1820    validate_axis_list(OP, "rhs_batch", &config.rhs_batch_dims, rhs_rank)?;
1821    validate_role_disjoint(
1822        OP,
1823        "lhs_contracting",
1824        &config.lhs_contracting_dims,
1825        "lhs_batch",
1826        &config.lhs_batch_dims,
1827    )?;
1828    validate_role_disjoint(
1829        OP,
1830        "rhs_contracting",
1831        &config.rhs_contracting_dims,
1832        "rhs_batch",
1833        &config.rhs_batch_dims,
1834    )?;
1835
1836    for (&lhs_axis, &rhs_axis) in config
1837        .lhs_contracting_dims
1838        .iter()
1839        .zip(&config.rhs_contracting_dims)
1840    {
1841        if lhs_shape[lhs_axis] != rhs_shape[rhs_axis] {
1842            return Err(Error::validation(
1843                OP,
1844                tenferro_tensor::ShapeMismatch::ContractedDimensions {
1845                    lhs_axis,
1846                    lhs_size: lhs_shape[lhs_axis],
1847                    rhs_axis,
1848                    rhs_size: rhs_shape[rhs_axis],
1849                }
1850                .into(),
1851            ));
1852        }
1853    }
1854
1855    for (&lhs_axis, &rhs_axis) in config.lhs_batch_dims.iter().zip(&config.rhs_batch_dims) {
1856        if lhs_shape[lhs_axis] != rhs_shape[rhs_axis] {
1857            return Err(Error::shape_mismatch(
1858                OP,
1859                lhs_shape.to_vec(),
1860                rhs_shape.to_vec(),
1861            ));
1862        }
1863    }
1864
1865    Ok(())
1866}