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}