Skip to main content

tenferro_tensor/
backend.rs

1use crate::config::{
2    CompareDir, DotGeneralConfig, GatherConfig, PadConfig, ScatterConfig, SliceConfig,
3};
4use crate::types::{
5    TensorRank, TensorScalar, TensorView, TensorViewMut, TypedTensor, TypedTensorView,
6    TypedTensorViewMut,
7};
8use crate::validate::validate_convert_dtype;
9use crate::{
10    AllocationDomainId, AllocationId, DType, Error, RuntimeCacheControl, SessionEntryError,
11    ShapeMismatch, Tensor, TensorRead, TensorValue, TensorWrite, ValidationError,
12};
13use num_complex::{Complex32, Complex64};
14
15#[cfg(test)]
16mod tests;
17
18fn read_boundary_error(op: &'static str) -> crate::Error {
19    crate::Error::unsupported(
20        op,
21        "backend does not accept borrowed tensor views at this execution boundary",
22    )
23}
24
25fn validation(op: &'static str, source: ValidationError) -> crate::Error {
26    Error::validation(op, source)
27}
28
29/// Gather any host-owned view, including strided or offset layouts, into a compact owner.
30///
31/// This backs the default [`TensorStructural::to_contiguous_read`] for host-owned
32/// views. Backend-owned storage must be rejected by the caller first; the typed
33/// `to_col_major` gather is metadata-driven element copying and never transfers.
34fn default_materialize_host_view(view: TensorView<'_>) -> crate::Result<Tensor> {
35    fn typed<T: TensorScalar>(view: &TypedTensorView<'_, T>) -> crate::Result<Tensor> {
36        view.to_col_major().map(Tensor::from_typed)
37    }
38
39    match view {
40        TensorView::F32(view) => typed(&view),
41        TensorView::F64(view) => typed(&view),
42        TensorView::I32(view) => typed(&view),
43        TensorView::I64(view) => typed(&view),
44        TensorView::Bool(view) => typed(&view),
45        TensorView::C32(view) => typed(&view),
46        TensorView::C64(view) => typed(&view),
47    }
48}
49
50fn invalid_argument(op: &'static str, argument: &'static str, message: impl Into<String>) -> Error {
51    Error::invalid_argument(op, argument, message)
52}
53
54fn read_tensor<'a>(op: &'static str, input: TensorRead<'a>) -> crate::Result<&'a Tensor> {
55    input.as_tensor().ok_or_else(|| read_boundary_error(op))
56}
57
58/// Return the owned tensor behind `input`, or the read-boundary error for a borrowed view.
59///
60/// This is the read-half default a backend receives when it implements only the
61/// one-shot form of an operation: an owned tensor is delegated to the one-shot
62/// method, and a borrowed view is rejected with [`Error::Unsupported`] instead
63/// of being silently materialized.
64///
65/// Issue #1926 makes the `_read` halves required, so an implementor that relied
66/// on the previous default reproduces it by calling this function:
67///
68/// ```text
69/// fn reduce_sum_read(&mut self, input: TensorRead<'_>, axes: &[usize]) -> Result<Tensor> {
70///     self.reduce_sum(read_owned_tensor("reduce_sum", input)?, axes)
71/// }
72/// ```
73#[doc(hidden)]
74pub fn read_owned_tensor<'a>(op: &'static str, input: TensorRead<'a>) -> crate::Result<&'a Tensor> {
75    read_tensor(op, input)
76}
77
78fn validate_axis_list(
79    op: &'static str,
80    role: &'static str,
81    axes: &[usize],
82    rank: usize,
83) -> crate::Result<()> {
84    let mut seen = vec![false; rank];
85    for &axis in axes {
86        if axis >= rank {
87            return Err(validation(
88                op,
89                ValidationError::AxisOutOfBounds { axis, rank },
90            ));
91        }
92        if seen[axis] {
93            return Err(validation(
94                op,
95                ValidationError::DuplicateAxis { axis, role },
96            ));
97        }
98        seen[axis] = true;
99    }
100    Ok(())
101}
102
103fn validate_role_disjoint(
104    op: &'static str,
105    first_role: &'static str,
106    first_axes: &[usize],
107    second_role: &'static str,
108    second_axes: &[usize],
109) -> crate::Result<()> {
110    for &axis in first_axes {
111        if second_axes.contains(&axis) {
112            return Err(validation(
113                op,
114                ValidationError::AxisRoleConflict {
115                    axis,
116                    first_role,
117                    second_role,
118                },
119            ));
120        }
121    }
122    Ok(())
123}
124
125/// Infer the output shape for a validated dot-general operation.
126#[doc(hidden)]
127pub fn dot_general_output_shape(
128    lhs_shape: &[usize],
129    rhs_shape: &[usize],
130    config: &DotGeneralConfig,
131    op: &'static str,
132) -> crate::Result<Vec<usize>> {
133    if config.lhs_contracting_dims.len() != config.rhs_contracting_dims.len() {
134        return Err(invalid_argument(
135            op,
136            "contracting_dims",
137            "lhs/rhs contracting dim counts differ",
138        ));
139    }
140    if config.lhs_batch_dims.len() != config.rhs_batch_dims.len() {
141        return Err(invalid_argument(
142            op,
143            "batch_dims",
144            "lhs/rhs batch dim counts differ",
145        ));
146    }
147
148    let lhs_rank = lhs_shape.len();
149    let rhs_rank = rhs_shape.len();
150    validate_axis_list(
151        op,
152        "lhs_contracting",
153        &config.lhs_contracting_dims,
154        lhs_rank,
155    )?;
156    validate_axis_list(
157        op,
158        "rhs_contracting",
159        &config.rhs_contracting_dims,
160        rhs_rank,
161    )?;
162    validate_axis_list(op, "lhs_batch", &config.lhs_batch_dims, lhs_rank)?;
163    validate_axis_list(op, "rhs_batch", &config.rhs_batch_dims, rhs_rank)?;
164    validate_role_disjoint(
165        op,
166        "lhs_contracting",
167        &config.lhs_contracting_dims,
168        "lhs_batch",
169        &config.lhs_batch_dims,
170    )?;
171    validate_role_disjoint(
172        op,
173        "rhs_contracting",
174        &config.rhs_contracting_dims,
175        "rhs_batch",
176        &config.rhs_batch_dims,
177    )?;
178
179    for (&lhs_axis, &rhs_axis) in config
180        .lhs_contracting_dims
181        .iter()
182        .zip(&config.rhs_contracting_dims)
183    {
184        if lhs_shape[lhs_axis] != rhs_shape[rhs_axis] {
185            return Err(validation(
186                op,
187                ShapeMismatch::ContractedDimensions {
188                    lhs_axis,
189                    lhs_size: lhs_shape[lhs_axis],
190                    rhs_axis,
191                    rhs_size: rhs_shape[rhs_axis],
192                }
193                .into(),
194            ));
195        }
196    }
197    for (&lhs_axis, &rhs_axis) in config.lhs_batch_dims.iter().zip(&config.rhs_batch_dims) {
198        if lhs_shape[lhs_axis] != rhs_shape[rhs_axis] {
199            return Err(validation(
200                op,
201                ShapeMismatch::ContractedDimensions {
202                    lhs_axis,
203                    lhs_size: lhs_shape[lhs_axis],
204                    rhs_axis,
205                    rhs_size: rhs_shape[rhs_axis],
206                }
207                .into(),
208            ));
209        }
210    }
211
212    let lhs_free = (0..lhs_rank)
213        .filter(|axis| {
214            !config.lhs_contracting_dims.contains(axis) && !config.lhs_batch_dims.contains(axis)
215        })
216        .map(|axis| lhs_shape[axis]);
217    let rhs_free = (0..rhs_rank)
218        .filter(|axis| {
219            !config.rhs_contracting_dims.contains(axis) && !config.rhs_batch_dims.contains(axis)
220        })
221        .map(|axis| rhs_shape[axis]);
222    let batch = config.lhs_batch_dims.iter().map(|&axis| lhs_shape[axis]);
223
224    Ok(lhs_free.chain(rhs_free).chain(batch).collect())
225}
226
227/// Validate output dtype and shape for dot-general read-into dispatch.
228#[doc(hidden)]
229pub fn validate_dot_general_read_into(
230    lhs: &TensorRead<'_>,
231    rhs: &TensorRead<'_>,
232    config: &DotGeneralConfig,
233    out: &TensorWrite<'_>,
234    op: &'static str,
235) -> crate::Result<Vec<usize>> {
236    if lhs.dtype() != rhs.dtype() {
237        return Err(validation(
238            op,
239            ValidationError::DTypeMismatch {
240                expected: lhs.dtype(),
241                actual: rhs.dtype(),
242            },
243        ));
244    }
245    if lhs.dtype() != out.dtype() {
246        return Err(validation(
247            op,
248            ValidationError::DTypeMismatch {
249                expected: lhs.dtype(),
250                actual: out.dtype(),
251            },
252        ));
253    }
254    let expected = dot_general_output_shape(lhs.shape(), rhs.shape(), config, op)?;
255    if out.shape() != expected.as_slice() {
256        return Err(validation(
257            op,
258            ShapeMismatch::ExpectedActual {
259                expected: expected.clone().into(),
260                actual: out.shape().to_vec().into(),
261            }
262            .into(),
263        ));
264    }
265    Ok(expected)
266}
267
268/// Scalar coefficient accepted by contraction accumulation backends.
269///
270/// `ContractionScalar` is intentionally narrower than [`crate::TensorScalar`]:
271/// dot-general accumulation is only defined for floating and complex tensor
272/// dtypes.
273///
274/// # Examples
275///
276/// ```rust
277/// use tenferro_tensor::{ContractionScalar, DType};
278///
279/// let alpha = ContractionScalar::F64(2.0);
280/// assert_eq!(alpha.dtype(), DType::F64);
281/// ```
282#[derive(Clone, Copy, Debug, PartialEq)]
283pub enum ContractionScalar {
284    F32(f32),
285    F64(f64),
286    C32(Complex32),
287    C64(Complex64),
288}
289
290impl ContractionScalar {
291    /// Return this scalar's tensor dtype.
292    ///
293    /// # Examples
294    ///
295    /// ```rust
296    /// use tenferro_tensor::{ContractionScalar, DType};
297    ///
298    /// assert_eq!(ContractionScalar::F32(1.0).dtype(), DType::F32);
299    /// ```
300    pub fn dtype(self) -> DType {
301        match self {
302            Self::F32(_) => DType::F32,
303            Self::F64(_) => DType::F64,
304            Self::C32(_) => DType::C32,
305            Self::C64(_) => DType::C64,
306        }
307    }
308
309    /// Return the multiplicative identity for a supported contraction dtype.
310    ///
311    /// # Examples
312    ///
313    /// ```rust
314    /// use tenferro_tensor::{ContractionScalar, DType};
315    ///
316    /// assert_eq!(ContractionScalar::one(DType::F64).unwrap(), ContractionScalar::F64(1.0));
317    /// assert!(ContractionScalar::one(DType::I32).is_err());
318    /// ```
319    /// # Errors
320    ///
321    /// Returns [`crate::Error::Validation`] with a
322    /// [`crate::ValidationError::DTypeMismatch`] source when `dtype` is `I32`,
323    /// `I64`, or `Bool`, which do not support contraction scalar identities.
324    pub fn one(dtype: DType) -> crate::Result<Self> {
325        match dtype {
326            DType::F32 => Ok(Self::F32(1.0)),
327            DType::F64 => Ok(Self::F64(1.0)),
328            DType::C32 => Ok(Self::C32(Complex32::new(1.0, 0.0))),
329            DType::C64 => Ok(Self::C64(Complex64::new(1.0, 0.0))),
330            DType::I32 | DType::I64 | DType::Bool | DType::External(_) => Err(validation(
331                "ContractionScalar::one",
332                ValidationError::DTypeMismatch {
333                    expected: dtype,
334                    actual: DType::F32,
335                },
336            )),
337        }
338    }
339
340    /// Return the additive identity for a supported contraction dtype.
341    ///
342    /// # Examples
343    ///
344    /// ```rust
345    /// use tenferro_tensor::{ContractionScalar, DType};
346    ///
347    /// assert_eq!(ContractionScalar::zero(DType::F64).unwrap(), ContractionScalar::F64(0.0));
348    /// ```
349    /// # Errors
350    ///
351    /// Returns [`crate::Error::Validation`] with a
352    /// [`crate::ValidationError::DTypeMismatch`] source when `dtype` is `I32`,
353    /// `I64`, or `Bool`, which do not support contraction scalar identities.
354    pub fn zero(dtype: DType) -> crate::Result<Self> {
355        match dtype {
356            DType::F32 => Ok(Self::F32(0.0)),
357            DType::F64 => Ok(Self::F64(0.0)),
358            DType::C32 => Ok(Self::C32(Complex32::new(0.0, 0.0))),
359            DType::C64 => Ok(Self::C64(Complex64::new(0.0, 0.0))),
360            DType::I32 | DType::I64 | DType::Bool | DType::External(_) => Err(validation(
361                "ContractionScalar::zero",
362                ValidationError::DTypeMismatch {
363                    expected: dtype,
364                    actual: DType::F32,
365                },
366            )),
367        }
368    }
369}
370
371/// Output-update semantics for dot-general accumulation.
372///
373/// This keeps contraction axes in [`DotGeneralConfig`] and output update
374/// semantics here, so cached and non-cached backend traits can share the same
375/// accumulation contract.
376///
377/// # Examples
378///
379/// ```rust
380/// use tenferro_tensor::{ContractionScalar, DotGeneralAccumulation, DType};
381///
382/// let accum = DotGeneralAccumulation::overwrite(DType::F64).unwrap();
383/// assert_eq!(accum.alpha, ContractionScalar::F64(1.0));
384/// assert_eq!(accum.beta, ContractionScalar::F64(0.0));
385/// ```
386#[derive(Clone, Copy, Debug, PartialEq)]
387pub struct DotGeneralAccumulation {
388    pub lhs_conj: bool,
389    pub rhs_conj: bool,
390    pub alpha: ContractionScalar,
391    pub beta: ContractionScalar,
392}
393
394/// One matrix multiply in a grouped GEMM over shared flat buffers.
395///
396/// Offsets are element offsets into the corresponding shared lhs, rhs, and
397/// output buffers. Each job computes a column-major `rows x cols` output block
398/// from a column-major `rows x contracted` lhs block and a column-major
399/// `contracted x cols` rhs block.
400///
401/// Provider implementations receive these descriptors through the public
402/// grouped-GEMM request accessor. The engine validates ranges and pairwise
403/// output disjointness before provider entry.
404#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
405pub struct GroupedGemmJob {
406    out_offset: usize,
407    lhs_offset: usize,
408    rhs_offset: usize,
409    rows: usize,
410    contracted: usize,
411    cols: usize,
412}
413
414impl GroupedGemmJob {
415    /// Construct a column-major grouped-GEMM job over shared flat buffers.
416    #[allow(clippy::too_many_arguments)]
417    pub fn new(
418        out_offset: usize,
419        lhs_offset: usize,
420        rhs_offset: usize,
421        rows: usize,
422        contracted: usize,
423        cols: usize,
424    ) -> Self {
425        Self {
426            out_offset,
427            lhs_offset,
428            rhs_offset,
429            rows,
430            contracted,
431            cols,
432        }
433    }
434
435    /// Return the output element offset.
436    pub fn out_offset(&self) -> usize {
437        self.out_offset
438    }
439
440    /// Return the left-input element offset.
441    pub fn lhs_offset(&self) -> usize {
442        self.lhs_offset
443    }
444
445    /// Return the right-input element offset.
446    pub fn rhs_offset(&self) -> usize {
447        self.rhs_offset
448    }
449
450    /// Return the output row count.
451    pub fn rows(&self) -> usize {
452        self.rows
453    }
454
455    /// Return the contracted dimension.
456    pub fn contracted(&self) -> usize {
457        self.contracted
458    }
459
460    /// Return the output column count.
461    pub fn cols(&self) -> usize {
462        self.cols
463    }
464}
465
466/// Shared scalar/update metadata for grouped GEMM execution.
467#[doc(hidden)]
468#[derive(Clone, Copy, Debug, PartialEq)]
469pub struct GroupedGemmConfig<'a> {
470    jobs: &'a [GroupedGemmJob],
471    accumulation: DotGeneralAccumulation,
472}
473
474impl<'a> GroupedGemmConfig<'a> {
475    pub fn new(jobs: &'a [GroupedGemmJob], accumulation: DotGeneralAccumulation) -> Self {
476        Self { jobs, accumulation }
477    }
478
479    pub fn jobs(&self) -> &'a [GroupedGemmJob] {
480        self.jobs
481    }
482
483    pub fn accumulation(&self) -> DotGeneralAccumulation {
484        self.accumulation
485    }
486}
487
488impl DotGeneralAccumulation {
489    fn identity(
490        op: &'static str,
491        dtype: DType,
492        multiplicative: bool,
493    ) -> crate::Result<ContractionScalar> {
494        let result = if multiplicative {
495            ContractionScalar::one(dtype)
496        } else {
497            ContractionScalar::zero(dtype)
498        };
499        result.map_err(|error| match error {
500            Error::Validation { source, .. } => validation(op, source),
501            error => error,
502        })
503    }
504
505    /// Return overwrite semantics, `out = lhs dot rhs`, for `dtype`.
506    ///
507    /// # Examples
508    ///
509    /// ```rust
510    /// use tenferro_tensor::{ContractionScalar, DotGeneralAccumulation, DType};
511    ///
512    /// let accum = DotGeneralAccumulation::overwrite(DType::F64).unwrap();
513    /// assert_eq!(accum.alpha, ContractionScalar::F64(1.0));
514    /// assert_eq!(accum.beta, ContractionScalar::F64(0.0));
515    /// ```
516    ///
517    /// # Errors
518    ///
519    /// Returns [`crate::Error::Validation`] with a
520    /// [`crate::ValidationError::DTypeMismatch`] source when `dtype` does not
521    /// support contraction scalar identities.
522    pub fn overwrite(dtype: DType) -> crate::Result<Self> {
523        Ok(Self {
524            lhs_conj: false,
525            rhs_conj: false,
526            alpha: Self::identity("DotGeneralAccumulation::overwrite", dtype, true)?,
527            beta: Self::identity("DotGeneralAccumulation::overwrite", dtype, false)?,
528        })
529    }
530
531    /// Return additive update semantics, `out += lhs dot rhs`, for `dtype`.
532    ///
533    /// # Examples
534    ///
535    /// ```rust
536    /// use tenferro_tensor::{ContractionScalar, DType, DotGeneralAccumulation};
537    ///
538    /// let accum = DotGeneralAccumulation::add_to(DType::F64)?;
539    /// assert_eq!(accum.alpha, ContractionScalar::F64(1.0));
540    /// assert_eq!(accum.beta, ContractionScalar::F64(1.0));
541    /// # Ok::<(), tenferro_tensor::Error>(())
542    /// ```
543    /// # Errors
544    ///
545    /// Returns [`crate::Error::Validation`] with a
546    /// [`crate::ValidationError::DTypeMismatch`] source when `dtype` does not
547    /// support contraction scalar identities.
548    pub fn add_to(dtype: DType) -> crate::Result<Self> {
549        Ok(Self {
550            lhs_conj: false,
551            rhs_conj: false,
552            alpha: Self::identity("DotGeneralAccumulation::add_to", dtype, true)?,
553            beta: Self::identity("DotGeneralAccumulation::add_to", dtype, true)?,
554        })
555    }
556
557    /// Return scaled update semantics, `out = alpha * lhs dot rhs + beta * out`.
558    ///
559    /// # Examples
560    ///
561    /// ```rust
562    /// use tenferro_tensor::{ContractionScalar, DotGeneralAccumulation};
563    ///
564    /// let accum = DotGeneralAccumulation::scaled(
565    ///     ContractionScalar::F32(0.5),
566    ///     ContractionScalar::F32(2.0),
567    /// )?;
568    /// assert_eq!(accum.alpha, ContractionScalar::F32(0.5));
569    /// # Ok::<(), tenferro_tensor::Error>(())
570    /// ```
571    /// # Errors
572    ///
573    /// Returns [`crate::Error::Validation`] with a
574    /// [`crate::ValidationError::DTypeMismatch`] source when `alpha` and `beta`
575    /// have different dtypes.
576    pub fn scaled(alpha: ContractionScalar, beta: ContractionScalar) -> crate::Result<Self> {
577        if alpha.dtype() != beta.dtype() {
578            return Err(validation(
579                "DotGeneralAccumulation::scaled",
580                ValidationError::DTypeMismatch {
581                    expected: alpha.dtype(),
582                    actual: beta.dtype(),
583                },
584            ));
585        }
586        Ok(Self {
587            lhs_conj: false,
588            rhs_conj: false,
589            alpha,
590            beta,
591        })
592    }
593
594    fn validate_for_dtype(self, dtype: DType) -> crate::Result<()> {
595        for scalar in [self.alpha, self.beta] {
596            if scalar.dtype() != dtype {
597                return Err(validation(
598                    "dot_general",
599                    ValidationError::DTypeMismatch {
600                        expected: scalar.dtype(),
601                        actual: dtype,
602                    },
603                ));
604            }
605        }
606        Ok(())
607    }
608}
609
610#[doc(hidden)]
611pub fn validate_dot_general_accumulation(
612    lhs: &TensorRead<'_>,
613    rhs: &TensorRead<'_>,
614    config: &DotGeneralConfig,
615    accumulation: DotGeneralAccumulation,
616    out: &TensorWrite<'_>,
617    op: &'static str,
618) -> crate::Result<Vec<usize>> {
619    let shape = validate_dot_general_read_into(lhs, rhs, config, out, op)?;
620    accumulation.validate_for_dtype(lhs.dtype())?;
621    Ok(shape)
622}
623
624#[doc(hidden)]
625pub fn dot_general_accum_via_temp<B: TensorDot + ?Sized>(
626    backend: &mut B,
627    lhs: TensorRead<'_>,
628    rhs: TensorRead<'_>,
629    config: &DotGeneralConfig,
630    accumulation: DotGeneralAccumulation,
631    mut out: TensorWrite<'_>,
632) -> crate::Result<()> {
633    validate_dot_general_accumulation(&lhs, &rhs, config, accumulation, &out, "dot_general")?;
634    let dot = backend.dot_general_with_conj_read(
635        lhs,
636        rhs,
637        config,
638        accumulation.lhs_conj,
639        accumulation.rhs_conj,
640    )?;
641    accumulate_dot_result_into(&dot, accumulation, &mut out)
642}
643
644fn grouped_checked_product(
645    op: &'static str,
646    role: &'static str,
647    dims: &[usize],
648) -> crate::Result<usize> {
649    dims.iter().try_fold(1usize, |acc, &dim| {
650        acc.checked_mul(dim).ok_or_else(|| {
651            invalid_argument(
652                op,
653                role,
654                format!("logical element count overflows usize for shape {dims:?}"),
655            )
656        })
657    })
658}
659
660fn checked_gemm_span(
661    op: &'static str,
662    role: &'static str,
663    offset: usize,
664    rows: usize,
665    cols: usize,
666) -> crate::Result<Option<std::ops::Range<usize>>> {
667    let len = rows.checked_mul(cols).ok_or_else(|| {
668        invalid_argument(
669            op,
670            role,
671            format!("matrix element count overflows usize: rows={rows} cols={cols}"),
672        )
673    })?;
674    if len == 0 {
675        return Ok(None);
676    }
677    let end = offset.checked_add(len).ok_or_else(|| {
678        invalid_argument(
679            op,
680            role,
681            format!("matrix range overflows usize: offset={offset} len={len}"),
682        )
683    })?;
684    Ok(Some(offset..end))
685}
686
687fn validate_grouped_gemm_range(
688    op: &'static str,
689    role: &'static str,
690    len: usize,
691    range: Option<std::ops::Range<usize>>,
692) -> crate::Result<()> {
693    let Some(range) = range else {
694        return Ok(());
695    };
696    if range.end > len {
697        return Err(invalid_argument(
698            op,
699            role,
700            format!(
701                "matrix range {}..{} exceeds shared buffer logical length {len}",
702                range.start, range.end
703            ),
704        ));
705    }
706    Ok(())
707}
708
709#[doc(hidden)]
710pub fn validate_grouped_gemm(
711    lhs: &TensorRead<'_>,
712    rhs: &TensorRead<'_>,
713    out: &TensorWrite<'_>,
714    config: &GroupedGemmConfig<'_>,
715    op: &'static str,
716) -> crate::Result<()> {
717    if lhs.dtype() != rhs.dtype() {
718        return Err(validation(
719            op,
720            ValidationError::DTypeMismatch {
721                expected: lhs.dtype(),
722                actual: rhs.dtype(),
723            },
724        ));
725    }
726    if lhs.dtype() != out.dtype() {
727        return Err(validation(
728            op,
729            ValidationError::DTypeMismatch {
730                expected: lhs.dtype(),
731                actual: out.dtype(),
732            },
733        ));
734    }
735    config.accumulation.validate_for_dtype(lhs.dtype())?;
736
737    let lhs_len = grouped_checked_product(op, "lhs", lhs.shape())?;
738    let rhs_len = grouped_checked_product(op, "rhs", rhs.shape())?;
739    let out_len = grouped_checked_product(op, "out", out.shape())?;
740    // Grouped GEMM job count is runtime-controlled and can be large. Keep the
741    // validation ranges in a reserved Vec, not SmallVec, so arbitrary batches
742    // avoid inline-capacity tuning and can be sorted for O(n log n) overlap
743    // validation.
744    let mut out_ranges = Vec::<(usize, std::ops::Range<usize>)>::with_capacity(config.jobs.len());
745    for (idx, job) in config.jobs.iter().enumerate() {
746        validate_grouped_gemm_range(
747            op,
748            "lhs",
749            lhs_len,
750            checked_gemm_span(op, "lhs", job.lhs_offset, job.rows, job.contracted)?,
751        )?;
752        validate_grouped_gemm_range(
753            op,
754            "rhs",
755            rhs_len,
756            checked_gemm_span(op, "rhs", job.rhs_offset, job.contracted, job.cols)?,
757        )?;
758        let out_range = checked_gemm_span(op, "out", job.out_offset, job.rows, job.cols)?;
759        validate_grouped_gemm_range(op, "out", out_len, out_range.clone())?;
760        if let Some(out_range) = out_range {
761            out_ranges.push((idx, out_range));
762        }
763    }
764    out_ranges.sort_unstable_by_key(|(_, range)| range.start);
765    for pair in out_ranges.windows(2) {
766        let (prev_idx, previous) = &pair[0];
767        let (idx, current) = &pair[1];
768        if previous.end > current.start {
769            return Err(invalid_argument(
770                op,
771                "jobs",
772                format!(
773                    "grouped GEMM output range for job {idx} overlaps job {prev_idx} range {}..{}",
774                    previous.start, previous.end
775                ),
776            ));
777        }
778    }
779    Ok(())
780}
781
782fn add_element_offsets(
783    op: &'static str,
784    base: isize,
785    offset: usize,
786    role: &'static str,
787) -> crate::Result<isize> {
788    let offset = isize::try_from(offset).map_err(|_| {
789        invalid_argument(op, role, format!("offset {offset} does not fit in isize"))
790    })?;
791    base.checked_add(offset).ok_or_else(|| {
792        invalid_argument(
793            op,
794            role,
795            format!("offset overflows isize: base={base} offset={offset}"),
796        )
797    })
798}
799
800fn dim_stride(op: &'static str, dim: usize, role: &'static str) -> crate::Result<isize> {
801    isize::try_from(dim).map_err(|_| {
802        invalid_argument(
803            op,
804            role,
805            format!("leading dimension {dim} does not fit in isize"),
806        )
807    })
808}
809
810fn typed_read_storage<'a, T: crate::TensorScalar>(
811    tensor: &'a TypedTensor<T>,
812    op: &'static str,
813) -> crate::Result<(&'a [T], isize)> {
814    tensor.host_data().map(|data| (data, 0)).map_err(|_| {
815        crate::Error::runtime_state(
816            op,
817            "grouped GEMM default path requires host-backed tensor storage",
818        )
819    })
820}
821
822fn grouped_gemm_default_config() -> DotGeneralConfig {
823    // DotGeneralConfig owns Vec fields, so this rank-2 fallback config follows
824    // that API boundary rather than introducing SmallVec locally.
825    DotGeneralConfig {
826        lhs_contracting_dims: [1].as_slice().into(),
827        rhs_contracting_dims: [0].as_slice().into(),
828        lhs_batch_dims: Default::default(),
829        rhs_batch_dims: Default::default(),
830    }
831}
832
833trait GroupedGemmDType<T> {
834    fn wrap_read(view: TypedTensorView<'_, T>) -> TensorView<'_>;
835    fn wrap_write(view: TypedTensorViewMut<'_, T>) -> TensorViewMut<'_>;
836}
837
838struct GroupedF32;
839struct GroupedF64;
840struct GroupedC32;
841struct GroupedC64;
842
843impl GroupedGemmDType<f32> for GroupedF32 {
844    fn wrap_read(view: TypedTensorView<'_, f32>) -> TensorView<'_> {
845        TensorView::F32(view)
846    }
847
848    fn wrap_write(view: TypedTensorViewMut<'_, f32>) -> TensorViewMut<'_> {
849        TensorViewMut::F32(view)
850    }
851}
852
853impl GroupedGemmDType<f64> for GroupedF64 {
854    fn wrap_read(view: TypedTensorView<'_, f64>) -> TensorView<'_> {
855        TensorView::F64(view)
856    }
857
858    fn wrap_write(view: TypedTensorViewMut<'_, f64>) -> TensorViewMut<'_> {
859        TensorViewMut::F64(view)
860    }
861}
862
863impl GroupedGemmDType<Complex32> for GroupedC32 {
864    fn wrap_read(view: TypedTensorView<'_, Complex32>) -> TensorView<'_> {
865        TensorView::C32(view)
866    }
867
868    fn wrap_write(view: TypedTensorViewMut<'_, Complex32>) -> TensorViewMut<'_> {
869        TensorViewMut::C32(view)
870    }
871}
872
873impl GroupedGemmDType<Complex64> for GroupedC64 {
874    fn wrap_read(view: TypedTensorView<'_, Complex64>) -> TensorView<'_> {
875        TensorView::C64(view)
876    }
877
878    fn wrap_write(view: TypedTensorViewMut<'_, Complex64>) -> TensorViewMut<'_> {
879        TensorViewMut::C64(view)
880    }
881}
882
883#[allow(clippy::too_many_arguments)]
884fn grouped_gemm_default_loop<B, T, V>(
885    backend: &mut B,
886    lhs_data: &[T],
887    lhs_base: isize,
888    rhs_data: &[T],
889    rhs_base: isize,
890    out_view: &mut TypedTensorViewMut<'_, T>,
891    config: &GroupedGemmConfig<'_>,
892) -> crate::Result<()>
893where
894    B: TensorDot + ?Sized,
895    T: 'static,
896    V: GroupedGemmDType<T>,
897{
898    let op = "grouped_gemm";
899    let dot_config = grouped_gemm_default_config();
900    for job in config.jobs {
901        let lhs_offset = add_element_offsets(op, lhs_base, job.lhs_offset, "lhs")?;
902        let rhs_offset = add_element_offsets(op, rhs_base, job.rhs_offset, "rhs")?;
903        let out_offset = add_element_offsets(op, out_view.offset(), job.out_offset, "out")?;
904        let lhs_rows = dim_stride(op, job.rows, "lhs")?;
905        let rhs_rows = dim_stride(op, job.contracted, "rhs")?;
906        let out_rows = dim_stride(op, job.rows, "out")?;
907        // TypedTensorView constructors own Vec shape/stride metadata. These
908        // fallback rank-2 views are short-lived, but SmallVec is not usable
909        // without changing the view API.
910        let lhs_matrix = TypedTensorView::from_slice(
911            vec![job.rows, job.contracted],
912            vec![1, lhs_rows],
913            lhs_offset,
914            lhs_data,
915        )?;
916        let rhs_matrix = TypedTensorView::from_slice(
917            vec![job.contracted, job.cols],
918            vec![1, rhs_rows],
919            rhs_offset,
920            rhs_data,
921        )?;
922        let out_storage = out_view.host_storage_mut()?;
923        let out_matrix = TypedTensorViewMut::from_slice(
924            vec![job.rows, job.cols],
925            vec![1, out_rows],
926            out_offset,
927            out_storage,
928        )?;
929        backend.dot_general_read_into_accum(
930            TensorRead::from_view(V::wrap_read(lhs_matrix)),
931            TensorRead::from_view(V::wrap_read(rhs_matrix)),
932            &dot_config,
933            config.accumulation,
934            TensorWrite::from_view(V::wrap_write(out_matrix)),
935        )?;
936    }
937    Ok(())
938}
939
940#[doc(hidden)]
941pub fn grouped_gemm_via_sequential<B>(
942    backend: &mut B,
943    lhs: TensorRead<'_>,
944    rhs: TensorRead<'_>,
945    config: &GroupedGemmConfig<'_>,
946    mut out: TensorWrite<'_>,
947) -> crate::Result<()>
948where
949    B: TensorDot + ?Sized,
950{
951    validate_grouped_gemm(&lhs, &rhs, &out, config, "grouped_gemm")?;
952    macro_rules! dispatch {
953        ($variant:ident, $scalar:ty, $wrapper:ty) => {
954            match (&lhs, &rhs, &mut out) {
955                (TensorRead::Tensor(a), TensorRead::Tensor(b), TensorWrite::Tensor(c))
956                    if a.dtype() == <$scalar as crate::TensorScalar>::dtype()
957                        && b.dtype() == <$scalar as crate::TensorScalar>::dtype()
958                        && c.dtype() == <$scalar as crate::TensorScalar>::dtype() =>
959                {
960                    let a = a
961                        .as_typed::<$scalar>()
962                        .expect("the dtype guard selects this arm");
963                    let b = b
964                        .as_typed::<$scalar>()
965                        .expect("the dtype guard selects this arm");
966                    let c = c
967                        .as_typed_mut::<$scalar>()
968                        .expect("the dtype guard selects this arm");
969                    let (a_data, a_base) = typed_read_storage(a, "grouped_gemm")?;
970                    let (b_data, b_base) = typed_read_storage(b, "grouped_gemm")?;
971                    let mut c_view = c.as_view_mut();
972                    return grouped_gemm_default_loop::<_, _, $wrapper>(
973                        backend,
974                        a_data,
975                        a_base,
976                        b_data,
977                        b_base,
978                        &mut c_view,
979                        config,
980                    );
981                }
982                (
983                    TensorRead::Tensor(a),
984                    TensorRead::View(TensorView::$variant(b)),
985                    TensorWrite::Tensor(c),
986                ) if a.dtype() == <$scalar as crate::TensorScalar>::dtype()
987                    && c.dtype() == <$scalar as crate::TensorScalar>::dtype() =>
988                {
989                    let a = a
990                        .as_typed::<$scalar>()
991                        .expect("the dtype guard selects this arm");
992                    let c = c
993                        .as_typed_mut::<$scalar>()
994                        .expect("the dtype guard selects this arm");
995                    let (a_data, a_base) = typed_read_storage(a, "grouped_gemm")?;
996                    let mut c_view = c.as_view_mut();
997                    return grouped_gemm_default_loop::<_, _, $wrapper>(
998                        backend,
999                        a_data,
1000                        a_base,
1001                        b.host_storage()?,
1002                        b.offset(),
1003                        &mut c_view,
1004                        config,
1005                    );
1006                }
1007                (
1008                    TensorRead::View(TensorView::$variant(a)),
1009                    TensorRead::Tensor(b),
1010                    TensorWrite::Tensor(c),
1011                ) if b.dtype() == <$scalar as crate::TensorScalar>::dtype()
1012                    && c.dtype() == <$scalar as crate::TensorScalar>::dtype() =>
1013                {
1014                    let b = b
1015                        .as_typed::<$scalar>()
1016                        .expect("the dtype guard selects this arm");
1017                    let c = c
1018                        .as_typed_mut::<$scalar>()
1019                        .expect("the dtype guard selects this arm");
1020                    let (b_data, b_base) = typed_read_storage(b, "grouped_gemm")?;
1021                    let mut c_view = c.as_view_mut();
1022                    return grouped_gemm_default_loop::<_, _, $wrapper>(
1023                        backend,
1024                        a.host_storage()?,
1025                        a.offset(),
1026                        b_data,
1027                        b_base,
1028                        &mut c_view,
1029                        config,
1030                    );
1031                }
1032                (
1033                    TensorRead::View(TensorView::$variant(a)),
1034                    TensorRead::View(TensorView::$variant(b)),
1035                    TensorWrite::Tensor(c),
1036                ) if c.dtype() == <$scalar as crate::TensorScalar>::dtype() => {
1037                    let c = c
1038                        .as_typed_mut::<$scalar>()
1039                        .expect("the dtype guard selects this arm");
1040                    let mut c_view = c.as_view_mut();
1041                    return grouped_gemm_default_loop::<_, _, $wrapper>(
1042                        backend,
1043                        a.host_storage()?,
1044                        a.offset(),
1045                        b.host_storage()?,
1046                        b.offset(),
1047                        &mut c_view,
1048                        config,
1049                    );
1050                }
1051                (
1052                    TensorRead::Tensor(a),
1053                    TensorRead::Tensor(b),
1054                    TensorWrite::View(TensorViewMut::$variant(c)),
1055                ) if a.dtype() == <$scalar as crate::TensorScalar>::dtype()
1056                    && b.dtype() == <$scalar as crate::TensorScalar>::dtype() =>
1057                {
1058                    let a = a
1059                        .as_typed::<$scalar>()
1060                        .expect("the dtype guard selects this arm");
1061                    let b = b
1062                        .as_typed::<$scalar>()
1063                        .expect("the dtype guard selects this arm");
1064                    let (a_data, a_base) = typed_read_storage(a, "grouped_gemm")?;
1065                    let (b_data, b_base) = typed_read_storage(b, "grouped_gemm")?;
1066                    return grouped_gemm_default_loop::<_, _, $wrapper>(
1067                        backend, a_data, a_base, b_data, b_base, c, config,
1068                    );
1069                }
1070                (
1071                    TensorRead::Tensor(a),
1072                    TensorRead::View(TensorView::$variant(b)),
1073                    TensorWrite::View(TensorViewMut::$variant(c)),
1074                ) if a.dtype() == <$scalar as crate::TensorScalar>::dtype() => {
1075                    let a = a
1076                        .as_typed::<$scalar>()
1077                        .expect("the dtype guard selects this arm");
1078                    let (a_data, a_base) = typed_read_storage(a, "grouped_gemm")?;
1079                    return grouped_gemm_default_loop::<_, _, $wrapper>(
1080                        backend,
1081                        a_data,
1082                        a_base,
1083                        b.host_storage()?,
1084                        b.offset(),
1085                        c,
1086                        config,
1087                    );
1088                }
1089                (
1090                    TensorRead::View(TensorView::$variant(a)),
1091                    TensorRead::Tensor(b),
1092                    TensorWrite::View(TensorViewMut::$variant(c)),
1093                ) if b.dtype() == <$scalar as crate::TensorScalar>::dtype() => {
1094                    let b = b
1095                        .as_typed::<$scalar>()
1096                        .expect("the dtype guard selects this arm");
1097                    let (b_data, b_base) = typed_read_storage(b, "grouped_gemm")?;
1098                    return grouped_gemm_default_loop::<_, _, $wrapper>(
1099                        backend,
1100                        a.host_storage()?,
1101                        a.offset(),
1102                        b_data,
1103                        b_base,
1104                        c,
1105                        config,
1106                    );
1107                }
1108                (
1109                    TensorRead::View(TensorView::$variant(a)),
1110                    TensorRead::View(TensorView::$variant(b)),
1111                    TensorWrite::View(TensorViewMut::$variant(c)),
1112                ) => {
1113                    return grouped_gemm_default_loop::<_, _, $wrapper>(
1114                        backend,
1115                        a.host_storage()?,
1116                        a.offset(),
1117                        b.host_storage()?,
1118                        b.offset(),
1119                        c,
1120                        config,
1121                    );
1122                }
1123                _ => {}
1124            }
1125        };
1126    }
1127
1128    dispatch!(F32, f32, GroupedF32);
1129    dispatch!(F64, f64, GroupedF64);
1130    dispatch!(C32, Complex32, GroupedC32);
1131    dispatch!(C64, Complex64, GroupedC64);
1132    Err(validation(
1133        "grouped_gemm",
1134        ValidationError::DTypeMismatch {
1135            expected: lhs.dtype(),
1136            actual: out.dtype(),
1137        },
1138    ))
1139}
1140
1141fn grouped_gemm_default<B>(
1142    backend: &mut B,
1143    lhs: TensorRead<'_>,
1144    rhs: TensorRead<'_>,
1145    config: &GroupedGemmConfig<'_>,
1146    out: TensorWrite<'_>,
1147) -> crate::Result<()>
1148where
1149    B: TensorDot + ?Sized,
1150{
1151    grouped_gemm_via_sequential(backend, lhs, rhs, config, out)
1152}
1153
1154#[doc(hidden)]
1155pub fn accumulate_dot_result_into(
1156    dot: &Tensor,
1157    accumulation: DotGeneralAccumulation,
1158    out: &mut TensorWrite<'_>,
1159) -> crate::Result<()> {
1160    macro_rules! dispatch {
1161        ($variant:ident, $ty:ty) => {
1162            if dot.dtype() == <$ty as crate::TensorScalar>::dtype() {
1163                let (ContractionScalar::$variant(alpha), ContractionScalar::$variant(beta)) =
1164                    (accumulation.alpha, accumulation.beta)
1165                else {
1166                    return Err(validation(
1167                        "dot_general",
1168                        ValidationError::DTypeMismatch {
1169                            expected: dot.dtype(),
1170                            actual: accumulation.alpha.dtype(),
1171                        },
1172                    ));
1173                };
1174                let dot = dot
1175                    .as_typed::<$ty>()
1176                    .expect("the dtype guard selects this arm");
1177                match out {
1178                    TensorWrite::Tensor(out) => {
1179                        let out = out
1180                            .as_typed_mut::<$ty>()
1181                            .expect("the dtype guard selects this arm");
1182                        let mut out = out.as_view_mut();
1183                        accumulate_typed(dot.as_slice()?, alpha, beta, &mut out)?;
1184                        return Ok(());
1185                    }
1186                    TensorWrite::View(crate::TensorViewMut::$variant(out)) => {
1187                        accumulate_typed(dot.as_slice()?, alpha, beta, out)?;
1188                        return Ok(());
1189                    }
1190                    _ => {}
1191                }
1192            }
1193        };
1194    }
1195
1196    dispatch!(F32, f32);
1197    dispatch!(F64, f64);
1198    dispatch!(C32, Complex32);
1199    dispatch!(C64, Complex64);
1200
1201    Err(validation(
1202        "dot_general",
1203        ValidationError::DTypeMismatch {
1204            expected: accumulation.alpha.dtype(),
1205            actual: dot.dtype(),
1206        },
1207    ))
1208}
1209
1210fn accumulate_typed<T>(
1211    dot: &[T],
1212    alpha: T,
1213    beta: T,
1214    out: &mut TypedTensorViewMut<'_, T>,
1215) -> crate::Result<()>
1216where
1217    T: Copy
1218        + PartialEq
1219        + std::ops::Add<Output = T>
1220        + std::ops::Mul<Output = T>
1221        + num_traits::Zero
1222        + 'static,
1223{
1224    let beta_is_zero = beta == T::zero();
1225    if let Some(output) = compact_host_accumulation_slice(out, dot.len())? {
1226        for (output, dot_value) in output.iter_mut().zip(dot.iter().copied()) {
1227            // INVARIANT: beta == 0 follows BLAS GEMM semantics and does not read
1228            // the existing output element; beta != 0 requires an initialized
1229            // TensorWrite target and performs a read-modify-write update.
1230            *output = if beta_is_zero {
1231                alpha * dot_value
1232            } else {
1233                alpha * dot_value + beta * *output
1234            };
1235        }
1236        return Ok(());
1237    }
1238
1239    for (linear, dot_value) in dot.iter().copied().enumerate() {
1240        let indices = flat_to_multi_for_shape(out.shape(), linear);
1241        let output = out.get_mut(&indices).ok_or_else(|| {
1242            invalid_argument(
1243                "dot_general",
1244                "output",
1245                format!("index {indices:?} is outside accumulation target"),
1246            )
1247        })?;
1248        // INVARIANT: beta == 0 follows BLAS GEMM semantics and does not read
1249        // the existing output element; beta != 0 requires an initialized
1250        // TensorWrite target and performs a read-modify-write update.
1251        *output = if beta_is_zero {
1252            alpha * dot_value
1253        } else {
1254            alpha * dot_value + beta * *output
1255        };
1256    }
1257    Ok(())
1258}
1259
1260fn compact_host_accumulation_slice<'a, T: 'static>(
1261    out: &'a mut TypedTensorViewMut<'_, T>,
1262    expected_len: usize,
1263) -> crate::Result<Option<&'a mut [T]>> {
1264    if out.backend_buffer().is_some()
1265        || out.n_elements() != expected_len
1266        || !out.is_col_major_contiguous()?
1267    {
1268        return Ok(None);
1269    }
1270
1271    let start = usize::try_from(out.offset()).map_err(|_| {
1272        invalid_argument("dot_general", "output", "compact output offset is negative")
1273    })?;
1274    let end = start
1275        .checked_add(expected_len)
1276        .ok_or_else(|| validation("dot_general", ValidationError::IntegerOverflow))?;
1277    out.host_storage_mut()?
1278        .get_mut(start..end)
1279        .map(Some)
1280        .ok_or_else(|| {
1281            invalid_argument(
1282                "dot_general",
1283                "output",
1284                "compact output is outside its backing storage",
1285            )
1286        })
1287}
1288
1289fn flat_to_multi_for_shape(shape: &[usize], mut linear: usize) -> Vec<usize> {
1290    let mut indices = Vec::with_capacity(shape.len());
1291    for &dim in shape {
1292        if dim == 0 {
1293            indices.push(0);
1294        } else {
1295            indices.push(linear % dim);
1296            linear /= dim;
1297        }
1298    }
1299    indices
1300}
1301
1302/// Canonical elementwise fusion plan shared between segmented execution and backends.
1303#[doc(hidden)]
1304#[derive(Clone, Debug, Hash, PartialEq, Eq)]
1305pub struct ElementwiseFusionPlan {
1306    dtype: crate::DType,
1307    input_count: usize,
1308    // Keep view metadata in Vecs. A/B benchmarking on the broadcast_mul
1309    // path showed SmallVec made this metadata path about 6-7% slower.
1310    input_views: Vec<ElementwiseFusionInputView>,
1311    outputs: Vec<usize>,
1312    ops: Vec<ElementwiseFusionInst>,
1313}
1314
1315/// Metadata-only view applied to one backend fusion input.
1316#[doc(hidden)]
1317#[derive(Clone, Debug, Hash, PartialEq, Eq)]
1318pub enum ElementwiseFusionInputView {
1319    Identity,
1320    BroadcastInDim {
1321        // Vec is intentional here; see ElementwiseFusionPlan::input_views.
1322        shape: Vec<usize>,
1323        dims: Vec<usize>,
1324    },
1325}
1326
1327/// One node in a canonical elementwise fusion plan.
1328#[doc(hidden)]
1329#[derive(Clone, Debug, Hash, PartialEq, Eq)]
1330pub struct ElementwiseFusionInst {
1331    op: ElementwiseFusionOp,
1332    inputs: Vec<usize>,
1333}
1334
1335tenferro_core_ops::define_elementwise_fusion_op!();
1336
1337impl ElementwiseFusionPlan {
1338    /// Build a backend elementwise fusion plan.
1339    ///
1340    /// # Examples
1341    ///
1342    /// ```rust
1343    /// use tenferro_tensor::backend::{
1344    ///     ElementwiseFusionInst, ElementwiseFusionOp, ElementwiseFusionPlan,
1345    /// };
1346    /// use tenferro_tensor::DType;
1347    ///
1348    /// let plan = ElementwiseFusionPlan::new(
1349    ///     DType::F64,
1350    ///     2,
1351    ///     vec![2],
1352    ///     vec![ElementwiseFusionInst::new(ElementwiseFusionOp::Add, vec![0, 1])],
1353    /// );
1354    /// assert_eq!(plan.input_count(), 2);
1355    /// ```
1356    pub fn new(
1357        dtype: crate::DType,
1358        input_count: usize,
1359        outputs: Vec<usize>,
1360        ops: Vec<ElementwiseFusionInst>,
1361    ) -> Self {
1362        Self::with_input_views(
1363            dtype,
1364            vec![ElementwiseFusionInputView::Identity; input_count],
1365            outputs,
1366            ops,
1367        )
1368    }
1369
1370    /// Build a backend elementwise fusion plan with input view metadata.
1371    ///
1372    /// # Examples
1373    ///
1374    /// ```rust
1375    /// use tenferro_tensor::backend::{
1376    ///     ElementwiseFusionInputView, ElementwiseFusionInst, ElementwiseFusionOp,
1377    ///     ElementwiseFusionPlan,
1378    /// };
1379    /// use tenferro_tensor::DType;
1380    ///
1381    /// let plan = ElementwiseFusionPlan::with_input_views(
1382    ///     DType::F64,
1383    ///     vec![ElementwiseFusionInputView::broadcast_in_dim(vec![2, 3], vec![0])],
1384    ///     vec![1],
1385    ///     vec![ElementwiseFusionInst::new(ElementwiseFusionOp::Negate, vec![0])],
1386    /// );
1387    /// assert_eq!(plan.input_count(), 1);
1388    /// ```
1389    pub fn with_input_views(
1390        dtype: crate::DType,
1391        input_views: impl IntoIterator<Item = ElementwiseFusionInputView>,
1392        outputs: Vec<usize>,
1393        ops: Vec<ElementwiseFusionInst>,
1394    ) -> Self {
1395        let input_views = input_views.into_iter().collect::<Vec<_>>();
1396        let input_count = input_views.len();
1397        Self {
1398            dtype,
1399            input_count,
1400            input_views,
1401            outputs,
1402            ops,
1403        }
1404    }
1405
1406    /// Return the scalar dtype expected by this fusion plan.
1407    ///
1408    /// # Examples
1409    ///
1410    /// ```rust
1411    /// use tenferro_tensor::backend::ElementwiseFusionPlan;
1412    /// use tenferro_tensor::DType;
1413    ///
1414    /// let plan = ElementwiseFusionPlan::new(DType::F32, 0, Vec::new(), Vec::new());
1415    /// assert_eq!(plan.dtype(), DType::F32);
1416    /// ```
1417    pub fn dtype(&self) -> crate::DType {
1418        self.dtype
1419    }
1420
1421    /// Return the number of input tensors expected by this plan.
1422    ///
1423    /// # Examples
1424    ///
1425    /// ```rust
1426    /// use tenferro_tensor::backend::ElementwiseFusionPlan;
1427    /// use tenferro_tensor::DType;
1428    ///
1429    /// let plan = ElementwiseFusionPlan::new(DType::F64, 3, Vec::new(), Vec::new());
1430    /// assert_eq!(plan.input_count(), 3);
1431    /// ```
1432    pub fn input_count(&self) -> usize {
1433        self.input_count
1434    }
1435
1436    /// Return metadata views applied to fusion inputs before executing ops.
1437    ///
1438    /// # Examples
1439    ///
1440    /// ```rust
1441    /// use tenferro_tensor::backend::ElementwiseFusionPlan;
1442    /// use tenferro_tensor::DType;
1443    ///
1444    /// let plan = ElementwiseFusionPlan::new(DType::F64, 2, Vec::new(), Vec::new());
1445    /// assert_eq!(plan.input_views().len(), 2);
1446    /// ```
1447    pub fn input_views(&self) -> &[ElementwiseFusionInputView] {
1448        &self.input_views
1449    }
1450
1451    /// Return the value ids selected as fusion outputs.
1452    ///
1453    /// # Examples
1454    ///
1455    /// ```rust
1456    /// use tenferro_tensor::backend::ElementwiseFusionPlan;
1457    /// use tenferro_tensor::DType;
1458    ///
1459    /// let plan = ElementwiseFusionPlan::new(DType::F64, 0, vec![0], Vec::new());
1460    /// assert_eq!(plan.outputs(), &[0]);
1461    /// ```
1462    pub fn outputs(&self) -> &[usize] {
1463        &self.outputs
1464    }
1465
1466    /// Return the fused elementwise instruction sequence.
1467    ///
1468    /// # Examples
1469    ///
1470    /// ```rust
1471    /// use tenferro_tensor::backend::{
1472    ///     ElementwiseFusionInst, ElementwiseFusionOp, ElementwiseFusionPlan,
1473    /// };
1474    /// use tenferro_tensor::DType;
1475    ///
1476    /// let inst = ElementwiseFusionInst::new(ElementwiseFusionOp::Negate, vec![0]);
1477    /// let plan = ElementwiseFusionPlan::new(DType::F64, 1, vec![1], vec![inst]);
1478    /// assert_eq!(plan.ops().len(), 1);
1479    /// ```
1480    pub fn ops(&self) -> &[ElementwiseFusionInst] {
1481        &self.ops
1482    }
1483}
1484
1485impl ElementwiseFusionInputView {
1486    /// Build metadata for a `BroadcastInDim` fusion input view.
1487    ///
1488    /// # Examples
1489    ///
1490    /// ```rust
1491    /// use tenferro_tensor::backend::ElementwiseFusionInputView;
1492    ///
1493    /// let view = ElementwiseFusionInputView::broadcast_in_dim(vec![2, 3], vec![0]);
1494    /// assert!(matches!(view, ElementwiseFusionInputView::BroadcastInDim { .. }));
1495    /// ```
1496    pub fn broadcast_in_dim(
1497        shape: impl IntoIterator<Item = usize>,
1498        dims: impl IntoIterator<Item = usize>,
1499    ) -> Self {
1500        Self::BroadcastInDim {
1501            shape: shape.into_iter().collect(),
1502            dims: dims.into_iter().collect(),
1503        }
1504    }
1505
1506    /// Return true when this fusion input is an identity view.
1507    ///
1508    /// # Examples
1509    ///
1510    /// ```rust
1511    /// use tenferro_tensor::backend::ElementwiseFusionInputView;
1512    ///
1513    /// assert!(ElementwiseFusionInputView::Identity.is_identity());
1514    /// ```
1515    pub fn is_identity(&self) -> bool {
1516        matches!(self, Self::Identity)
1517    }
1518}
1519
1520impl ElementwiseFusionInst {
1521    /// Build a backend elementwise fusion instruction.
1522    ///
1523    /// # Examples
1524    ///
1525    /// ```rust
1526    /// use tenferro_tensor::backend::{ElementwiseFusionInst, ElementwiseFusionOp};
1527    ///
1528    /// let inst = ElementwiseFusionInst::new(ElementwiseFusionOp::Add, vec![0, 1]);
1529    /// assert_eq!(inst.inputs(), &[0, 1]);
1530    /// ```
1531    pub fn new(op: ElementwiseFusionOp, inputs: Vec<usize>) -> Self {
1532        Self { op, inputs }
1533    }
1534
1535    /// Return the elementwise op executed by this instruction.
1536    ///
1537    /// # Examples
1538    ///
1539    /// ```rust
1540    /// use tenferro_tensor::backend::{ElementwiseFusionInst, ElementwiseFusionOp};
1541    ///
1542    /// let inst = ElementwiseFusionInst::new(ElementwiseFusionOp::Negate, vec![0]);
1543    /// assert_eq!(inst.op(), ElementwiseFusionOp::Negate);
1544    /// ```
1545    pub fn op(&self) -> ElementwiseFusionOp {
1546        self.op
1547    }
1548
1549    /// Return this instruction's input value ids.
1550    ///
1551    /// # Examples
1552    ///
1553    /// ```rust
1554    /// use tenferro_tensor::backend::{ElementwiseFusionInst, ElementwiseFusionOp};
1555    ///
1556    /// let inst = ElementwiseFusionInst::new(ElementwiseFusionOp::Multiply, vec![2, 0]);
1557    /// assert_eq!(inst.inputs(), &[2, 0]);
1558    /// ```
1559    pub fn inputs(&self) -> &[usize] {
1560        &self.inputs
1561    }
1562}
1563
1564/// Runtime operation selected by [`TensorElementwise::elementwise_read_into`].
1565#[non_exhaustive]
1566#[derive(Clone, Copy, Debug, Eq, PartialEq)]
1567pub enum ElementwiseReadOp {
1568    /// Binary addition.
1569    Add,
1570    /// Binary subtraction.
1571    Subtract,
1572    /// Binary multiplication.
1573    Multiply,
1574    /// Unary negation.
1575    Negate,
1576    /// Unary conjugation.
1577    Conj,
1578    /// Binary division.
1579    Divide,
1580}
1581
1582impl ElementwiseReadOp {
1583    #[doc(hidden)]
1584    pub fn label(self) -> &'static str {
1585        match self {
1586            Self::Add => "add",
1587            Self::Subtract => "sub",
1588            Self::Multiply => "mul",
1589            Self::Negate => "neg",
1590            Self::Conj => "conj",
1591            Self::Divide => "div",
1592        }
1593    }
1594
1595    #[doc(hidden)]
1596    pub fn arity(self) -> usize {
1597        match self {
1598            Self::Negate | Self::Conj => 1,
1599            Self::Add | Self::Subtract | Self::Multiply | Self::Divide => 2,
1600        }
1601    }
1602}
1603
1604#[derive(Clone, Copy, Debug)]
1605enum StorageIdentity {
1606    Host {
1607        start: usize,
1608        end: usize,
1609    },
1610    Backend {
1611        domain: Option<AllocationDomainId>,
1612        allocation: Option<AllocationId>,
1613        family: &'static str,
1614        object: usize,
1615    },
1616}
1617
1618fn host_storage_identity<T>(data: &[T]) -> StorageIdentity {
1619    let start = data.as_ptr() as usize;
1620    let bytes = std::mem::size_of_val(data);
1621    StorageIdentity::Host {
1622        start,
1623        end: start.saturating_add(bytes),
1624    }
1625}
1626
1627fn backend_storage_identity<T: 'static>(buffer: &dyn crate::BackendStorage<T>) -> StorageIdentity {
1628    StorageIdentity::Backend {
1629        domain: buffer.allocation_domain(),
1630        allocation: buffer.allocation_id(),
1631        family: buffer.backend_family(),
1632        // INVARIANT: every backend buffer is borrowed from the single Box-owned
1633        // root allocation; the data pointer of this trait object is stable for
1634        // that owner and is used only as a fallback when provider identity is
1635        // unavailable.
1636        object: buffer as *const dyn crate::BackendStorage<T> as *const () as usize,
1637    }
1638}
1639
1640fn typed_tensor_storage_identity<T: crate::TensorScalar>(
1641    tensor: &TypedTensor<T>,
1642) -> crate::Result<StorageIdentity> {
1643    if tensor.backend_buffer().is_some() {
1644        let buffer = tensor.backend_buffer().ok_or_else(|| {
1645            crate::Error::runtime_state("typed_tensor_storage_identity", "backend buffer missing")
1646        })?;
1647        Ok(backend_storage_identity(buffer))
1648    } else {
1649        Ok(host_storage_identity(tensor.host_data()?))
1650    }
1651}
1652
1653fn typed_view_storage_identity<T: crate::TensorScalar + 'static>(
1654    view: &TypedTensorView<'_, T>,
1655) -> crate::Result<StorageIdentity> {
1656    match view.backend_buffer() {
1657        Some(buffer) => Ok(backend_storage_identity(buffer)),
1658        None => view.host_storage().map(host_storage_identity),
1659    }
1660}
1661
1662/// The typed tensor behind `value`, or the refusal this module reports for one.
1663///
1664/// Callers reach this from a match on the dtype, so `None` means the tag and the runtime dtype
1665/// disagree rather than a caller mistake.
1666fn identity_operand<T: TensorScalar>(value: &Tensor) -> crate::Result<&TypedTensor<T>> {
1667    value.as_typed::<T>().ok_or_else(|| {
1668        crate::Error::unsupported_dtype(
1669            "storage_identity",
1670            value.dtype(),
1671            "an externally defined payload has no storage identity",
1672        )
1673    })
1674}
1675
1676fn tensor_read_storage_identity(input: &TensorRead<'_>) -> crate::Result<StorageIdentity> {
1677    macro_rules! typed_identity {
1678        ($value:expr) => {
1679            match $value.dtype() {
1680                DType::F32 => typed_tensor_storage_identity(identity_operand::<f32>($value)?),
1681                DType::F64 => typed_tensor_storage_identity(identity_operand::<f64>($value)?),
1682                DType::I32 => typed_tensor_storage_identity(identity_operand::<i32>($value)?),
1683                DType::I64 => typed_tensor_storage_identity(identity_operand::<i64>($value)?),
1684                DType::Bool => typed_tensor_storage_identity(identity_operand::<bool>($value)?),
1685                DType::C32 => typed_tensor_storage_identity(identity_operand::<Complex32>($value)?),
1686                DType::C64 => typed_tensor_storage_identity(identity_operand::<Complex64>($value)?),
1687                // A caller-owned payload has no allocation identity, so it cannot
1688                // take part in an aliasing check.
1689                DType::External(_) => {
1690                    return Err(crate::Error::unsupported_dtype(
1691                        "storage_identity",
1692                        $value.dtype(),
1693                        "an externally defined payload has no storage identity",
1694                    ));
1695                }
1696            }
1697        };
1698    }
1699    macro_rules! view_identity {
1700        ($value:expr) => {
1701            match $value {
1702                TensorView::F32(value) => typed_view_storage_identity(value),
1703                TensorView::F64(value) => typed_view_storage_identity(value),
1704                TensorView::I32(value) => typed_view_storage_identity(value),
1705                TensorView::I64(value) => typed_view_storage_identity(value),
1706                TensorView::Bool(value) => typed_view_storage_identity(value),
1707                TensorView::C32(value) => typed_view_storage_identity(value),
1708                TensorView::C64(value) => typed_view_storage_identity(value),
1709            }
1710        };
1711    }
1712
1713    match input {
1714        TensorRead::Tensor(tensor) => typed_identity!(tensor),
1715        TensorRead::View(view) => view_identity!(view),
1716    }
1717}
1718
1719fn storage_overlaps(lhs: StorageIdentity, rhs: StorageIdentity) -> bool {
1720    match (lhs, rhs) {
1721        (
1722            StorageIdentity::Host {
1723                start: lhs_start,
1724                end: lhs_end,
1725            },
1726            StorageIdentity::Host {
1727                start: rhs_start,
1728                end: rhs_end,
1729            },
1730        ) => lhs_start < rhs_end && rhs_start < lhs_end,
1731        (
1732            StorageIdentity::Backend {
1733                domain: lhs_domain,
1734                allocation: lhs_allocation,
1735                family: lhs_family,
1736                object: lhs_object,
1737            },
1738            StorageIdentity::Backend {
1739                domain: rhs_domain,
1740                allocation: rhs_allocation,
1741                family: rhs_family,
1742                object: rhs_object,
1743            },
1744        ) => {
1745            lhs_object == rhs_object
1746                || matches!(
1747                    (lhs_domain, rhs_domain, lhs_allocation, rhs_allocation),
1748                    (Some(lhs_domain), Some(rhs_domain), Some(lhs), Some(rhs))
1749                        if lhs_domain == rhs_domain && lhs == rhs
1750                )
1751                || matches!(
1752                    (lhs_domain, rhs_domain, lhs_allocation, rhs_allocation),
1753                    (None, None, Some(lhs), Some(rhs)) if lhs_family == rhs_family && lhs == rhs
1754                )
1755        }
1756        _ => false,
1757    }
1758}
1759
1760/// Validate that a caller-owned destination does not overlap any read input.
1761///
1762/// The check is intentionally conservative for host views: two views backed by
1763/// the same host allocation are treated as overlapping because the allocation
1764/// identity is the only stable boundary contract available to erased backend
1765/// code. Backend allocations use their domain/allocation identity when the
1766/// provider exposes it.
1767///
1768/// # Errors
1769///
1770/// Returns `tenferro_tensor_core::ValidationError::InvalidArgument` when the
1771/// destination storage overlaps an input, or `Error::RuntimeState` when
1772/// storage identity cannot be established safely.
1773///
1774/// # Examples
1775///
1776/// ```rust
1777/// use tenferro_tensor::{Tensor, TensorRead, TensorWrite};
1778/// use tenferro_tensor::backend::validate_read_into_destination;
1779///
1780/// let input = Tensor::from_vec_col_major(vec![1], vec![1.0_f64])?;
1781/// let mut output = Tensor::from_vec_col_major(vec![1], vec![0.0_f64])?;
1782/// validate_read_into_destination(
1783///     "example",
1784///     &[TensorRead::from_tensor(&input)],
1785///     &TensorWrite::from_tensor(&mut output),
1786/// )?;
1787/// # Ok::<(), tenferro_tensor::Error>(())
1788/// ```
1789pub fn validate_read_into_destination(
1790    op: &'static str,
1791    inputs: &[TensorRead<'_>],
1792    out: &TensorWrite<'_>,
1793) -> crate::Result<()> {
1794    let output_identity = tensor_read_storage_identity(&out.as_read())?;
1795    for (index, input) in inputs.iter().enumerate() {
1796        if storage_overlaps(tensor_read_storage_identity(input)?, output_identity) {
1797            return Err(Error::invalid_argument(
1798                op,
1799                "out",
1800                format!("destination storage overlaps input {index}"),
1801            ));
1802        }
1803    }
1804    Ok(())
1805}
1806
1807/// Execute a read-into elementwise operation through a backend's allocating operations.
1808///
1809/// Backend implementations use this helper when their device-native operation
1810/// APIs already provide the correct placement and copy semantics. It deliberately
1811/// does not inspect or transfer storage across devices.
1812///
1813/// # Errors
1814///
1815/// Returns [`Error::Validation`] for wrong arity or overlapping output storage,
1816/// and propagates the selected backend operation and copy errors.
1817#[doc(hidden)]
1818pub fn elementwise_read_into_via_allocating_ops<B: TensorElementwise + ?Sized>(
1819    backend: &mut B,
1820    op: ElementwiseReadOp,
1821    inputs: &[TensorRead<'_>],
1822    out: TensorWrite<'_>,
1823) -> crate::Result<()> {
1824    if inputs.len() != op.arity() {
1825        return Err(Error::invalid_argument(
1826            op.label(),
1827            "inputs",
1828            format!("expected {} inputs, got {}", op.arity(), inputs.len()),
1829        ));
1830    }
1831    validate_read_into_destination(op.label(), inputs, &out)?;
1832    let result = match op {
1833        ElementwiseReadOp::Add => backend.add_read(inputs[0].clone(), inputs[1].clone())?,
1834        ElementwiseReadOp::Subtract => backend.sub_read(inputs[0].clone(), inputs[1].clone())?,
1835        ElementwiseReadOp::Multiply => backend.mul_read(inputs[0].clone(), inputs[1].clone())?,
1836        ElementwiseReadOp::Negate => backend.neg_read(inputs[0].clone())?,
1837        ElementwiseReadOp::Conj => backend.conj_read(inputs[0].clone())?,
1838        ElementwiseReadOp::Divide => backend.div_read(inputs[0].clone(), inputs[1].clone())?,
1839    };
1840    backend.copy_read_into(TensorRead::from_tensor(&result), out)
1841}
1842
1843/// Elementwise tensor operations.
1844///
1845/// # Examples
1846///
1847/// ```rust
1848/// use tenferro_tensor::TensorElementwise;
1849///
1850/// fn accepts_elementwise<B: TensorElementwise>(_backend: &mut B) {}
1851/// ```
1852pub trait TensorElementwise: TensorStructural {
1853    /// Execute an elementwise operation into caller-owned storage.
1854    ///
1855    /// Every backend must implement this hook with its explicit execution
1856    /// context, placement checks, and output-storage policy. The hook is the
1857    /// backend boundary for borrowed inputs and caller-owned output; it must not
1858    /// silently transfer data between devices.
1859    ///
1860    /// # Examples
1861    ///
1862    /// ```rust
1863    /// use tenferro_tensor::{ElementwiseReadOp, TensorElementwise, TensorRead, TensorWrite};
1864    ///
1865    /// fn run_into<B: TensorElementwise>(
1866    ///     backend: &mut B,
1867    ///     input: TensorRead<'_>,
1868    ///     out: TensorWrite<'_>,
1869    /// ) -> tenferro_tensor::Result<()> {
1870    ///     backend.elementwise_read_into(ElementwiseReadOp::Conj, &[input], out)
1871    /// }
1872    /// ```
1873    ///
1874    /// # Errors
1875    ///
1876    /// Returns [`crate::Error::Validation`] for invalid arity, metadata, or
1877    /// overlapping storage. Backend-specific placement, unsupported-operation,
1878    /// and execution errors are returned unchanged by the implementation.
1879    fn elementwise_read_into(
1880        &mut self,
1881        op: ElementwiseReadOp,
1882        inputs: &[TensorRead<'_>],
1883        out: TensorWrite<'_>,
1884    ) -> crate::Result<()>;
1885
1886    /// Elementwise addition accepting either owned tensors or borrowed views.
1887    ///
1888    /// Backends that implement this method must not silently move data across
1889    /// devices. A backend that cannot consume views should return an explicit
1890    /// backend error rather than materializing or transferring implicitly.
1891    ///
1892    /// # Examples
1893    ///
1894    /// ```rust
1895    /// use tenferro_tensor::{Tensor, TensorElementwise, TensorRead};
1896    ///
1897    /// fn add_owned<B: TensorElementwise>(
1898    ///     backend: &mut B,
1899    ///     lhs: &Tensor,
1900    ///     rhs: &Tensor,
1901    /// ) -> tenferro_tensor::Result<Tensor> {
1902    ///     backend.add_read(TensorRead::from_tensor(lhs), TensorRead::from_tensor(rhs))
1903    /// }
1904    /// ```
1905    /// # Errors
1906    ///
1907    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
1908    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
1909    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
1910    /// backend execution or storage access cannot provide the requested result.
1911    fn add_read(&mut self, lhs: TensorRead<'_>, rhs: TensorRead<'_>) -> crate::Result<Tensor>;
1912
1913    /// Overwrite caller-provided output with elementwise addition.
1914    ///
1915    /// `_into` methods never accumulate into the previous output value.
1916    ///
1917    /// # Examples
1918    ///
1919    /// ```rust
1920    /// use tenferro_tensor::{Tensor, TensorElementwise, TensorWrite};
1921    ///
1922    /// fn add_into<B: TensorElementwise>(
1923    ///     backend: &mut B,
1924    ///     lhs: &Tensor,
1925    ///     rhs: &Tensor,
1926    ///     mut out: Tensor,
1927    /// ) -> tenferro_tensor::Result<Tensor> {
1928    ///     backend.add_into(lhs, rhs, TensorWrite::from_tensor(&mut out))?;
1929    ///     Ok(out)
1930    /// }
1931    /// ```
1932    /// # Errors
1933    ///
1934    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
1935    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
1936    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
1937    /// backend execution or storage access cannot provide the requested result.
1938    fn add_into(&mut self, lhs: &Tensor, rhs: &Tensor, out: TensorWrite<'_>) -> crate::Result<()> {
1939        self.add_read_into(
1940            TensorRead::from_tensor(lhs),
1941            TensorRead::from_tensor(rhs),
1942            out,
1943        )
1944    }
1945
1946    /// Overwrite caller-provided output with elementwise addition from reads.
1947    ///
1948    /// # Examples
1949    ///
1950    /// ```rust
1951    /// use tenferro_tensor::{TensorElementwise, TensorRead, TensorWrite};
1952    ///
1953    /// fn add_read_into<B: TensorElementwise>(
1954    ///     backend: &mut B,
1955    ///     lhs: TensorRead<'_>,
1956    ///     rhs: TensorRead<'_>,
1957    ///     out: TensorWrite<'_>,
1958    /// ) -> tenferro_tensor::Result<()> {
1959    ///     backend.add_read_into(lhs, rhs, out)
1960    /// }
1961    /// ```
1962    /// # Errors
1963    ///
1964    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
1965    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
1966    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
1967    /// backend execution or storage access cannot provide the requested result.
1968    fn add_read_into(
1969        &mut self,
1970        lhs: TensorRead<'_>,
1971        rhs: TensorRead<'_>,
1972        out: TensorWrite<'_>,
1973    ) -> crate::Result<()> {
1974        self.elementwise_read_into(ElementwiseReadOp::Add, &[lhs, rhs], out)
1975    }
1976
1977    /// Elementwise subtraction accepting either owned tensors or borrowed views.
1978    ///
1979    /// # Examples
1980    ///
1981    /// ```rust
1982    /// use tenferro_tensor::{Tensor, TensorElementwise, TensorRead};
1983    ///
1984    /// fn sub_owned<B: TensorElementwise>(
1985    ///     backend: &mut B,
1986    ///     lhs: &Tensor,
1987    ///     rhs: &Tensor,
1988    /// ) -> tenferro_tensor::Result<Tensor> {
1989    ///     backend.sub_read(TensorRead::from_tensor(lhs), TensorRead::from_tensor(rhs))
1990    /// }
1991    /// ```
1992    /// # Errors
1993    ///
1994    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
1995    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
1996    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
1997    /// backend execution or storage access cannot provide the requested result.
1998    fn sub_read(&mut self, lhs: TensorRead<'_>, rhs: TensorRead<'_>) -> crate::Result<Tensor>;
1999
2000    /// Overwrite caller-provided output with elementwise subtraction.
2001    /// # Errors
2002    ///
2003    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2004    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2005    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2006    /// backend execution or storage access cannot provide the requested result.
2007    fn sub_into(&mut self, lhs: &Tensor, rhs: &Tensor, out: TensorWrite<'_>) -> crate::Result<()> {
2008        self.sub_read_into(
2009            TensorRead::from_tensor(lhs),
2010            TensorRead::from_tensor(rhs),
2011            out,
2012        )
2013    }
2014
2015    /// Overwrite caller-provided output with elementwise subtraction from reads.
2016    /// # Errors
2017    ///
2018    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2019    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2020    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2021    /// backend execution or storage access cannot provide the requested result.
2022    fn sub_read_into(
2023        &mut self,
2024        lhs: TensorRead<'_>,
2025        rhs: TensorRead<'_>,
2026        out: TensorWrite<'_>,
2027    ) -> crate::Result<()> {
2028        self.elementwise_read_into(ElementwiseReadOp::Subtract, &[lhs, rhs], out)
2029    }
2030
2031    /// # Examples
2032    ///
2033    /// ```rust
2034    /// use tenferro_tensor::{TensorElementwise, Tensor, TensorRead};
2035    ///
2036    /// fn mul_read_in_session<B: TensorElementwise>(
2037    ///     backend: &mut B,
2038    ///     lhs: TensorRead<'_>,
2039    ///     rhs: TensorRead<'_>,
2040    /// ) -> tenferro_tensor::Result<Tensor> {
2041    ///     backend.mul_read(lhs, rhs)
2042    /// }
2043    /// ```
2044    ///
2045    /// # Errors
2046    ///
2047    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2048    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2049    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2050    /// backend execution or storage access cannot provide the requested result.
2051    fn mul_read(&mut self, lhs: TensorRead<'_>, rhs: TensorRead<'_>) -> crate::Result<Tensor>;
2052
2053    /// Overwrite caller-provided output with elementwise multiplication.
2054    /// # Errors
2055    ///
2056    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2057    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2058    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2059    /// backend execution or storage access cannot provide the requested result.
2060    fn mul_into(&mut self, lhs: &Tensor, rhs: &Tensor, out: TensorWrite<'_>) -> crate::Result<()> {
2061        self.mul_read_into(
2062            TensorRead::from_tensor(lhs),
2063            TensorRead::from_tensor(rhs),
2064            out,
2065        )
2066    }
2067
2068    /// Overwrite caller-provided output with elementwise multiplication from reads.
2069    /// # Errors
2070    ///
2071    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2072    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2073    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2074    /// backend execution or storage access cannot provide the requested result.
2075    fn mul_read_into(
2076        &mut self,
2077        lhs: TensorRead<'_>,
2078        rhs: TensorRead<'_>,
2079        out: TensorWrite<'_>,
2080    ) -> crate::Result<()> {
2081        self.elementwise_read_into(ElementwiseReadOp::Multiply, &[lhs, rhs], out)
2082    }
2083
2084    /// # Examples
2085    ///
2086    /// ```rust
2087    /// use tenferro_tensor::{TensorElementwise, Tensor, TensorRead};
2088    ///
2089    /// fn neg_read_in_session<B: TensorElementwise>(
2090    ///     backend: &mut B,
2091    ///     input: TensorRead<'_>,
2092    /// ) -> tenferro_tensor::Result<Tensor> {
2093    ///     backend.neg_read(input)
2094    /// }
2095    /// ```
2096    ///
2097    /// # Errors
2098    ///
2099    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2100    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2101    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2102    /// backend execution or storage access cannot provide the requested result.
2103    fn neg_read(&mut self, input: TensorRead<'_>) -> crate::Result<Tensor>;
2104
2105    /// Overwrite caller-provided output with elementwise negation.
2106    /// # Errors
2107    ///
2108    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2109    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2110    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2111    /// backend execution or storage access cannot provide the requested result.
2112    fn neg_into(&mut self, input: &Tensor, out: TensorWrite<'_>) -> crate::Result<()> {
2113        self.neg_read_into(TensorRead::from_tensor(input), out)
2114    }
2115
2116    /// Overwrite caller-provided output with elementwise negation from a read.
2117    /// # Errors
2118    ///
2119    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2120    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2121    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2122    /// backend execution or storage access cannot provide the requested result.
2123    fn neg_read_into(&mut self, input: TensorRead<'_>, out: TensorWrite<'_>) -> crate::Result<()> {
2124        self.elementwise_read_into(ElementwiseReadOp::Negate, &[input], out)
2125    }
2126
2127    /// # Examples
2128    ///
2129    /// ```rust
2130    /// use tenferro_tensor::{TensorElementwise, Tensor, TensorRead};
2131    ///
2132    /// fn conj_read_in_session<B: TensorElementwise>(
2133    ///     backend: &mut B,
2134    ///     input: TensorRead<'_>,
2135    /// ) -> tenferro_tensor::Result<Tensor> {
2136    ///     backend.conj_read(input)
2137    /// }
2138    /// ```
2139    ///
2140    /// # Errors
2141    ///
2142    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2143    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2144    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2145    /// backend execution or storage access cannot provide the requested result.
2146    fn conj_read(&mut self, input: TensorRead<'_>) -> crate::Result<Tensor>;
2147
2148    /// Overwrite caller-provided output with elementwise conjugation.
2149    /// # Errors
2150    ///
2151    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2152    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2153    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2154    /// backend execution or storage access cannot provide the requested result.
2155    fn conj_into(&mut self, input: &Tensor, out: TensorWrite<'_>) -> crate::Result<()> {
2156        self.conj_read_into(TensorRead::from_tensor(input), out)
2157    }
2158
2159    /// Overwrite caller-provided output with elementwise conjugation from a read.
2160    /// # Errors
2161    ///
2162    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2163    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2164    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2165    /// backend execution or storage access cannot provide the requested result.
2166    fn conj_read_into(&mut self, input: TensorRead<'_>, out: TensorWrite<'_>) -> crate::Result<()> {
2167        self.elementwise_read_into(ElementwiseReadOp::Conj, &[input], out)
2168    }
2169
2170    /// # Examples
2171    ///
2172    /// ```rust
2173    /// use tenferro_tensor::{TensorElementwise, Tensor, TensorRead};
2174    ///
2175    /// fn div_read_in_session<B: TensorElementwise>(
2176    ///     backend: &mut B,
2177    ///     lhs: TensorRead<'_>,
2178    ///     rhs: TensorRead<'_>,
2179    /// ) -> tenferro_tensor::Result<Tensor> {
2180    ///     backend.div_read(lhs, rhs)
2181    /// }
2182    /// ```
2183    ///
2184    /// # Errors
2185    ///
2186    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2187    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2188    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2189    /// backend execution or storage access cannot provide the requested result.
2190    fn div_read(&mut self, lhs: TensorRead<'_>, rhs: TensorRead<'_>) -> crate::Result<Tensor>;
2191
2192    /// Overwrite caller-provided output with elementwise division.
2193    /// # Errors
2194    ///
2195    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2196    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2197    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2198    /// backend execution or storage access cannot provide the requested result.
2199    fn div_into(&mut self, lhs: &Tensor, rhs: &Tensor, out: TensorWrite<'_>) -> crate::Result<()> {
2200        self.div_read_into(
2201            TensorRead::from_tensor(lhs),
2202            TensorRead::from_tensor(rhs),
2203            out,
2204        )
2205    }
2206
2207    /// Overwrite caller-provided output with elementwise division from reads.
2208    /// # Errors
2209    ///
2210    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2211    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2212    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2213    /// backend execution or storage access cannot provide the requested result.
2214    fn div_read_into(
2215        &mut self,
2216        lhs: TensorRead<'_>,
2217        rhs: TensorRead<'_>,
2218        out: TensorWrite<'_>,
2219    ) -> crate::Result<()> {
2220        self.elementwise_read_into(ElementwiseReadOp::Divide, &[lhs, rhs], out)
2221    }
2222
2223    /// Elementwise remainder.
2224    ///
2225    /// The default is an explicit unsupported error so backend implementors can
2226    /// opt in without silent fallback.
2227    ///
2228    /// # Examples
2229    ///
2230    /// ```rust
2231    /// use tenferro_tensor::{Tensor, TensorElementwise};
2232    ///
2233    /// fn rem_owned<B: TensorElementwise>(
2234    ///     backend: &mut B,
2235    ///     lhs: &Tensor,
2236    ///     rhs: &Tensor,
2237    /// ) -> tenferro_tensor::Result<Tensor> {
2238    ///     backend.rem(lhs, rhs)
2239    /// }
2240    /// ```
2241    /// # Errors
2242    ///
2243    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2244    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2245    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2246    /// backend execution or storage access cannot provide the requested result.
2247    fn rem(&mut self, lhs: &Tensor, _rhs: &Tensor) -> crate::Result<Tensor> {
2248        Err(crate::Error::unsupported(
2249            "rem",
2250            format!("backend does not implement rem for dtype {:?}", lhs.dtype()),
2251        ))
2252    }
2253
2254    /// Elementwise remainder accepting owned tensors or borrowed views.
2255    ///
2256    /// # Examples
2257    ///
2258    /// ```rust
2259    /// use tenferro_tensor::{Tensor, TensorElementwise, TensorRead};
2260    ///
2261    /// fn rem_read<B: TensorElementwise>(
2262    ///     backend: &mut B,
2263    ///     lhs: &Tensor,
2264    ///     rhs: &Tensor,
2265    /// ) -> tenferro_tensor::Result<Tensor> {
2266    ///     backend.rem_read(TensorRead::from_tensor(lhs), TensorRead::from_tensor(rhs))
2267    /// }
2268    /// ```
2269    /// # Errors
2270    ///
2271    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2272    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2273    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2274    /// backend execution or storage access cannot provide the requested result.
2275    fn rem_read(&mut self, lhs: TensorRead<'_>, rhs: TensorRead<'_>) -> crate::Result<Tensor> {
2276        self.rem(read_tensor("rem", lhs)?, read_tensor("rem", rhs)?)
2277    }
2278
2279    /// # Examples
2280    ///
2281    /// ```rust
2282    /// use tenferro_tensor::{TensorElementwise, Tensor, TensorRead};
2283    ///
2284    /// fn abs_read_in_session<B: TensorElementwise>(
2285    ///     backend: &mut B,
2286    ///     input: TensorRead<'_>,
2287    /// ) -> tenferro_tensor::Result<Tensor> {
2288    ///     backend.abs_read(input)
2289    /// }
2290    /// ```
2291    ///
2292    /// # Errors
2293    ///
2294    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2295    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2296    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2297    /// backend execution or storage access cannot provide the requested result.
2298    fn abs_read(&mut self, input: TensorRead<'_>) -> crate::Result<Tensor>;
2299
2300    /// # Examples
2301    ///
2302    /// ```rust
2303    /// use tenferro_tensor::{TensorElementwise, Tensor, TensorRead};
2304    ///
2305    /// fn sign_read_in_session<B: TensorElementwise>(
2306    ///     backend: &mut B,
2307    ///     input: TensorRead<'_>,
2308    /// ) -> tenferro_tensor::Result<Tensor> {
2309    ///     backend.sign_read(input)
2310    /// }
2311    /// ```
2312    ///
2313    /// # Errors
2314    ///
2315    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2316    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2317    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2318    /// backend execution or storage access cannot provide the requested result.
2319    fn sign_read(&mut self, input: TensorRead<'_>) -> crate::Result<Tensor>;
2320
2321    /// # Examples
2322    ///
2323    /// ```rust
2324    /// use tenferro_tensor::{TensorElementwise, Tensor, TensorRead};
2325    ///
2326    /// fn maximum_read_in_session<B: TensorElementwise>(
2327    ///     backend: &mut B,
2328    ///     lhs: TensorRead<'_>,
2329    ///     rhs: TensorRead<'_>,
2330    /// ) -> tenferro_tensor::Result<Tensor> {
2331    ///     backend.maximum_read(lhs, rhs)
2332    /// }
2333    /// ```
2334    ///
2335    /// # Errors
2336    ///
2337    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2338    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2339    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2340    /// backend execution or storage access cannot provide the requested result.
2341    fn maximum_read(&mut self, lhs: TensorRead<'_>, rhs: TensorRead<'_>) -> crate::Result<Tensor>;
2342
2343    /// # Examples
2344    ///
2345    /// ```rust
2346    /// use tenferro_tensor::{TensorElementwise, Tensor, TensorRead};
2347    ///
2348    /// fn minimum_read_in_session<B: TensorElementwise>(
2349    ///     backend: &mut B,
2350    ///     lhs: TensorRead<'_>,
2351    ///     rhs: TensorRead<'_>,
2352    /// ) -> tenferro_tensor::Result<Tensor> {
2353    ///     backend.minimum_read(lhs, rhs)
2354    /// }
2355    /// ```
2356    ///
2357    /// # Errors
2358    ///
2359    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2360    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2361    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2362    /// backend execution or storage access cannot provide the requested result.
2363    fn minimum_read(&mut self, lhs: TensorRead<'_>, rhs: TensorRead<'_>) -> crate::Result<Tensor>;
2364
2365    /// # Errors
2366    ///
2367    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2368    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2369    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2370    /// backend execution or storage access cannot provide the requested result.
2371    fn compare_read(
2372        &mut self,
2373        lhs: TensorRead<'_>,
2374        rhs: TensorRead<'_>,
2375        dir: &CompareDir,
2376    ) -> crate::Result<Tensor>;
2377
2378    /// # Errors
2379    ///
2380    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2381    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2382    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2383    /// backend execution or storage access cannot provide the requested result.
2384    fn select_read(
2385        &mut self,
2386        pred: TensorRead<'_>,
2387        on_true: TensorRead<'_>,
2388        on_false: TensorRead<'_>,
2389    ) -> crate::Result<Tensor>;
2390
2391    /// # Errors
2392    ///
2393    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2394    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2395    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2396    /// backend execution or storage access cannot provide the requested result.
2397    fn clamp_read(
2398        &mut self,
2399        input: TensorRead<'_>,
2400        lower: TensorRead<'_>,
2401        upper: TensorRead<'_>,
2402    ) -> crate::Result<Tensor>;
2403}
2404
2405/// Analytic unary and binary tensor operations.
2406///
2407/// # Examples
2408///
2409/// ```rust
2410/// use tenferro_tensor::TensorAnalytic;
2411///
2412/// fn accepts_analytic<B: TensorAnalytic>(_backend: &mut B) {}
2413/// ```
2414pub trait TensorAnalytic {
2415    /// # Examples
2416    ///
2417    /// ```rust
2418    /// use tenferro_tensor::{TensorAnalytic, Tensor, TensorRead};
2419    ///
2420    /// fn exp_read_in_session<B: TensorAnalytic>(
2421    ///     backend: &mut B,
2422    ///     input: TensorRead<'_>,
2423    /// ) -> tenferro_tensor::Result<Tensor> {
2424    ///     backend.exp_read(input)
2425    /// }
2426    /// ```
2427    ///
2428    /// # Errors
2429    ///
2430    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2431    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2432    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2433    /// backend execution or storage access cannot provide the requested result.
2434    fn exp_read(&mut self, input: TensorRead<'_>) -> crate::Result<Tensor>;
2435
2436    /// # Examples
2437    ///
2438    /// ```rust
2439    /// use tenferro_tensor::{TensorAnalytic, Tensor, TensorRead};
2440    ///
2441    /// fn log_read_in_session<B: TensorAnalytic>(
2442    ///     backend: &mut B,
2443    ///     input: TensorRead<'_>,
2444    /// ) -> tenferro_tensor::Result<Tensor> {
2445    ///     backend.log_read(input)
2446    /// }
2447    /// ```
2448    ///
2449    /// # Errors
2450    ///
2451    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2452    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2453    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2454    /// backend execution or storage access cannot provide the requested result.
2455    fn log_read(&mut self, input: TensorRead<'_>) -> crate::Result<Tensor>;
2456
2457    /// # Examples
2458    ///
2459    /// ```rust
2460    /// use tenferro_tensor::{TensorAnalytic, Tensor, TensorRead};
2461    ///
2462    /// fn sin_read_in_session<B: TensorAnalytic>(
2463    ///     backend: &mut B,
2464    ///     input: TensorRead<'_>,
2465    /// ) -> tenferro_tensor::Result<Tensor> {
2466    ///     backend.sin_read(input)
2467    /// }
2468    /// ```
2469    ///
2470    /// # Errors
2471    ///
2472    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2473    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2474    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2475    /// backend execution or storage access cannot provide the requested result.
2476    fn sin_read(&mut self, input: TensorRead<'_>) -> crate::Result<Tensor>;
2477
2478    /// # Examples
2479    ///
2480    /// ```rust
2481    /// use tenferro_tensor::{TensorAnalytic, Tensor, TensorRead};
2482    ///
2483    /// fn cos_read_in_session<B: TensorAnalytic>(
2484    ///     backend: &mut B,
2485    ///     input: TensorRead<'_>,
2486    /// ) -> tenferro_tensor::Result<Tensor> {
2487    ///     backend.cos_read(input)
2488    /// }
2489    /// ```
2490    ///
2491    /// # Errors
2492    ///
2493    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2494    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2495    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2496    /// backend execution or storage access cannot provide the requested result.
2497    fn cos_read(&mut self, input: TensorRead<'_>) -> crate::Result<Tensor>;
2498
2499    /// # Examples
2500    ///
2501    /// ```rust
2502    /// use tenferro_tensor::{TensorAnalytic, Tensor, TensorRead};
2503    ///
2504    /// fn tanh_read_in_session<B: TensorAnalytic>(
2505    ///     backend: &mut B,
2506    ///     input: TensorRead<'_>,
2507    /// ) -> tenferro_tensor::Result<Tensor> {
2508    ///     backend.tanh_read(input)
2509    /// }
2510    /// ```
2511    ///
2512    /// # Errors
2513    ///
2514    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2515    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2516    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2517    /// backend execution or storage access cannot provide the requested result.
2518    fn tanh_read(&mut self, input: TensorRead<'_>) -> crate::Result<Tensor>;
2519
2520    /// # Examples
2521    ///
2522    /// ```rust
2523    /// use tenferro_tensor::{TensorAnalytic, Tensor, TensorRead};
2524    ///
2525    /// fn sqrt_read_in_session<B: TensorAnalytic>(
2526    ///     backend: &mut B,
2527    ///     input: TensorRead<'_>,
2528    /// ) -> tenferro_tensor::Result<Tensor> {
2529    ///     backend.sqrt_read(input)
2530    /// }
2531    /// ```
2532    ///
2533    /// # Errors
2534    ///
2535    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2536    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2537    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2538    /// backend execution or storage access cannot provide the requested result.
2539    fn sqrt_read(&mut self, input: TensorRead<'_>) -> crate::Result<Tensor>;
2540
2541    /// # Examples
2542    ///
2543    /// ```rust
2544    /// use tenferro_tensor::{TensorAnalytic, Tensor, TensorRead};
2545    ///
2546    /// fn rsqrt_read_in_session<B: TensorAnalytic>(
2547    ///     backend: &mut B,
2548    ///     input: TensorRead<'_>,
2549    /// ) -> tenferro_tensor::Result<Tensor> {
2550    ///     backend.rsqrt_read(input)
2551    /// }
2552    /// ```
2553    ///
2554    /// # Errors
2555    ///
2556    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2557    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2558    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2559    /// backend execution or storage access cannot provide the requested result.
2560    fn rsqrt_read(&mut self, input: TensorRead<'_>) -> crate::Result<Tensor>;
2561
2562    /// # Examples
2563    ///
2564    /// ```rust
2565    /// use tenferro_tensor::{TensorAnalytic, Tensor, TensorRead};
2566    ///
2567    /// fn pow_read_in_session<B: TensorAnalytic>(
2568    ///     backend: &mut B,
2569    ///     lhs: TensorRead<'_>,
2570    ///     rhs: TensorRead<'_>,
2571    /// ) -> tenferro_tensor::Result<Tensor> {
2572    ///     backend.pow_read(lhs, rhs)
2573    /// }
2574    /// ```
2575    ///
2576    /// # Errors
2577    ///
2578    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2579    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2580    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2581    /// backend execution or storage access cannot provide the requested result.
2582    fn pow_read(&mut self, lhs: TensorRead<'_>, rhs: TensorRead<'_>) -> crate::Result<Tensor>;
2583
2584    /// # Examples
2585    ///
2586    /// ```rust
2587    /// use tenferro_tensor::{TensorAnalytic, Tensor, TensorRead};
2588    ///
2589    /// fn expm1_read_in_session<B: TensorAnalytic>(
2590    ///     backend: &mut B,
2591    ///     input: TensorRead<'_>,
2592    /// ) -> tenferro_tensor::Result<Tensor> {
2593    ///     backend.expm1_read(input)
2594    /// }
2595    /// ```
2596    ///
2597    /// # Errors
2598    ///
2599    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2600    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2601    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2602    /// backend execution or storage access cannot provide the requested result.
2603    fn expm1_read(&mut self, input: TensorRead<'_>) -> crate::Result<Tensor>;
2604
2605    /// # Examples
2606    ///
2607    /// ```rust
2608    /// use tenferro_tensor::{TensorAnalytic, Tensor, TensorRead};
2609    ///
2610    /// fn log1p_read_in_session<B: TensorAnalytic>(
2611    ///     backend: &mut B,
2612    ///     input: TensorRead<'_>,
2613    /// ) -> tenferro_tensor::Result<Tensor> {
2614    ///     backend.log1p_read(input)
2615    /// }
2616    /// ```
2617    ///
2618    /// # Errors
2619    ///
2620    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2621    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2622    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2623    /// backend execution or storage access cannot provide the requested result.
2624    fn log1p_read(&mut self, input: TensorRead<'_>) -> crate::Result<Tensor>;
2625
2626    /// Elementwise error function `erf(x) = 2/sqrt(pi) * integral_0^x exp(-t^2) dt`.
2627    ///
2628    /// Defined for real `F32` and `F64` input only; the result has the input
2629    /// dtype and shape. `erf(+-0) = +-0`, `erf(+-inf) = +-1`, and `NaN` stays
2630    /// `NaN`.
2631    ///
2632    /// # Examples
2633    ///
2634    /// ```rust
2635    /// use tenferro_tensor::{TensorAnalytic, Tensor, TensorRead};
2636    ///
2637    /// fn erf_read_in_session<B: TensorAnalytic>(
2638    ///     backend: &mut B,
2639    ///     input: TensorRead<'_>,
2640    /// ) -> tenferro_tensor::Result<Tensor> {
2641    ///     backend.erf_read(input)
2642    /// }
2643    /// ```
2644    ///
2645    /// # Errors
2646    ///
2647    /// Returns [`crate::Error::UnsupportedDType`] for a complex, integer, or `Bool`
2648    /// input dtype, [`crate::Error::Validation`] with a typed `ValidationError`
2649    /// source for invalid shapes or output metadata, and
2650    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2651    /// backend execution or storage access cannot provide the requested result.
2652    fn erf_read(&mut self, input: TensorRead<'_>) -> crate::Result<Tensor>;
2653}
2654
2655/// Shape, layout, and dtype transformation operations.
2656///
2657/// # Examples
2658///
2659/// ```rust
2660/// use tenferro_tensor::TensorStructural;
2661///
2662/// fn accepts_structural<B: TensorStructural>(_backend: &mut B) {}
2663/// ```
2664pub trait TensorStructural {
2665    /// Materialize an owned tensor or borrowed view into fresh compact storage.
2666    ///
2667    /// The result has the input's shape and dtype, uses compact column-major
2668    /// layout, and remains in the input's placement. This operation is a
2669    /// same-placement canonicalization boundary, never an implicit host/device
2670    /// transfer. The conservative default accepts host-owned tensors and host
2671    /// views; a view is gathered over its layout, so transposed, sliced, and
2672    /// negative-stride views materialize. It rejects backend buffers and device
2673    /// placement because only an owning backend can materialize those safely.
2674    ///
2675    /// Backend overrides may also accept backend-owned strided views. CUDA
2676    /// accepts numeric and complex views on its active device, including
2677    /// arbitrary valid strides; for `Bool` it copies owned tensors but
2678    /// reports an explicit unsupported-dtype error for strided views.
2679    ///
2680    /// # Examples
2681    ///
2682    /// ```rust
2683    /// use tenferro_tensor::{DType, Tensor, TensorRead, TensorStructural};
2684    ///
2685    /// struct HostDefaults;
2686    /// impl TensorStructural for HostDefaults {
2687    ///     fn transpose_read(&mut self, _: TensorRead<'_>, _: &[usize]) -> tenferro_tensor::Result<Tensor> { unimplemented!() }
2688    ///     fn reshape_read(&mut self, _: TensorRead<'_>, _: &[usize]) -> tenferro_tensor::Result<Tensor> { unimplemented!() }
2689    ///     fn broadcast_in_dim_read(&mut self, _: TensorRead<'_>, _: &[usize], _: &[usize]) -> tenferro_tensor::Result<Tensor> { unimplemented!() }
2690    ///     fn cast(&mut self, _: &Tensor, _: DType) -> tenferro_tensor::Result<Tensor> { unimplemented!() }
2691    ///     fn extract_diagonal(&mut self, _: &Tensor, _: usize, _: usize) -> tenferro_tensor::Result<Tensor> { unimplemented!() }
2692    ///     fn embed_diagonal(&mut self, _: &Tensor, _: usize, _: usize) -> tenferro_tensor::Result<Tensor> { unimplemented!() }
2693    ///     fn tril(&mut self, _: &Tensor, _: i64) -> tenferro_tensor::Result<Tensor> { unimplemented!() }
2694    ///     fn triu(&mut self, _: &Tensor, _: i64) -> tenferro_tensor::Result<Tensor> { unimplemented!() }
2695    /// }
2696    ///
2697    /// let input = Tensor::from_vec_col_major(vec![2], vec![1_i32, 2])?;
2698    /// let mut backend = HostDefaults;
2699    /// let structural: &mut dyn TensorStructural = &mut backend;
2700    /// let output = structural.to_contiguous_read(TensorRead::from_tensor(&input))?;
2701    /// assert_eq!(output.shape(), &[2]);
2702    /// assert_eq!(output.as_slice::<i32>()?, &[1, 2]);
2703    /// # Ok::<(), tenferro_tensor::Error>(())
2704    /// ```
2705    /// # Errors
2706    ///
2707    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2708    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2709    /// [`crate::Error::RuntimeState`] for backend-owned or device-placed input,
2710    /// which only the owning backend can materialize, or
2711    /// [`crate::Error::BackendFailure`] / [`crate::Error::BackendSource`] when
2712    /// backend execution or storage access cannot provide the requested result.
2713    fn to_contiguous_read(&mut self, input: TensorRead<'_>) -> crate::Result<Tensor> {
2714        match input {
2715            TensorRead::Tensor(input) => {
2716                if input.is_backend_buffer()
2717                    || !matches!(
2718                        input.placement().memory_kind,
2719                        crate::MemoryKind::PinnedHost | crate::MemoryKind::UnpinnedHost
2720                    )
2721                {
2722                    return Err(crate::Error::runtime_state(
2723                        "to_contiguous_read",
2724                        "default materialization accepts only host-owned tensors; use the storage's owning backend",
2725                    ));
2726                }
2727                input.duplicate()
2728            }
2729            TensorRead::View(view) => {
2730                if view.backend_family().is_some()
2731                    || !matches!(
2732                        view.placement().memory_kind,
2733                        crate::MemoryKind::PinnedHost | crate::MemoryKind::UnpinnedHost
2734                    )
2735                {
2736                    return Err(crate::Error::runtime_state(
2737                        "to_contiguous_read",
2738                        "default materialization accepts only host-owned tensors; use the storage's owning backend",
2739                    ));
2740                }
2741                default_materialize_host_view(view)
2742            }
2743        }
2744    }
2745
2746    /// Overwrite caller-provided storage from a readable tensor or view.
2747    ///
2748    /// Source and destination must have identical dtype and shape and belong to
2749    /// the executing backend's placement. The destination is not resized, and
2750    /// every logical destination element is overwritten without reading its old
2751    /// value. Source and destination allocations must not alias. Implementations
2752    /// must not materialize through host memory or perform an implicit transfer.
2753    ///
2754    /// CPU accepts arbitrary valid source and destination strides and performs
2755    /// no tensor allocation. CUDA currently accepts only a compact column-major
2756    /// source with offset zero covering its full allocation; CUDA destinations
2757    /// may be arbitrary valid non-overlapping views. CUDA rejects aliased
2758    /// allocations and currently reports an explicit unsupported-dtype error
2759    /// for `Bool`. The conservative default is explicitly unsupported.
2760    ///
2761    /// # Examples
2762    ///
2763    /// ```rust
2764    /// use tenferro_tensor::{DType, Tensor, TensorRead, TensorStructural, TensorWrite};
2765    ///
2766    /// struct ConservativeDefaults;
2767    /// impl TensorStructural for ConservativeDefaults {
2768    ///     fn transpose_read(&mut self, _: TensorRead<'_>, _: &[usize]) -> tenferro_tensor::Result<Tensor> { unimplemented!() }
2769    ///     fn reshape_read(&mut self, _: TensorRead<'_>, _: &[usize]) -> tenferro_tensor::Result<Tensor> { unimplemented!() }
2770    ///     fn broadcast_in_dim_read(&mut self, _: TensorRead<'_>, _: &[usize], _: &[usize]) -> tenferro_tensor::Result<Tensor> { unimplemented!() }
2771    ///     fn cast(&mut self, _: &Tensor, _: DType) -> tenferro_tensor::Result<Tensor> { unimplemented!() }
2772    ///     fn extract_diagonal(&mut self, _: &Tensor, _: usize, _: usize) -> tenferro_tensor::Result<Tensor> { unimplemented!() }
2773    ///     fn embed_diagonal(&mut self, _: &Tensor, _: usize, _: usize) -> tenferro_tensor::Result<Tensor> { unimplemented!() }
2774    ///     fn tril(&mut self, _: &Tensor, _: i64) -> tenferro_tensor::Result<Tensor> { unimplemented!() }
2775    ///     fn triu(&mut self, _: &Tensor, _: i64) -> tenferro_tensor::Result<Tensor> { unimplemented!() }
2776    /// }
2777    ///
2778    /// let src = Tensor::from_vec_col_major(vec![2], vec![1_i32, 2])?;
2779    /// let mut dst = Tensor::from_vec_col_major(vec![2], vec![0_i32, 0])?;
2780    /// let mut backend = ConservativeDefaults;
2781    /// let structural: &mut dyn TensorStructural = &mut backend;
2782    /// let error = structural.copy_read_into(
2783    ///     TensorRead::from_tensor(&src),
2784    ///     TensorWrite::from_tensor(&mut dst),
2785    /// ).unwrap_err();
2786    /// assert!(error.to_string().contains("unsupported"));
2787    /// assert_eq!(dst.as_slice::<i32>()?, &[0, 0]);
2788    /// # Ok::<(), tenferro_tensor::Error>(())
2789    /// ```
2790    /// # Errors
2791    ///
2792    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2793    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2794    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2795    /// backend execution or storage access cannot provide the requested result.
2796    fn copy_read_into(&mut self, _src: TensorRead<'_>, _dst: TensorWrite<'_>) -> crate::Result<()> {
2797        Err(crate::Error::unsupported(
2798            "copy_read_into",
2799            "backend-owned runtime copy is unsupported by this backend",
2800        ))
2801    }
2802
2803    /// # Examples
2804    ///
2805    /// ```rust
2806    /// use tenferro_tensor::{TensorStructural, Tensor, TensorRead};
2807    ///
2808    /// fn transpose_read_in_session<B: TensorStructural>(
2809    ///     backend: &mut B,
2810    ///     input: TensorRead<'_>,
2811    ///     perm: &[usize],
2812    /// ) -> tenferro_tensor::Result<Tensor> {
2813    ///     backend.transpose_read(input, perm)
2814    /// }
2815    /// ```
2816    ///
2817    /// # Errors
2818    ///
2819    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2820    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2821    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2822    /// backend execution or storage access cannot provide the requested result.
2823    fn transpose_read(&mut self, input: TensorRead<'_>, perm: &[usize]) -> crate::Result<Tensor>;
2824
2825    /// # Examples
2826    ///
2827    /// ```rust
2828    /// use tenferro_tensor::{TensorStructural, Tensor, TensorRead};
2829    ///
2830    /// fn reshape_read_in_session<B: TensorStructural>(
2831    ///     backend: &mut B,
2832    ///     input: TensorRead<'_>,
2833    ///     shape: &[usize],
2834    /// ) -> tenferro_tensor::Result<Tensor> {
2835    ///     backend.reshape_read(input, shape)
2836    /// }
2837    /// ```
2838    ///
2839    /// # Errors
2840    ///
2841    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2842    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2843    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2844    /// backend execution or storage access cannot provide the requested result.
2845    fn reshape_read(&mut self, input: TensorRead<'_>, shape: &[usize]) -> crate::Result<Tensor>;
2846
2847    /// # Errors
2848    ///
2849    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2850    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2851    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2852    /// backend execution or storage access cannot provide the requested result.
2853    fn broadcast_in_dim_read(
2854        &mut self,
2855        input: TensorRead<'_>,
2856        shape: &[usize],
2857        dims: &[usize],
2858    ) -> crate::Result<Tensor>;
2859
2860    /// Cast a tensor to another dtype using explicit dtype projection.
2861    ///
2862    /// Backends may truncate, narrow precision, project complex values, or use
2863    /// boolean truthiness according to their documented cast support.
2864    ///
2865    /// # Examples
2866    ///
2867    /// ```rust
2868    /// use tenferro_tensor::{DType, Tensor, TensorStructural};
2869    ///
2870    /// fn cast_to_i32<B: TensorStructural>(
2871    ///     backend: &mut B,
2872    ///     input: &Tensor,
2873    /// ) -> tenferro_tensor::Result<Tensor> {
2874    ///     backend.cast(input, DType::I32)
2875    /// }
2876    /// ```
2877    /// # Errors
2878    ///
2879    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2880    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2881    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2882    /// backend execution or storage access cannot provide the requested result.
2883    fn cast(&mut self, input: &Tensor, to: crate::DType) -> crate::Result<Tensor>;
2884
2885    /// Convert a tensor to another dtype using checked dtype conversion.
2886    ///
2887    /// `convert` accepts only conversions allowed by tenferro's dtype-promotion
2888    /// lattice. Use [`TensorStructural::cast`] for explicit lossy projection.
2889    ///
2890    /// # Examples
2891    ///
2892    /// ```rust
2893    /// use tenferro_tensor::{DType, Tensor, TensorStructural};
2894    ///
2895    /// fn convert_to_f64<B: TensorStructural>(
2896    ///     backend: &mut B,
2897    ///     input: &Tensor,
2898    /// ) -> tenferro_tensor::Result<Tensor> {
2899    ///     backend.convert(input, DType::F64)
2900    /// }
2901    /// ```
2902    /// # Errors
2903    ///
2904    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2905    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2906    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2907    /// backend execution or storage access cannot provide the requested result.
2908    fn convert(&mut self, input: &Tensor, to: crate::DType) -> crate::Result<Tensor> {
2909        validate_convert_dtype("convert", input.dtype(), to)?;
2910        self.cast(input, to)
2911    }
2912
2913    /// # Errors
2914    ///
2915    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2916    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2917    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2918    /// backend execution or storage access cannot provide the requested result.
2919    fn extract_diagonal(
2920        &mut self,
2921        input: &Tensor,
2922        axis_a: usize,
2923        axis_b: usize,
2924    ) -> crate::Result<Tensor>;
2925    /// # Errors
2926    ///
2927    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2928    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2929    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2930    /// backend execution or storage access cannot provide the requested result.
2931    fn embed_diagonal(
2932        &mut self,
2933        input: &Tensor,
2934        axis_a: usize,
2935        axis_b: usize,
2936    ) -> crate::Result<Tensor>;
2937    /// # Errors
2938    ///
2939    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2940    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2941    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2942    /// backend execution or storage access cannot provide the requested result.
2943    fn tril(&mut self, input: &Tensor, k: i64) -> crate::Result<Tensor>;
2944    /// # Errors
2945    ///
2946    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2947    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2948    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2949    /// backend execution or storage access cannot provide the requested result.
2950    fn triu(&mut self, input: &Tensor, k: i64) -> crate::Result<Tensor>;
2951}
2952
2953/// Reduction operations.
2954///
2955/// Reducing over an axis whose extent is zero returns an error for every
2956/// reduction operation. Passing an empty `axes` slice is a no-op for the public
2957/// reductions and returns the input values unchanged. Internal mapped
2958/// reductions document their own empty-axis semantics.
2959///
2960/// # Examples
2961///
2962/// ```rust
2963/// use tenferro_tensor::TensorReduction;
2964///
2965/// fn accepts_reduction<B: TensorReduction>(_backend: &mut B) {}
2966/// ```
2967pub trait TensorReduction {
2968    /// Sum elements across axes from an owned tensor or borrowed view.
2969    ///
2970    /// # Examples
2971    ///
2972    /// ```rust
2973    /// use tenferro_tensor::{Tensor, TensorRead, TensorReduction};
2974    ///
2975    /// fn sum_owned<B: TensorReduction>(
2976    ///     backend: &mut B,
2977    ///     input: &Tensor,
2978    /// ) -> tenferro_tensor::Result<Tensor> {
2979    ///     backend.reduce_sum_read(TensorRead::from_tensor(input), &[0])
2980    /// }
2981    /// ```
2982    /// # Errors
2983    ///
2984    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
2985    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
2986    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
2987    /// backend execution or storage access cannot provide the requested result.
2988    fn reduce_sum_read(&mut self, input: TensorRead<'_>, axes: &[usize]) -> crate::Result<Tensor>;
2989
2990    /// Sum elementwise squares across axes.
2991    ///
2992    /// This execution hook is used by composite operations that avoid a
2993    /// materialized square. Empty axes produce an elementwise square. Backends
2994    /// that support this optimized path must override the hook directly.
2995    ///
2996    /// # Errors
2997    ///
2998    /// Returns the typed validation, unsupported, runtime-state, or backend
2999    /// error produced by multiplication or reduction.
3000    #[doc(hidden)]
3001    fn reduce_sum_squares_read(
3002        &mut self,
3003        _input: TensorRead<'_>,
3004        _axes: &[usize],
3005    ) -> crate::Result<Tensor> {
3006        Err(crate::Error::unsupported(
3007            "reduce_sum_squares",
3008            "backend does not implement fused sum-of-squares reduction",
3009        ))
3010    }
3011
3012    /// Multiply elements across axes from an owned tensor or borrowed view.
3013    ///
3014    /// # Examples
3015    ///
3016    /// ```rust
3017    /// use tenferro_tensor::{Tensor, TensorRead, TensorReduction};
3018    ///
3019    /// fn prod_owned<B: TensorReduction>(
3020    ///     backend: &mut B,
3021    ///     input: &Tensor,
3022    /// ) -> tenferro_tensor::Result<Tensor> {
3023    ///     backend.reduce_prod_read(TensorRead::from_tensor(input), &[0])
3024    /// }
3025    /// ```
3026    /// # Errors
3027    ///
3028    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
3029    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
3030    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
3031    /// backend execution or storage access cannot provide the requested result.
3032    fn reduce_prod_read(&mut self, input: TensorRead<'_>, axes: &[usize]) -> crate::Result<Tensor>;
3033
3034    /// Take maximum values across axes from an owned tensor or borrowed view.
3035    ///
3036    /// # Examples
3037    ///
3038    /// ```rust
3039    /// use tenferro_tensor::{Tensor, TensorRead, TensorReduction};
3040    ///
3041    /// fn max_owned<B: TensorReduction>(
3042    ///     backend: &mut B,
3043    ///     input: &Tensor,
3044    /// ) -> tenferro_tensor::Result<Tensor> {
3045    ///     backend.reduce_max_read(TensorRead::from_tensor(input), &[0])
3046    /// }
3047    /// ```
3048    /// # Errors
3049    ///
3050    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
3051    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
3052    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
3053    /// backend execution or storage access cannot provide the requested result.
3054    fn reduce_max_read(&mut self, input: TensorRead<'_>, axes: &[usize]) -> crate::Result<Tensor>;
3055
3056    /// Take minimum values across axes from an owned tensor or borrowed view.
3057    ///
3058    /// # Examples
3059    ///
3060    /// ```rust
3061    /// use tenferro_tensor::{Tensor, TensorRead, TensorReduction};
3062    ///
3063    /// fn min_owned<B: TensorReduction>(
3064    ///     backend: &mut B,
3065    ///     input: &Tensor,
3066    /// ) -> tenferro_tensor::Result<Tensor> {
3067    ///     backend.reduce_min_read(TensorRead::from_tensor(input), &[0])
3068    /// }
3069    /// ```
3070    /// # Errors
3071    ///
3072    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
3073    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
3074    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
3075    /// backend execution or storage access cannot provide the requested result.
3076    fn reduce_min_read(&mut self, input: TensorRead<'_>, axes: &[usize]) -> crate::Result<Tensor>;
3077}
3078
3079/// Dot-general operations.
3080///
3081/// # Examples
3082///
3083/// ```rust
3084/// use tenferro_tensor::TensorDot;
3085///
3086/// fn accepts_dot<B: TensorDot>(_backend: &mut B) {}
3087/// ```
3088pub trait TensorDot: TensorElementwise {
3089    #[doc(hidden)]
3090    fn dot_general_read(
3091        &mut self,
3092        lhs: TensorRead<'_>,
3093        rhs: TensorRead<'_>,
3094        config: &DotGeneralConfig,
3095    ) -> crate::Result<Tensor>;
3096
3097    /// Overwrite caller-provided output with dot-general from read inputs.
3098    ///
3099    /// This is the dot/GEMM spelling of `_into`: the previous output value is
3100    /// not read. Use [`TensorDot::dot_general_read_into_accum`] for explicit
3101    /// read-modify-write accumulation.
3102    ///
3103    /// # Examples
3104    ///
3105    /// ```rust
3106    /// use tenferro_tensor::{DotGeneralConfig, TensorDot, TensorRead, TensorWrite};
3107    ///
3108    /// fn dot_into<B: TensorDot>(
3109    ///     backend: &mut B,
3110    ///     lhs: TensorRead<'_>,
3111    ///     rhs: TensorRead<'_>,
3112    ///     config: &DotGeneralConfig,
3113    ///     out: TensorWrite<'_>,
3114    /// ) -> tenferro_tensor::Result<()> {
3115    ///     backend.dot_general_read_into(lhs, rhs, config, out)
3116    /// }
3117    /// ```
3118    /// # Errors
3119    ///
3120    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
3121    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
3122    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
3123    /// backend execution or storage access cannot provide the requested result.
3124    fn dot_general_read_into(
3125        &mut self,
3126        lhs: TensorRead<'_>,
3127        rhs: TensorRead<'_>,
3128        config: &DotGeneralConfig,
3129        out: TensorWrite<'_>,
3130    ) -> crate::Result<()> {
3131        let accumulation = DotGeneralAccumulation::overwrite(lhs.dtype())?;
3132        self.dot_general_read_into_accum(lhs, rhs, config, accumulation, out)
3133    }
3134
3135    #[doc(hidden)]
3136    fn dot_general_with_conj(
3137        &mut self,
3138        lhs: &Tensor,
3139        rhs: &Tensor,
3140        config: &DotGeneralConfig,
3141        lhs_conj: bool,
3142        rhs_conj: bool,
3143    ) -> crate::Result<Tensor> {
3144        if !lhs_conj && !rhs_conj {
3145            return self.dot_general_read(
3146                TensorRead::from_tensor(lhs),
3147                TensorRead::from_tensor(rhs),
3148                config,
3149            );
3150        }
3151
3152        let lhs_tmp;
3153        let lhs_ref = if lhs_conj {
3154            lhs_tmp = self.conj_read(TensorRead::from_tensor(lhs))?;
3155            &lhs_tmp
3156        } else {
3157            lhs
3158        };
3159        let rhs_tmp;
3160        let rhs_ref = if rhs_conj {
3161            rhs_tmp = self.conj_read(TensorRead::from_tensor(rhs))?;
3162            &rhs_tmp
3163        } else {
3164            rhs
3165        };
3166        self.dot_general_read(
3167            TensorRead::from_tensor(lhs_ref),
3168            TensorRead::from_tensor(rhs_ref),
3169            config,
3170        )
3171    }
3172
3173    #[allow(clippy::too_many_arguments)]
3174    #[doc(hidden)]
3175    fn dot_general_with_conj_read(
3176        &mut self,
3177        lhs: TensorRead<'_>,
3178        rhs: TensorRead<'_>,
3179        config: &DotGeneralConfig,
3180        lhs_conj: bool,
3181        rhs_conj: bool,
3182    ) -> crate::Result<Tensor> {
3183        if !lhs_conj && !rhs_conj {
3184            return self.dot_general_read(lhs, rhs, config);
3185        }
3186
3187        let lhs_tmp;
3188        let lhs_ref = if let Some(tensor) = lhs.as_tensor() {
3189            tensor
3190        } else {
3191            lhs_tmp = self.to_contiguous_read(lhs)?;
3192            &lhs_tmp
3193        };
3194        let rhs_tmp;
3195        let rhs_ref = if let Some(tensor) = rhs.as_tensor() {
3196            tensor
3197        } else {
3198            rhs_tmp = self.to_contiguous_read(rhs)?;
3199            &rhs_tmp
3200        };
3201        self.dot_general_with_conj(lhs_ref, rhs_ref, config, lhs_conj, rhs_conj)
3202    }
3203
3204    /// Apply scaled dot-general accumulation into caller-provided output.
3205    ///
3206    /// This is explicitly read-modify-write when `accumulation.beta` is nonzero:
3207    /// `out = alpha * dot_general(lhs, rhs) + beta * out`.
3208    ///
3209    /// # Examples
3210    ///
3211    /// ```rust
3212    /// use tenferro_tensor::{
3213    ///     DotGeneralAccumulation, DotGeneralConfig, TensorDot, TensorRead, TensorWrite,
3214    /// };
3215    ///
3216    /// fn dot_add_to<B: TensorDot>(
3217    ///     backend: &mut B,
3218    ///     lhs: TensorRead<'_>,
3219    ///     rhs: TensorRead<'_>,
3220    ///     config: &DotGeneralConfig,
3221    ///     out: TensorWrite<'_>,
3222    /// ) -> tenferro_tensor::Result<()> {
3223    ///     let accumulation = DotGeneralAccumulation::add_to(lhs.dtype())?;
3224    ///     backend.dot_general_read_into_accum(lhs, rhs, config, accumulation, out)
3225    /// }
3226    /// ```
3227    /// # Errors
3228    ///
3229    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
3230    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
3231    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
3232    /// backend execution or storage access cannot provide the requested result.
3233    fn dot_general_read_into_accum(
3234        &mut self,
3235        lhs: TensorRead<'_>,
3236        rhs: TensorRead<'_>,
3237        config: &DotGeneralConfig,
3238        accumulation: DotGeneralAccumulation,
3239        out: TensorWrite<'_>,
3240    ) -> crate::Result<()> {
3241        dot_general_accum_via_temp(self, lhs, rhs, config, accumulation, out)
3242    }
3243}
3244
3245/// Session-scoped cached dot-general operations.
3246///
3247/// # Examples
3248///
3249/// ```rust
3250/// use tenferro_tensor::BackendSession;
3251///
3252/// fn accepts_session_dot<S: BackendSession + ?Sized>(_session: &mut S) {}
3253/// ```
3254pub trait SessionCachedDot: TensorDot {
3255    #[doc(hidden)]
3256    fn dot_general_cached(
3257        &mut self,
3258        _cache_slot: Option<usize>,
3259        lhs: &Tensor,
3260        rhs: &Tensor,
3261        config: &DotGeneralConfig,
3262    ) -> crate::Result<Tensor> {
3263        self.dot_general_read(
3264            TensorRead::from_tensor(lhs),
3265            TensorRead::from_tensor(rhs),
3266            config,
3267        )
3268    }
3269
3270    #[doc(hidden)]
3271    fn dot_general_read_cached(
3272        &mut self,
3273        cache_slot: Option<usize>,
3274        lhs: TensorRead<'_>,
3275        rhs: TensorRead<'_>,
3276        config: &DotGeneralConfig,
3277    ) -> crate::Result<Tensor> {
3278        match (lhs.as_tensor(), rhs.as_tensor()) {
3279            (Some(lhs), Some(rhs)) => self.dot_general_cached(cache_slot, lhs, rhs, config),
3280            _ => {
3281                let lhs = self.to_contiguous_read(lhs)?;
3282                let rhs = self.to_contiguous_read(rhs)?;
3283                self.dot_general_cached(cache_slot, &lhs, &rhs, config)
3284            }
3285        }
3286    }
3287
3288    // Mirrors the dot-general signature plus runtime-cache metadata.
3289    #[allow(clippy::too_many_arguments)]
3290    #[doc(hidden)]
3291    fn dot_general_with_conj_cached(
3292        &mut self,
3293        _cache_slot: Option<usize>,
3294        lhs: &Tensor,
3295        rhs: &Tensor,
3296        config: &DotGeneralConfig,
3297        lhs_conj: bool,
3298        rhs_conj: bool,
3299    ) -> crate::Result<Tensor> {
3300        self.dot_general_with_conj(lhs, rhs, config, lhs_conj, rhs_conj)
3301    }
3302
3303    // Mirrors the dot-general read signature plus runtime-cache metadata.
3304    #[allow(clippy::too_many_arguments)]
3305    #[doc(hidden)]
3306    fn dot_general_with_conj_read_cached(
3307        &mut self,
3308        cache_slot: Option<usize>,
3309        lhs: TensorRead<'_>,
3310        rhs: TensorRead<'_>,
3311        config: &DotGeneralConfig,
3312        lhs_conj: bool,
3313        rhs_conj: bool,
3314    ) -> crate::Result<Tensor> {
3315        if !lhs_conj && !rhs_conj {
3316            return self.dot_general_read_cached(cache_slot, lhs, rhs, config);
3317        }
3318
3319        let lhs_tmp;
3320        let lhs_ref = if let Some(tensor) = lhs.as_tensor() {
3321            tensor
3322        } else {
3323            lhs_tmp = self.to_contiguous_read(lhs)?;
3324            &lhs_tmp
3325        };
3326        let rhs_tmp;
3327        let rhs_ref = if let Some(tensor) = rhs.as_tensor() {
3328            tensor
3329        } else {
3330            rhs_tmp = self.to_contiguous_read(rhs)?;
3331            &rhs_tmp
3332        };
3333        self.dot_general_with_conj_cached(cache_slot, lhs_ref, rhs_ref, config, lhs_conj, rhs_conj)
3334    }
3335
3336    /// Apply session-cached scaled dot-general accumulation into output.
3337    ///
3338    /// The cache slot is session-local metadata; `accumulation` still controls
3339    /// overwrite versus read-modify-write semantics.
3340    ///
3341    /// # Examples
3342    ///
3343    /// ```rust
3344    /// use tenferro_tensor::{
3345    ///     DotGeneralAccumulation, DotGeneralConfig, SessionCachedDot, TensorRead, TensorWrite,
3346    /// };
3347    ///
3348    /// fn session_cached_dot_add_to<S: SessionCachedDot + ?Sized>(
3349    ///     session: &mut S,
3350    ///     lhs: TensorRead<'_>,
3351    ///     rhs: TensorRead<'_>,
3352    ///     config: &DotGeneralConfig,
3353    ///     out: TensorWrite<'_>,
3354    /// ) -> tenferro_tensor::Result<()> {
3355    ///     let accumulation = DotGeneralAccumulation::add_to(lhs.dtype())?;
3356    ///     session.dot_general_read_into_accum_cached(
3357    ///         Some(0),
3358    ///         lhs,
3359    ///         rhs,
3360    ///         config,
3361    ///         accumulation,
3362    ///         out,
3363    ///     )
3364    /// }
3365    /// ```
3366    /// # Errors
3367    ///
3368    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
3369    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
3370    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
3371    /// backend execution or storage access cannot provide the requested result.
3372    fn dot_general_read_into_accum_cached(
3373        &mut self,
3374        _cache_slot: Option<usize>,
3375        lhs: TensorRead<'_>,
3376        rhs: TensorRead<'_>,
3377        config: &DotGeneralConfig,
3378        accumulation: DotGeneralAccumulation,
3379        out: TensorWrite<'_>,
3380    ) -> crate::Result<()> {
3381        self.dot_general_read_into_accum(lhs, rhs, config, accumulation, out)
3382    }
3383
3384    #[doc(hidden)]
3385    fn grouped_gemm_cached(
3386        &mut self,
3387        _cache_slot: Option<usize>,
3388        lhs: TensorRead<'_>,
3389        rhs: TensorRead<'_>,
3390        config: &GroupedGemmConfig<'_>,
3391        out: TensorWrite<'_>,
3392    ) -> crate::Result<()> {
3393        grouped_gemm_default(self, lhs, rhs, config, out)
3394    }
3395}
3396
3397/// Indexing, slicing, and padding operations.
3398///
3399/// # Examples
3400///
3401/// ```rust
3402/// use tenferro_tensor::TensorIndexing;
3403///
3404/// fn accepts_indexing<B: TensorIndexing>(_backend: &mut B) {}
3405/// ```
3406pub trait TensorIndexing {
3407    /// # Errors
3408    ///
3409    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
3410    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
3411    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
3412    /// backend execution or storage access cannot provide the requested result.
3413    fn gather(
3414        &mut self,
3415        operand: &Tensor,
3416        start_indices: &Tensor,
3417        config: &GatherConfig,
3418    ) -> crate::Result<Tensor>;
3419    /// # Errors
3420    ///
3421    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
3422    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
3423    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
3424    /// backend execution or storage access cannot provide the requested result.
3425    fn scatter(
3426        &mut self,
3427        operand: &Tensor,
3428        scatter_indices: &Tensor,
3429        updates: &Tensor,
3430        config: &ScatterConfig,
3431    ) -> crate::Result<Tensor>;
3432    /// # Errors
3433    ///
3434    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
3435    /// for invalid shapes, ranks, axes, dtypes, or output metadata. In
3436    /// particular, a limit greater than the corresponding input dimension is
3437    /// reported as [`crate::ValidationError::InvalidArgument`] with the
3438    /// `"configuration"` argument. It returns [`crate::Error::BackendFailure`]
3439    /// or [`crate::Error::BackendSource`] when backend execution or storage
3440    /// access cannot provide the requested result.
3441    fn slice(&mut self, input: &Tensor, config: &SliceConfig) -> crate::Result<Tensor>;
3442    /// # Errors
3443    ///
3444    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
3445    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
3446    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
3447    /// backend execution or storage access cannot provide the requested result.
3448    fn dynamic_slice(
3449        &mut self,
3450        input: &Tensor,
3451        starts: &Tensor,
3452        slice_sizes: &[usize],
3453    ) -> crate::Result<Tensor>;
3454    /// # Errors
3455    ///
3456    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
3457    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
3458    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
3459    /// backend execution or storage access cannot provide the requested result.
3460    fn dynamic_update_slice(
3461        &mut self,
3462        operand: &Tensor,
3463        update: &Tensor,
3464        starts: &Tensor,
3465    ) -> crate::Result<Tensor>;
3466    /// # Errors
3467    ///
3468    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
3469    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
3470    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
3471    /// backend execution or storage access cannot provide the requested result.
3472    fn pad(&mut self, input: &Tensor, config: &PadConfig) -> crate::Result<Tensor>;
3473    /// # Errors
3474    ///
3475    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
3476    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
3477    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
3478    /// backend execution or storage access cannot provide the requested result.
3479    fn concatenate(&mut self, inputs: &[&Tensor], axis: usize) -> crate::Result<Tensor>;
3480    /// # Errors
3481    ///
3482    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
3483    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
3484    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
3485    /// backend execution or storage access cannot provide the requested result.
3486    fn reverse(&mut self, input: &Tensor, axes: &[usize]) -> crate::Result<Tensor>;
3487}
3488
3489/// Backend-owned canonicalization for typed tensor views.
3490///
3491/// Implementations must preserve the input placement family. CPU backends
3492/// canonicalize host views through explicit host copies and reject backend
3493/// buffers with a diagnostic that asks the caller to download first. GPU
3494/// backends canonicalize GPU-resident views on the same device and reject host
3495/// buffers with an upload hint.
3496///
3497/// [`TensorViewCanonicalization::copy_into`] requires source and destination
3498/// shapes, scalar dtypes, and placement families to match. The destination
3499/// view must be internally non-overlapping, and source and destination backing
3500/// allocations must not alias unless an implementation explicitly documents
3501/// and supports that case. Implementations may reject layouts their native
3502/// kernels cannot consume.
3503///
3504/// CUDA currently accepts only a compact column-major source view with offset
3505/// zero that covers its full allocation; arbitrary-stride destinations remain
3506/// supported. Canonicalization and copying are same-placement operations: they
3507/// must not perform hidden host/device transfers or silently materialize an
3508/// unsupported source layout.
3509///
3510/// This trait is intentionally separate from [`BackendSession`] so generic
3511/// typed methods do not change the object-safety contract of `dyn BackendSession`.
3512///
3513/// # Examples
3514///
3515/// ```rust
3516/// use tenferro_tensor::{DynRank, TensorViewCanonicalization, TypedTensor};
3517///
3518/// fn compact_i32<B: TensorViewCanonicalization<i32, DynRank>>(
3519///     backend: &mut B,
3520///     tensor: &TypedTensor<i32>,
3521/// ) -> tenferro_tensor::Result<TypedTensor<i32>> {
3522///     backend.to_contiguous(&tensor.as_view())
3523/// }
3524///
3525/// fn copy_i32<B: TensorViewCanonicalization<i32, DynRank>>(
3526///     backend: &mut B,
3527///     src: &TypedTensor<i32>,
3528///     dst: &mut TypedTensor<i32>,
3529/// ) -> tenferro_tensor::Result<()> {
3530///     backend.copy_into(&src.as_view(), &mut dst.as_view_mut())
3531/// }
3532/// ```
3533pub trait TensorViewCanonicalization<T: TensorScalar, R: TensorRank> {
3534    /// # Errors
3535    ///
3536    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
3537    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
3538    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
3539    /// backend execution or storage access cannot provide the requested result.
3540    fn to_contiguous(
3541        &mut self,
3542        view: &TypedTensorView<'_, T, R>,
3543    ) -> crate::Result<TypedTensor<T, R>>;
3544
3545    /// # Errors
3546    ///
3547    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
3548    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
3549    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
3550    /// backend execution or storage access cannot provide the requested result.
3551    fn copy_into(
3552        &mut self,
3553        src: &TypedTensorView<'_, T, R>,
3554        dst: &mut TypedTensorViewMut<'_, T, R>,
3555    ) -> crate::Result<()>;
3556}
3557
3558/// Optional elementwise fusion execution.
3559///
3560/// # Examples
3561///
3562/// ```rust
3563/// use tenferro_tensor::TensorFusion;
3564///
3565/// fn accepts_fusion<B: TensorFusion>(_backend: &mut B) {}
3566/// ```
3567pub trait TensorFusion {
3568    #[doc(hidden)]
3569    fn execute_elementwise_fusion(
3570        &mut self,
3571        _inputs: &[&Tensor],
3572        _plan: &ElementwiseFusionPlan,
3573    ) -> crate::Result<Option<Vec<Tensor>>> {
3574        Ok(None)
3575    }
3576
3577    #[doc(hidden)]
3578    #[allow(clippy::too_many_arguments)]
3579    fn execute_broadcast_multiply(
3580        &mut self,
3581        _lhs: TensorRead<'_>,
3582        _lhs_shape: &[usize],
3583        _lhs_dims: &[usize],
3584        _rhs: TensorRead<'_>,
3585        _rhs_shape: &[usize],
3586        _rhs_dims: &[usize],
3587    ) -> crate::Result<Option<Tensor>> {
3588        Ok(None)
3589    }
3590
3591    #[doc(hidden)]
3592    #[allow(clippy::too_many_arguments)]
3593    fn execute_broadcast_multiply_value(
3594        &mut self,
3595        lhs: TensorRead<'_>,
3596        lhs_shape: &[usize],
3597        lhs_dims: &[usize],
3598        rhs: TensorRead<'_>,
3599        rhs_shape: &[usize],
3600        rhs_dims: &[usize],
3601    ) -> crate::Result<Option<TensorValue>> {
3602        self.execute_broadcast_multiply(lhs, lhs_shape, lhs_dims, rhs, rhs_shape, rhs_dims)
3603            .map(|tensor| tensor.map(TensorValue::from_tensor))
3604    }
3605}
3606
3607/// Backend buffer lifecycle operations.
3608///
3609/// # Examples
3610///
3611/// ```rust
3612/// use tenferro_tensor::TensorBuffer;
3613///
3614/// fn accepts_buffer<B: TensorBuffer>(_backend: &mut B) {}
3615/// ```
3616pub trait TensorBuffer {
3617    fn reclaim_buffer(&mut self, _tensor: Tensor) {}
3618}
3619
3620/// Device transfer operations on backend boundaries.
3621///
3622/// # Examples
3623///
3624/// ```rust
3625/// use tenferro_tensor::TensorDeviceTransfer;
3626///
3627/// fn accepts_transfer<B: TensorDeviceTransfer>(_backend: &mut B) {}
3628/// ```
3629pub trait TensorDeviceTransfer {
3630    /// Explicitly copy a provider-owned read target into host storage.
3631    ///
3632    /// Implementations must not return the input unchanged or stage through an
3633    /// unrelated provider. A backend that cannot transfer the requested read
3634    /// target returns a typed unsupported error.
3635    ///
3636    /// # Errors
3637    ///
3638    /// Returns [`crate::Error::Unsupported`] when the implementation cannot
3639    /// perform the requested transfer, or a typed validation/backend error when
3640    /// the source cannot be read.
3641    fn download_to_host(&mut self, tensor: TensorRead<'_>) -> crate::Result<Tensor>;
3642
3643    /// Explicitly copy a host read target into provider storage.
3644    ///
3645    /// # Errors
3646    ///
3647    /// Returns [`crate::Error::Unsupported`] when the implementation cannot
3648    /// perform the requested transfer, or a typed validation/backend error when
3649    /// the source cannot be read.
3650    fn upload_host_tensor(&mut self, tensor: TensorRead<'_>) -> crate::Result<Tensor>;
3651}
3652
3653/// Runtime cache associated with a backend.
3654///
3655/// # Examples
3656///
3657/// ```rust
3658/// use tenferro_tensor::BackendRuntimeCache;
3659///
3660/// fn accepts_runtime_cache<B: BackendRuntimeCache>(_backend: &B) {}
3661/// ```
3662pub trait BackendRuntimeCache {
3663    #[doc(hidden)]
3664    type RuntimeCache: RuntimeCacheControl + Send + Sync + 'static;
3665}
3666
3667/// Backend-owned cached dot-general operations.
3668///
3669/// # Examples
3670///
3671/// ```rust
3672/// use tenferro_tensor::BackendCachedDot;
3673///
3674/// fn accepts_backend_cached_dot<B: BackendCachedDot>(_backend: &mut B) {}
3675/// ```
3676pub trait BackendCachedDot: BackendRuntimeCache + TensorDot {
3677    #[doc(hidden)]
3678    fn dot_general_cached(
3679        &mut self,
3680        _cache: &mut Self::RuntimeCache,
3681        _cache_slot: Option<usize>,
3682        lhs: &Tensor,
3683        rhs: &Tensor,
3684        config: &DotGeneralConfig,
3685    ) -> crate::Result<Tensor> {
3686        self.dot_general_read(
3687            TensorRead::from_tensor(lhs),
3688            TensorRead::from_tensor(rhs),
3689            config,
3690        )
3691    }
3692
3693    #[doc(hidden)]
3694    fn dot_general_read_cached(
3695        &mut self,
3696        cache: &mut Self::RuntimeCache,
3697        cache_slot: Option<usize>,
3698        lhs: TensorRead<'_>,
3699        rhs: TensorRead<'_>,
3700        config: &DotGeneralConfig,
3701    ) -> crate::Result<Tensor> {
3702        match (lhs.as_tensor(), rhs.as_tensor()) {
3703            (Some(lhs), Some(rhs)) => self.dot_general_cached(cache, cache_slot, lhs, rhs, config),
3704            _ => {
3705                let lhs = self.to_contiguous_read(lhs)?;
3706                let rhs = self.to_contiguous_read(rhs)?;
3707                self.dot_general_cached(cache, cache_slot, &lhs, &rhs, config)
3708            }
3709        }
3710    }
3711
3712    // Mirrors the dot-general signature plus runtime-cache metadata.
3713    #[allow(clippy::too_many_arguments)]
3714    #[doc(hidden)]
3715    fn dot_general_with_conj_cached(
3716        &mut self,
3717        _cache: &mut Self::RuntimeCache,
3718        _cache_slot: Option<usize>,
3719        lhs: &Tensor,
3720        rhs: &Tensor,
3721        config: &DotGeneralConfig,
3722        lhs_conj: bool,
3723        rhs_conj: bool,
3724    ) -> crate::Result<Tensor> {
3725        self.dot_general_with_conj(lhs, rhs, config, lhs_conj, rhs_conj)
3726    }
3727
3728    // Mirrors the dot-general read signature plus runtime-cache metadata.
3729    #[allow(clippy::too_many_arguments)]
3730    #[doc(hidden)]
3731    fn dot_general_with_conj_read_cached(
3732        &mut self,
3733        cache: &mut Self::RuntimeCache,
3734        cache_slot: Option<usize>,
3735        lhs: TensorRead<'_>,
3736        rhs: TensorRead<'_>,
3737        config: &DotGeneralConfig,
3738        lhs_conj: bool,
3739        rhs_conj: bool,
3740    ) -> crate::Result<Tensor> {
3741        if !lhs_conj && !rhs_conj {
3742            return self.dot_general_read_cached(cache, cache_slot, lhs, rhs, config);
3743        }
3744
3745        let lhs_tmp;
3746        let lhs_ref = if let Some(tensor) = lhs.as_tensor() {
3747            tensor
3748        } else {
3749            lhs_tmp = self.to_contiguous_read(lhs)?;
3750            &lhs_tmp
3751        };
3752        let rhs_tmp;
3753        let rhs_ref = if let Some(tensor) = rhs.as_tensor() {
3754            tensor
3755        } else {
3756            rhs_tmp = self.to_contiguous_read(rhs)?;
3757            &rhs_tmp
3758        };
3759        self.dot_general_with_conj_cached(
3760            cache, cache_slot, lhs_ref, rhs_ref, config, lhs_conj, rhs_conj,
3761        )
3762    }
3763
3764    /// Apply cached scaled dot-general accumulation into caller-provided output.
3765    ///
3766    /// The cache slot identifies backend-local analysis metadata only; output
3767    /// semantics are still fully described by `accumulation`.
3768    ///
3769    /// # Examples
3770    ///
3771    /// ```rust
3772    /// use tenferro_tensor::{
3773    ///     BackendCachedDot, BackendRuntimeCache, DotGeneralAccumulation, DotGeneralConfig,
3774    ///     TensorRead, TensorWrite,
3775    /// };
3776    ///
3777    /// fn cached_dot_add_to<B: BackendCachedDot>(
3778    ///     backend: &mut B,
3779    ///     cache: &mut B::RuntimeCache,
3780    ///     lhs: TensorRead<'_>,
3781    ///     rhs: TensorRead<'_>,
3782    ///     config: &DotGeneralConfig,
3783    ///     out: TensorWrite<'_>,
3784    /// ) -> tenferro_tensor::Result<()>
3785    /// where
3786    ///     B: BackendRuntimeCache,
3787    /// {
3788    ///     let accumulation = DotGeneralAccumulation::add_to(lhs.dtype())?;
3789    ///     backend.dot_general_read_into_accum_cached(
3790    ///         cache,
3791    ///         Some(0),
3792    ///         lhs,
3793    ///         rhs,
3794    ///         config,
3795    ///         accumulation,
3796    ///         out,
3797    ///     )
3798    /// }
3799    /// ```
3800    #[allow(clippy::too_many_arguments)]
3801    /// # Errors
3802    ///
3803    /// Returns [`crate::Error::Validation`] with a typed `ValidationError` source
3804    /// for invalid shapes, ranks, axes, dtypes, or output metadata. It returns
3805    /// [`crate::Error::BackendFailure`] or [`crate::Error::BackendSource`] when
3806    /// backend execution or storage access cannot provide the requested result.
3807    fn dot_general_read_into_accum_cached(
3808        &mut self,
3809        _cache: &mut Self::RuntimeCache,
3810        _cache_slot: Option<usize>,
3811        lhs: TensorRead<'_>,
3812        rhs: TensorRead<'_>,
3813        config: &DotGeneralConfig,
3814        accumulation: DotGeneralAccumulation,
3815        out: TensorWrite<'_>,
3816    ) -> crate::Result<()> {
3817        self.dot_general_read_into_accum(lhs, rhs, config, accumulation, out)
3818    }
3819
3820    #[doc(hidden)]
3821    fn grouped_gemm_cached(
3822        &mut self,
3823        _cache: &mut Self::RuntimeCache,
3824        _cache_slot: Option<usize>,
3825        lhs: TensorRead<'_>,
3826        rhs: TensorRead<'_>,
3827        config: &GroupedGemmConfig<'_>,
3828        out: TensorWrite<'_>,
3829    ) -> crate::Result<()> {
3830        grouped_gemm_default(self, lhs, rhs, config, out)
3831    }
3832}
3833
3834/// Backend execution-session entry points.
3835///
3836/// `with_backend_session` is the canonical user entry and the only way an
3837/// operation is reached from a backend; `with_backend_session_cached` is the
3838/// runtime-cache-aware entry used by the runtime layer.
3839///
3840/// # Examples
3841///
3842/// ```rust
3843/// use tenferro_tensor::BackendSessionHost;
3844///
3845/// fn accepts_session_host<B: BackendSessionHost>(_backend: &mut B) {}
3846/// ```
3847///
3848/// The session factory this trait's default bodies used to call was deleted
3849/// with the one-shot spellings, so naming it does not compile:
3850///
3851/// ```compile_fail
3852/// use tenferro_tensor::default_backend_session;
3853/// ```
3854pub trait BackendSessionHost: BackendRuntimeCache {
3855    /// Open one backend session and run `f` inside it.
3856    ///
3857    /// The session is built by the backend rather than by coercing the owner, so
3858    /// this is the only way an operation is reached from a backend.
3859    ///
3860    /// Admission happens before `f` runs. A backend may wait for resources that
3861    /// another thread will release (CPU resource permits are granted in FIFO
3862    /// order), but it never runs `f` twice, never retries `f` through a fallback
3863    /// path, and never reports a failure of `f` as an admission failure: the
3864    /// callback's own value, including its own `Result`, is returned unchanged in
3865    /// `Ok`.
3866    ///
3867    /// # Errors
3868    ///
3869    /// Returns [`SessionEntryError`] without running `f` when the backend cannot
3870    /// admit the session: [`SessionEntryError::Reentered`] for a nested entry on
3871    /// the calling thread, [`SessionEntryError::Contended`] for a busy resource
3872    /// that admission cannot wait for, [`SessionEntryError::IncompatibleContext`]
3873    /// for a mismatched execution scope or executor declaration,
3874    /// [`SessionEntryError::ResourcePoisoned`] for poisoned admission state, and
3875    /// [`SessionEntryError::Executor`] when the executor cannot be entered.
3876    fn with_backend_session<R: Send>(
3877        &mut self,
3878        f: impl FnOnce(&mut dyn BackendSession) -> R + Send,
3879    ) -> Result<R, SessionEntryError>;
3880
3881    #[doc(hidden)]
3882    fn with_backend_session_cached<R: Send>(
3883        &mut self,
3884        _cache: &mut Self::RuntimeCache,
3885        f: impl FnOnce(&mut dyn BackendSession) -> R + Send,
3886    ) -> Result<R, SessionEntryError> {
3887        self.with_backend_session(f)
3888    }
3889}
3890
3891/// Operation capabilities shared by backends and backend sessions.
3892#[doc(hidden)]
3893pub trait TensorBackendOps:
3894    TensorElementwise
3895    + TensorAnalytic
3896    + TensorStructural
3897    + TensorReduction
3898    + TensorIndexing
3899    + TensorDot
3900    + TensorFusion
3901    + TensorBuffer
3902{
3903}
3904
3905impl<T> TensorBackendOps for T where
3906    T: TensorElementwise
3907        + TensorAnalytic
3908        + TensorStructural
3909        + TensorReduction
3910        + TensorIndexing
3911        + TensorDot
3912        + TensorFusion
3913        + TensorBuffer
3914        + ?Sized
3915{
3916}
3917
3918/// Validate the shared input contract for [`BackendSession::vdot_read`].
3919#[doc(hidden)]
3920pub fn validate_vdot_read(lhs: &TensorRead<'_>, rhs: &TensorRead<'_>) -> crate::Result<()> {
3921    let op = "BackendSession::vdot_read";
3922    validate_supported_blas1_dtype(op, lhs.dtype())?;
3923    if rhs.dtype() != lhs.dtype() {
3924        return Err(Error::dtype_mismatch(op, lhs.dtype(), rhs.dtype()));
3925    }
3926    if lhs.shape() != rhs.shape() {
3927        return Err(Error::shape_mismatch(
3928            op,
3929            lhs.shape().to_vec(),
3930            rhs.shape().to_vec(),
3931        ));
3932    }
3933    validate_shape_product(op, lhs.shape())?;
3934    validate_compatible_placement(op, lhs, rhs)
3935}
3936
3937/// Validate the shared input contract for [`BackendSession::norm_squared_read`].
3938#[doc(hidden)]
3939pub fn validate_norm_squared_read(input: &TensorRead<'_>) -> crate::Result<()> {
3940    validate_supported_blas1_dtype("BackendSession::norm_squared_read", input.dtype())?;
3941    validate_shape_product("BackendSession::norm_squared_read", input.shape()).map(|_| ())
3942}
3943
3944/// Validate the shared input and destination contract for
3945/// [`BackendSession::axpby_read_into_accum`].
3946#[doc(hidden)]
3947pub fn validate_axpby_read_into_accum(
3948    alpha: ContractionScalar,
3949    x: &TensorRead<'_>,
3950    beta: ContractionScalar,
3951    y: &TensorWrite<'_>,
3952) -> crate::Result<()> {
3953    let op = "BackendSession::axpby_read_into_accum";
3954    validate_supported_blas1_dtype(op, x.dtype())?;
3955    if y.dtype() != x.dtype() {
3956        return Err(Error::dtype_mismatch(op, x.dtype(), y.dtype()));
3957    }
3958    if alpha.dtype() != x.dtype() {
3959        return Err(Error::dtype_mismatch(op, x.dtype(), alpha.dtype()));
3960    }
3961    if beta.dtype() != x.dtype() {
3962        return Err(Error::dtype_mismatch(op, x.dtype(), beta.dtype()));
3963    }
3964    if x.shape() != y.shape() {
3965        return Err(Error::shape_mismatch(
3966            op,
3967            x.shape().to_vec(),
3968            y.shape().to_vec(),
3969        ));
3970    }
3971    validate_shape_product(op, x.shape())?;
3972    if !y.as_read().is_col_major_contiguous()? {
3973        return Err(Error::invalid_argument(
3974            op,
3975            "y",
3976            "destination must be compact column-major and injective",
3977        ));
3978    }
3979    validate_compatible_placement(op, x, &y.as_read())?;
3980    validate_read_into_destination(op, std::slice::from_ref(x), y)
3981}
3982
3983/// # Errors
3984///
3985/// Returns [`Error::Unsupported`] when `dtype` is not floating or complex.
3986fn validate_supported_blas1_dtype(op: &'static str, dtype: DType) -> crate::Result<()> {
3987    if matches!(dtype, DType::F32 | DType::F64 | DType::C32 | DType::C64) {
3988        Ok(())
3989    } else {
3990        Err(Error::unsupported(
3991            op,
3992            format!("unsupported dtype {dtype:?}; supported dtypes: F32/F64/C32/C64"),
3993        ))
3994    }
3995}
3996
3997/// # Errors
3998///
3999/// Returns [`Error::Validation`] with [`ValidationError::IntegerOverflow`] when
4000/// the checked shape product overflows.
4001fn validate_shape_product(op: &'static str, shape: &[usize]) -> crate::Result<usize> {
4002    crate::validate::checked_shape_product(op, "shape", shape)
4003}
4004
4005/// # Errors
4006///
4007/// Returns [`Error::Validation`] with [`ValidationError::InvalidArgument`] when
4008/// placements differ.
4009fn validate_compatible_placement(
4010    op: &'static str,
4011    lhs: &TensorRead<'_>,
4012    rhs: &TensorRead<'_>,
4013) -> crate::Result<()> {
4014    if lhs.placement() != rhs.placement() {
4015        return Err(Error::invalid_argument(
4016            op,
4017            "placement",
4018            "tensor inputs must have compatible placement",
4019        ));
4020    }
4021    Ok(())
4022}
4023
4024/// Execution session surface for dense tensor backends.
4025///
4026/// All operations run within a backend-owned execution scope such as a CPU
4027/// thread policy or a GPU stream. Individual ops must not try to re-enter that
4028/// scope.
4029///
4030/// # Examples
4031///
4032/// ```rust
4033/// use tenferro_tensor::{BackendSessionHost, Tensor, TensorRead, TypedTensor};
4034///
4035/// fn add_in_session<B: BackendSessionHost>(
4036///     backend: &mut B,
4037///     a: &Tensor,
4038///     b: &Tensor,
4039/// ) -> tenferro_tensor::Result<Tensor>
4040/// where
4041///     B: tenferro_tensor::TensorBackend,
4042/// {
4043///     // Admission failure converts into `tenferro_tensor::Error::SessionEntry`;
4044///     // the operation's own result is returned unchanged.
4045///     backend.with_backend_session(|exec| {
4046///         exec.add_read(TensorRead::from_tensor(a), TensorRead::from_tensor(b))
4047///     })?
4048/// }
4049/// ```
4050///
4051/// The operation one-shot spelling is gone; a session only answers to the read
4052/// form:
4053///
4054/// ```compile_fail
4055/// use tenferro_tensor::{BackendSessionHost, Tensor, TensorBackend};
4056///
4057/// fn add_in_session<B: BackendSessionHost + TensorBackend>(
4058///     backend: &mut B,
4059///     a: &Tensor,
4060///     b: &Tensor,
4061/// ) {
4062///     backend.with_backend_session(|exec| {
4063///         let _ = exec.add(a, b);
4064///     });
4065/// }
4066/// ```
4067pub trait BackendSession: TensorBackendOps + SessionCachedDot + TensorDeviceTransfer {
4068    /// Compute the all-axis conjugating dot product without transferring either input.
4069    ///
4070    /// The result is a rank-0 tensor with the input dtype and has the value
4071    /// `sum(conj(lhs) * rhs)`. Borrowed views remain borrowed through the
4072    /// backend's existing same-placement planning boundary.
4073    ///
4074    /// # Examples
4075    ///
4076    /// ```rust
4077    /// use tenferro_tensor::{BackendSession, TensorRead};
4078    ///
4079    /// fn vdot(session: &mut dyn BackendSession, x: TensorRead<'_>, y: TensorRead<'_>)
4080    ///     -> tenferro_tensor::Result<tenferro_tensor::Tensor>
4081    /// {
4082    ///     session.vdot_read(x, y)
4083    /// }
4084    /// ```
4085    ///
4086    /// # Errors
4087    ///
4088    /// Returns [`Error::Unsupported`] when the backend does not override this
4089    /// capability or the dtype is unsupported; [`Error::Validation`] with
4090    /// `DTypeMismatch`, `ShapeMismatch`, or `InvalidArgument` when dtype, shape,
4091    /// or placement differs; [`Error::RuntimeState`] for inaccessible backend
4092    /// storage; or [`Error::BackendSource`] when provider execution fails.
4093    fn vdot_read(&mut self, _lhs: TensorRead<'_>, _rhs: TensorRead<'_>) -> crate::Result<Tensor> {
4094        Err(Error::unsupported(
4095            "BackendSession::vdot_read",
4096            "backend session does not implement vdot_read",
4097        ))
4098    }
4099
4100    /// Compute the all-axis sum of squared magnitudes without taking a square root.
4101    ///
4102    /// The result is rank 0 and is F32 for F32/C32 input or F64 for F64/C64
4103    /// input. No transfer or full-size algebra temporary is implied by this
4104    /// session contract.
4105    ///
4106    /// # Examples
4107    ///
4108    /// ```rust
4109    /// use tenferro_tensor::{BackendSession, TensorRead};
4110    ///
4111    /// fn norm_squared(session: &mut dyn BackendSession, x: TensorRead<'_>)
4112    ///     -> tenferro_tensor::Result<tenferro_tensor::Tensor>
4113    /// {
4114    ///     session.norm_squared_read(x)
4115    /// }
4116    /// ```
4117    ///
4118    /// # Errors
4119    ///
4120    /// Returns [`Error::Unsupported`] when the backend does not override this
4121    /// capability or the dtype is unsupported, [`Error::RuntimeState`] when
4122    /// backend storage is not host-accessible, or [`Error::BackendSource`] when
4123    /// reduction execution fails.
4124    fn norm_squared_read(&mut self, _input: TensorRead<'_>) -> crate::Result<Tensor> {
4125        Err(Error::unsupported(
4126            "BackendSession::norm_squared_read",
4127            "backend session does not implement norm_squared_read",
4128        ))
4129    }
4130
4131    /// Apply `y <- alpha * x + beta * y` in one pass into caller-owned storage.
4132    ///
4133    /// Scalars must have the exact tensor dtype; real coefficients for complex
4134    /// vectors are represented as complex values with zero imaginary part. The
4135    /// destination must be compact and injective, and any x/y storage overlap
4136    /// is rejected before mutation.
4137    ///
4138    /// # Examples
4139    ///
4140    /// ```rust
4141    /// use tenferro_tensor::{BackendSession, ContractionScalar, TensorRead, TensorWrite};
4142    ///
4143    /// fn axpby(session: &mut dyn BackendSession, x: TensorRead<'_>, y: TensorWrite<'_>)
4144    ///     -> tenferro_tensor::Result<()>
4145    /// {
4146    ///     session.axpby_read_into_accum(
4147    ///         ContractionScalar::F64(1.0), x, ContractionScalar::F64(0.0), y,
4148    ///     )
4149    /// }
4150    /// ```
4151    ///
4152    /// # Errors
4153    ///
4154    /// Returns [`Error::Unsupported`] when the backend does not override this
4155    /// capability or the dtype is unsupported; [`Error::Validation`] with
4156    /// `DTypeMismatch`, `ShapeMismatch`, or `InvalidArgument` for scalar/dtype,
4157    /// shape, placement, compactness, injectivity, or overlap failures; or
4158    /// [`Error::RuntimeState`] for inaccessible backend storage. Invalid
4159    /// requests leave the destination unchanged.
4160    fn axpby_read_into_accum(
4161        &mut self,
4162        _alpha: ContractionScalar,
4163        _x: TensorRead<'_>,
4164        _beta: ContractionScalar,
4165        _y: TensorWrite<'_>,
4166    ) -> crate::Result<()> {
4167        Err(Error::unsupported(
4168            "BackendSession::axpby_read_into_accum",
4169            "backend session does not implement axpby_read_into_accum",
4170        ))
4171    }
4172
4173    /// Return this session's backend-leaf native capability, if it has one.
4174    ///
4175    /// Standard CPU, CUDA and WebGPU sessions return a token their own safe
4176    /// visitors (`with_cpu_exec_session`, `with_cuda_exec_session`,
4177    /// `with_webgpu_exec_session`) recover. The default is `None`: a custom
4178    /// session has no native services unless it forwards the token of a
4179    /// standard session it owns. A wrapper that overrides operation dispatch
4180    /// (for example a custom GEMM) should keep the default, because an
4181    /// operation family that finds a native token may call the delegate's
4182    /// native services directly and so bypass the wrapper's override.
4183    ///
4184    /// # Examples
4185    ///
4186    /// ```rust
4187    /// use tenferro_tensor::BackendSession;
4188    ///
4189    /// fn native_services_available(session: &mut dyn BackendSession) -> bool {
4190    ///     session.native_session().is_some()
4191    /// }
4192    /// ```
4193    fn native_session(&mut self) -> Option<crate::NativeSessionRef<'_>> {
4194        None
4195    }
4196}
4197
4198/// Standard runtime backend over dynamic [`Tensor`] values.
4199///
4200/// # Examples
4201///
4202/// ```rust
4203/// use tenferro_tensor::TensorBackend;
4204///
4205/// fn accepts_backend<B: TensorBackend>(_backend: &mut B) {}
4206/// ```
4207pub trait TensorBackend: BackendRuntimeCache + TensorDeviceTransfer + BackendSessionHost {}
4208
4209impl<T> SessionCachedDot for T where T: TensorBackend + TensorDot + ?Sized {}
4210
4211thread_local! {
4212    /// Tracks whether a session-entry closure is currently running on this
4213    /// thread, so nested session entry is caught in debug builds even for
4214    /// backends that do not override
4215    /// [`BackendSessionHost::with_backend_session`].
4216    static IN_SESSION: std::cell::Cell<bool> = const { std::cell::Cell::new(false) };
4217}
4218
4219/// Sets the in-session flag for the duration of a session closure and restores
4220/// it on exit, including on panic.
4221struct InSessionGuard;
4222
4223impl InSessionGuard {
4224    /// Set the in-session flag for one session closure.
4225    ///
4226    /// # Errors
4227    ///
4228    /// Returns [`SessionEntryError::Reentered`] when a session closure is
4229    /// already running on this thread.
4230    fn enter(backend: &'static str) -> Result<Self, SessionEntryError> {
4231        if IN_SESSION.get() {
4232            return Err(SessionEntryError::Reentered { backend });
4233        }
4234        IN_SESSION.set(true);
4235        Ok(InSessionGuard)
4236    }
4237}
4238
4239impl Drop for InSessionGuard {
4240    fn drop(&mut self) {
4241        IN_SESSION.set(false);
4242    }
4243}
4244
4245/// Whether a portable backend-session closure is running on the current thread.
4246///
4247/// Hosts that enter through [`with_session_entry_guard`] (CUDA, WebGPU and
4248/// backends keeping the default entry) set this flag for the duration of their
4249/// session closure. An owner that serializes callers with a blocking lock must
4250/// check it, together with its backend's own admission state, *before* waiting
4251/// on that lock: a thread holding a session must not wait on an owner another
4252/// thread holds while that thread waits for the same session's resources.
4253///
4254/// # Examples
4255///
4256/// ```
4257/// assert!(!tenferro_tensor::has_active_backend_session());
4258/// tenferro_tensor::with_session_entry_guard("doc", || {
4259///     assert!(tenferro_tensor::has_active_backend_session());
4260/// })?;
4261/// assert!(!tenferro_tensor::has_active_backend_session());
4262/// # Ok::<(), tenferro_tensor::SessionEntryError>(())
4263/// ```
4264pub fn has_active_backend_session() -> bool {
4265    IN_SESSION.get()
4266}
4267
4268/// Run `f` with the thread-local in-session flag set, restoring it on exit
4269/// including on panic.
4270///
4271/// This is the portable nested-entry guard shared by backend-session entry
4272/// points that have no resource permit of their own. A nested entry on the same
4273/// thread is rejected with [`SessionEntryError::Reentered`] before `f` runs, in
4274/// every build profile. CPU admission tracks reentry through its execution
4275/// owner instead.
4276#[doc(hidden)]
4277pub fn with_session_entry_guard<R>(
4278    backend: &'static str,
4279    f: impl FnOnce() -> R,
4280) -> Result<R, SessionEntryError> {
4281    let _guard = InSessionGuard::enter(backend)?;
4282    Ok(f())
4283}