Skip to main content

tenferro_einsum/
concrete.rs

1//! Public concrete tensor einsum extension API.
2
3use smallvec::SmallVec;
4use tenferro_tensor::{
5    BackendSession, DType, DotGeneralAccumulation, Tensor, TensorRead, TensorScalar, TensorWrite,
6    TypedTensor, TypedTensorView, TypedTensorWrite,
7};
8
9use crate::binary_dot::BinaryDotOperandOrder;
10use crate::eager::{
11    binary_dot_config_for_into, binary_dot_plan_for_shapes, eager_einsum_exec,
12    eager_einsum_exec_read, eager_einsum_exec_read_into, eager_einsum_exec_read_into_accum,
13    eager_einsum_read_subscripts_on_session, eager_einsum_subscripts_on_session,
14    execute_binary_dot_read_into, execute_binary_dot_read_into_accum, plan_subscripts,
15};
16use crate::ellipsis::resolve_einsum_notation;
17use crate::TensorDotAxes;
18use crate::{
19    parse_einsum_notation, ContractionTree, EinsumNotation, EinsumSubscripts, Error, Result,
20    Subscripts,
21};
22
23const TENSOR_EINSUM_INTO_OP: &str = "TensorEinsumIntoExt::einsum_into";
24const TENSOR_READ_EINSUM_INTO_OP: &str = "TensorReadEinsumIntoExt::einsum_read_into";
25const TYPED_TENSOR_EINSUM_OP: &str = "TypedTensorEinsumExt::einsum";
26const TYPED_TENSOR_EINSUM_INTO_OP: &str = "TypedTensorEinsumIntoExt::einsum_into";
27const TYPED_TENSOR_READ_EINSUM_OP: &str = "TypedTensorReadEinsumExt::einsum_read";
28const TYPED_TENSOR_READ_EINSUM_INTO_OP: &str = "TypedTensorReadEinsumIntoExt::einsum_read_into";
29const PLAN_EXECUTE_OP: &str = "ConcreteEinsumPlan::execute";
30const TYPED_TENSOR_TENSORDOT_OP: &str = "TypedTensorTensordotExt::tensordot";
31
32/// Backend-explicit tensordot sugar for dtype-erased concrete tensors.
33pub trait TensorTensordotExt {
34    /// Contract this tensor with `rhs` over explicit axes or an axis count.
35    ///
36    /// # Errors
37    ///
38    /// Returns [`Error::Validation`] with `InvalidArgument`, `RankMismatch`,
39    /// `AxisOutOfBounds`, or `ShapeMismatch` for invalid axes or incompatible
40    /// contracting dimensions, or [`Error::Tensor`] when backend execution
41    /// fails.
42    fn tensordot(
43        &self,
44        rhs: &Tensor,
45        axes: TensorDotAxes<'_>,
46        session: &mut dyn BackendSession,
47    ) -> Result<Tensor>;
48}
49
50impl TensorTensordotExt for Tensor {
51    fn tensordot(
52        &self,
53        rhs: &Tensor,
54        axes: TensorDotAxes<'_>,
55        session: &mut dyn BackendSession,
56    ) -> Result<Tensor> {
57        let config =
58            crate::tensordot::dot_general_config(axes, self.shape().len(), rhs.shape().len())?;
59        crate::tensordot::validate_concrete_contract_dims(self.shape(), rhs.shape(), &config)?;
60        session
61            .dot_general_read(
62                TensorRead::from_tensor(self),
63                TensorRead::from_tensor(rhs),
64                &config,
65            )
66            .map_err(Error::from)
67    }
68}
69
70/// Backend-explicit tensordot sugar for typed concrete tensors.
71pub trait TypedTensorTensordotExt<T: TensorScalar> {
72    /// Contract this tensor with `rhs` while preserving its scalar type.
73    ///
74    /// # Errors
75    ///
76    /// Returns [`Error::Validation`] with `InvalidArgument`, `RankMismatch`,
77    /// `AxisOutOfBounds`, or `ShapeMismatch` for invalid axes or incompatible
78    /// contracting dimensions, or [`Error::Tensor`] when backend execution
79    /// fails.
80    fn tensordot(
81        &self,
82        rhs: &TypedTensor<T>,
83        axes: TensorDotAxes<'_>,
84        session: &mut dyn BackendSession,
85    ) -> Result<TypedTensor<T>>;
86}
87
88impl<T: TensorScalar> TypedTensorTensordotExt<T> for TypedTensor<T> {
89    fn tensordot(
90        &self,
91        rhs: &TypedTensor<T>,
92        axes: TensorDotAxes<'_>,
93        session: &mut dyn BackendSession,
94    ) -> Result<TypedTensor<T>> {
95        let config =
96            crate::tensordot::dot_general_config(axes, self.shape().len(), rhs.shape().len())?;
97        crate::tensordot::validate_concrete_contract_dims(self.shape(), rhs.shape(), &config)?;
98        let result = session
99            .dot_general_read(T::tensor_read(self), T::tensor_read(rhs), &config)
100            .map_err(Error::from)?;
101        into_typed_result(result, TYPED_TENSOR_TENSORDOT_OP)
102    }
103}
104
105/// Backend-explicit einsum methods for dtype-erased concrete tensors.
106///
107/// Implementations are provided for slices and fixed-size arrays of
108/// [`Tensor`] references, so both `inputs.as_slice().einsum(...)` and
109/// `[&lhs, &rhs].einsum(...)` work.
110///
111/// # Examples
112///
113/// ```
114/// use tenferro_cpu::CpuBackend;
115/// use tenferro_einsum::TensorEinsumExt;
116/// use tenferro_tensor::{BackendSessionHost, Tensor};
117///
118/// let lhs = Tensor::from_vec_col_major(vec![2, 3], vec![1.0_f64; 6]).unwrap();
119/// let rhs = Tensor::from_vec_col_major(vec![3, 4], vec![1.0_f64; 12]).unwrap();
120/// let mut backend = CpuBackend::new();
121///
122/// let out = backend.with_backend_session(|session| {
123///     [&lhs, &rhs].einsum("ij,jk->ik", session)
124/// })??;
125/// assert_eq!(out.shape(), &[2, 4]);
126/// # Ok::<(), tenferro_einsum::Error>(())
127/// ```
128pub trait TensorEinsumExt {
129    /// Execute an einsum from string notation.
130    ///
131    /// # Errors
132    ///
133    /// Returns [`Error::InvalidSubscripts`] for malformed notation,
134    /// [`Error::Validation`] with shape, rank, or dtype payloads for an invalid
135    /// contraction, or [`Error::Tensor`] for a typed backend failure.
136    fn einsum(&self, subscripts: &str, session: &mut dyn BackendSession) -> Result<Tensor>;
137
138    /// Execute an einsum from rank-unresolved string/programmatic notation.
139    ///
140    /// # Examples
141    ///
142    /// ```
143    /// use tenferro_einsum::{EinsumAxis, EinsumNotation};
144    /// let notation = EinsumNotation::new(&[&[EinsumAxis::Ellipsis]], &[]);
145    /// assert_eq!(notation.input_count(), 1);
146    /// ```
147    ///
148    /// # Errors
149    ///
150    /// Returns a typed validation or backend error when notation, shapes, or execution are invalid.
151    fn einsum_notation(
152        &self,
153        notation: &EinsumNotation,
154        session: &mut dyn BackendSession,
155    ) -> Result<Tensor>;
156
157    /// Execute an einsum from parsed integer-label subscripts.
158    ///
159    /// # Errors
160    ///
161    /// Returns [`Error::Validation`] with shape, rank, or dtype payloads for an
162    /// invalid contraction, or [`Error::Tensor`] for a typed backend failure.
163    fn einsum_subscripts(
164        &self,
165        subscripts: &EinsumSubscripts,
166        session: &mut dyn BackendSession,
167    ) -> Result<Tensor>;
168}
169
170impl TensorEinsumExt for [&Tensor] {
171    fn einsum(&self, subscripts: &str, session: &mut dyn BackendSession) -> Result<Tensor> {
172        let notation = parse_einsum_notation(subscripts)?;
173        self.einsum_notation(&notation, session)
174    }
175
176    fn einsum_notation(
177        &self,
178        notation: &EinsumNotation,
179        session: &mut dyn BackendSession,
180    ) -> Result<Tensor> {
181        let subscripts = resolve_tensor_notation(self, notation)?;
182        eager_einsum_subscripts_on_session(session, self, &subscripts).map_err(Error::from)
183    }
184
185    fn einsum_subscripts(
186        &self,
187        subscripts: &EinsumSubscripts,
188        session: &mut dyn BackendSession,
189    ) -> Result<Tensor> {
190        let subscripts = Subscripts::from(subscripts);
191        eager_einsum_subscripts_on_session(session, self, &subscripts).map_err(Error::from)
192    }
193}
194
195impl<const N: usize> TensorEinsumExt for [&Tensor; N] {
196    fn einsum(&self, subscripts: &str, session: &mut dyn BackendSession) -> Result<Tensor> {
197        self.as_slice().einsum(subscripts, session)
198    }
199
200    fn einsum_notation(
201        &self,
202        notation: &EinsumNotation,
203        session: &mut dyn BackendSession,
204    ) -> Result<Tensor> {
205        self.as_slice().einsum_notation(notation, session)
206    }
207
208    fn einsum_subscripts(
209        &self,
210        subscripts: &EinsumSubscripts,
211        session: &mut dyn BackendSession,
212    ) -> Result<Tensor> {
213        self.as_slice().einsum_subscripts(subscripts, session)
214    }
215}
216
217/// Backend-explicit preallocated-output einsum methods for dtype-erased tensors.
218pub trait TensorEinsumIntoExt {
219    /// Execute an einsum from string notation into caller-provided output.
220    ///
221    /// # Errors
222    ///
223    /// Returns [`Error::InvalidSubscripts`] for malformed notation,
224    /// [`Error::Validation`] with a shape, rank, or dtype payload when inputs or
225    /// output do not match, or [`Error::Tensor`] for a typed backend failure.
226    fn einsum_into(
227        &self,
228        subscripts: &str,
229        session: &mut dyn BackendSession,
230        out: TensorWrite<'_>,
231    ) -> Result<()>;
232
233    /// Execute an einsum from rank-unresolved notation into caller-provided output.
234    ///
235    /// # Examples
236    ///
237    /// ```
238    /// use tenferro_einsum::{EinsumAxis, EinsumNotation};
239    /// let notation = EinsumNotation::new(&[&[EinsumAxis::Ellipsis]], &[]);
240    /// assert_eq!(notation.input_count(), 1);
241    /// ```
242    ///
243    /// # Errors
244    ///
245    /// Returns a typed validation or backend error when notation, shapes, or output are invalid.
246    fn einsum_into_notation(
247        &self,
248        notation: &EinsumNotation,
249        session: &mut dyn BackendSession,
250        out: TensorWrite<'_>,
251    ) -> Result<()>;
252
253    /// Execute an einsum from parsed integer-label subscripts into caller-provided output.
254    ///
255    /// # Errors
256    ///
257    /// Returns [`Error::Validation`] with a shape, rank, or dtype payload when
258    /// inputs or output do not match, or [`Error::Tensor`] for a typed backend
259    /// failure.
260    fn einsum_into_subscripts(
261        &self,
262        subscripts: &EinsumSubscripts,
263        session: &mut dyn BackendSession,
264        out: TensorWrite<'_>,
265    ) -> Result<()>;
266}
267
268impl TensorEinsumIntoExt for [&Tensor] {
269    fn einsum_into(
270        &self,
271        subscripts: &str,
272        session: &mut dyn BackendSession,
273        out: TensorWrite<'_>,
274    ) -> Result<()> {
275        if let ([lhs, rhs], Some((a, b, c))) = (self, parse_fast_ascii_binary_labels(subscripts)) {
276            let reads = [TensorRead::from_tensor(lhs), TensorRead::from_tensor(rhs)];
277            if let Some((order, config)) = read_binary_dot_config_for_labels(&reads, a, b, c, &out)
278            {
279                return execute_binary_dot_config_read_into(session, &reads, order, &config, out);
280            }
281        }
282        let notation = parse_einsum_notation(subscripts)?;
283        self.einsum_into_notation(&notation, session, out)
284    }
285
286    fn einsum_into_notation(
287        &self,
288        notation: &EinsumNotation,
289        session: &mut dyn BackendSession,
290        out: TensorWrite<'_>,
291    ) -> Result<()> {
292        let subscripts = resolve_tensor_notation(self, notation)?;
293        tensor_einsum_into_subscripts(session, self, &subscripts, out, TENSOR_EINSUM_INTO_OP)
294    }
295
296    fn einsum_into_subscripts(
297        &self,
298        subscripts: &EinsumSubscripts,
299        session: &mut dyn BackendSession,
300        out: TensorWrite<'_>,
301    ) -> Result<()> {
302        let subscripts = Subscripts::from(subscripts);
303        tensor_einsum_into_subscripts(session, self, &subscripts, out, TENSOR_EINSUM_INTO_OP)
304    }
305}
306
307impl<const N: usize> TensorEinsumIntoExt for [&Tensor; N] {
308    fn einsum_into(
309        &self,
310        subscripts: &str,
311        session: &mut dyn BackendSession,
312        out: TensorWrite<'_>,
313    ) -> Result<()> {
314        self.as_slice().einsum_into(subscripts, session, out)
315    }
316
317    fn einsum_into_notation(
318        &self,
319        notation: &EinsumNotation,
320        session: &mut dyn BackendSession,
321        out: TensorWrite<'_>,
322    ) -> Result<()> {
323        self.as_slice().einsum_into_notation(notation, session, out)
324    }
325
326    fn einsum_into_subscripts(
327        &self,
328        subscripts: &EinsumSubscripts,
329        session: &mut dyn BackendSession,
330        out: TensorWrite<'_>,
331    ) -> Result<()> {
332        self.as_slice()
333            .einsum_into_subscripts(subscripts, session, out)
334    }
335}
336
337/// Backend-explicit einsum methods for typed concrete tensors.
338///
339/// The result keeps the same scalar type as the inputs. Mixed dtypes should use
340/// [`TensorEinsumExt`] on dtype-erased [`Tensor`] values instead.
341///
342/// # Examples
343///
344/// ```
345/// use tenferro_cpu::CpuBackend;
346/// use tenferro_einsum::TypedTensorEinsumExt;
347/// use tenferro_tensor::{BackendSessionHost, TypedTensor};
348///
349/// let lhs = TypedTensor::<f64>::from_vec_col_major(vec![2, 3], vec![1.0; 6]).unwrap();
350/// let rhs = TypedTensor::<f64>::from_vec_col_major(vec![3, 4], vec![1.0; 12]).unwrap();
351/// let mut backend = CpuBackend::new();
352///
353/// let out = backend.with_backend_session(|session| {
354///     [&lhs, &rhs].einsum("ij,jk->ik", session)
355/// })??;
356/// assert_eq!(out.shape(), &[2, 4]);
357/// # Ok::<(), tenferro_einsum::Error>(())
358/// ```
359pub trait TypedTensorEinsumExt<T: TensorScalar> {
360    /// Execute an einsum from string notation.
361    ///
362    /// # Errors
363    ///
364    /// Returns [`Error::InvalidSubscripts`] for malformed notation,
365    /// [`Error::Validation`] with shape or rank payloads for an invalid
366    /// contraction, or [`Error::Tensor`] for a typed backend failure.
367    fn einsum(&self, subscripts: &str, session: &mut dyn BackendSession) -> Result<TypedTensor<T>>;
368
369    /// Execute an einsum from rank-unresolved notation.
370    ///
371    /// # Examples
372    ///
373    /// ```
374    /// use tenferro_einsum::{EinsumAxis, EinsumNotation};
375    /// let notation = EinsumNotation::new(&[&[EinsumAxis::Ellipsis]], &[]);
376    /// assert_eq!(notation.input_count(), 1);
377    /// ```
378    ///
379    /// # Errors
380    ///
381    /// Returns a typed validation or backend error when notation, shapes, or execution are invalid.
382    fn einsum_notation(
383        &self,
384        notation: &EinsumNotation,
385        session: &mut dyn BackendSession,
386    ) -> Result<TypedTensor<T>>;
387
388    /// Execute an einsum from parsed integer-label subscripts.
389    ///
390    /// # Errors
391    ///
392    /// Returns [`Error::Validation`] with shape or rank payloads for an invalid
393    /// contraction, or [`Error::Tensor`] for a typed backend failure.
394    fn einsum_subscripts(
395        &self,
396        subscripts: &EinsumSubscripts,
397        session: &mut dyn BackendSession,
398    ) -> Result<TypedTensor<T>>;
399}
400
401impl<T: TensorScalar> TypedTensorEinsumExt<T> for [&TypedTensor<T>] {
402    fn einsum(&self, subscripts: &str, session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
403        let notation = parse_einsum_notation(subscripts)?;
404        self.einsum_notation(&notation, session)
405    }
406
407    fn einsum_notation(
408        &self,
409        notation: &EinsumNotation,
410        session: &mut dyn BackendSession,
411    ) -> Result<TypedTensor<T>> {
412        let subscripts = resolve_typed_notation(self, notation)?;
413        typed_einsum_subscripts(session, self, &subscripts, TYPED_TENSOR_EINSUM_OP)
414    }
415
416    fn einsum_subscripts(
417        &self,
418        subscripts: &EinsumSubscripts,
419        session: &mut dyn BackendSession,
420    ) -> Result<TypedTensor<T>> {
421        let subscripts = Subscripts::from(subscripts);
422        typed_einsum_subscripts(session, self, &subscripts, TYPED_TENSOR_EINSUM_OP)
423    }
424}
425
426impl<T: TensorScalar, const N: usize> TypedTensorEinsumExt<T> for [&TypedTensor<T>; N] {
427    fn einsum(&self, subscripts: &str, session: &mut dyn BackendSession) -> Result<TypedTensor<T>> {
428        self.as_slice().einsum(subscripts, session)
429    }
430
431    fn einsum_notation(
432        &self,
433        notation: &EinsumNotation,
434        session: &mut dyn BackendSession,
435    ) -> Result<TypedTensor<T>> {
436        self.as_slice().einsum_notation(notation, session)
437    }
438
439    fn einsum_subscripts(
440        &self,
441        subscripts: &EinsumSubscripts,
442        session: &mut dyn BackendSession,
443    ) -> Result<TypedTensor<T>> {
444        self.as_slice().einsum_subscripts(subscripts, session)
445    }
446}
447
448/// Backend-explicit einsum methods for typed borrowed views.
449///
450/// The `_read` suffix distinguishes this borrowed-view surface from
451/// [`TypedTensorEinsumExt`], whose unsuffixed methods accept only owned compact
452/// typed tensors.
453///
454/// # Examples
455///
456/// ```
457/// use tenferro_cpu::CpuBackend;
458/// use tenferro_einsum::TypedTensorReadEinsumExt;
459/// use tenferro_tensor::{BackendSessionHost, TypedTensor};
460///
461/// let lhs = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![1.0, 2.0])?;
462/// let rhs = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![3.0, 4.0])?;
463/// let mut backend = CpuBackend::new();
464/// let result = backend.with_backend_session(|session| {
465///     [lhs.as_view(), rhs.as_view()].einsum_read("i,i->", session)
466/// })??;
467/// assert_eq!(result.as_slice()?, &[11.0]);
468/// # Ok::<(), tenferro_einsum::Error>(())
469/// ```
470pub trait TypedTensorReadEinsumExt<T: TensorScalar> {
471    /// Execute an einsum from string notation over typed borrowed views.
472    ///
473    /// # Examples
474    ///
475    /// ```
476    /// use tenferro_cpu::CpuBackend;
477    /// use tenferro_einsum::TypedTensorReadEinsumExt;
478    /// use tenferro_tensor::{BackendSessionHost, TypedTensor};
479    ///
480    /// let input = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![2.0, 3.0])?;
481    /// let mut backend = CpuBackend::new();
482    /// let result = backend.with_backend_session(|session| {
483    ///     [input.as_view()].einsum_read("i->i", session)
484    /// })??;
485    /// assert_eq!(result.as_slice()?, &[2.0, 3.0]);
486    /// # Ok::<(), tenferro_einsum::Error>(())
487    /// ```
488    ///
489    /// # Errors
490    ///
491    /// Returns [`Error::InvalidSubscripts`] for malformed notation,
492    /// [`Error::Validation`] with shape or rank payloads for incompatible
493    /// views, or [`Error::Tensor`] for a typed backend failure.
494    fn einsum_read(
495        &self,
496        subscripts: &str,
497        session: &mut dyn BackendSession,
498    ) -> Result<TypedTensor<T>>;
499
500    /// Execute an einsum from rank-unresolved notation over typed views.
501    ///
502    /// # Examples
503    ///
504    /// ```
505    /// use tenferro_einsum::{EinsumAxis, EinsumNotation};
506    /// let notation = EinsumNotation::new(&[&[EinsumAxis::Ellipsis]], &[]);
507    /// assert_eq!(notation.input_count(), 1);
508    /// ```
509    ///
510    /// # Errors
511    ///
512    /// Returns a typed validation or backend error when notation, views, or execution are invalid.
513    fn einsum_read_notation(
514        &self,
515        notation: &EinsumNotation,
516        session: &mut dyn BackendSession,
517    ) -> Result<TypedTensor<T>>;
518
519    /// Execute an einsum from parsed integer-label subscripts over typed
520    /// borrowed views.
521    ///
522    /// # Examples
523    ///
524    /// ```
525    /// use tenferro_cpu::CpuBackend;
526    /// use tenferro_einsum::{EinsumSubscripts, TypedTensorReadEinsumExt};
527    /// use tenferro_tensor::{BackendSessionHost, TypedTensor};
528    ///
529    /// let input = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![2.0, 3.0])?;
530    /// let subscripts = EinsumSubscripts::new(&[&[0]], &[0]);
531    /// let mut backend = CpuBackend::new();
532    /// let result = backend.with_backend_session(|session| {
533    ///     [input.as_view()].einsum_read_subscripts(&subscripts, session)
534    /// })??;
535    /// assert_eq!(result.as_slice()?, &[2.0, 3.0]);
536    /// # Ok::<(), tenferro_einsum::Error>(())
537    /// ```
538    ///
539    /// # Errors
540    ///
541    /// Returns [`Error::Validation`] with shape or rank payloads for
542    /// incompatible views, or [`Error::Tensor`] for a typed backend failure.
543    fn einsum_read_subscripts(
544        &self,
545        subscripts: &EinsumSubscripts,
546        session: &mut dyn BackendSession,
547    ) -> Result<TypedTensor<T>>;
548}
549
550impl<'a, T: TensorScalar> TypedTensorReadEinsumExt<T> for [TypedTensorView<'a, T>] {
551    fn einsum_read(
552        &self,
553        subscripts: &str,
554        session: &mut dyn BackendSession,
555    ) -> Result<TypedTensor<T>> {
556        let notation = parse_einsum_notation(subscripts)?;
557        self.einsum_read_notation(&notation, session)
558    }
559
560    fn einsum_read_notation(
561        &self,
562        notation: &EinsumNotation,
563        session: &mut dyn BackendSession,
564    ) -> Result<TypedTensor<T>> {
565        let subscripts = resolve_view_notation(self, notation)?;
566        typed_view_einsum_subscripts(session, self, &subscripts, TYPED_TENSOR_READ_EINSUM_OP)
567    }
568
569    fn einsum_read_subscripts(
570        &self,
571        subscripts: &EinsumSubscripts,
572        session: &mut dyn BackendSession,
573    ) -> Result<TypedTensor<T>> {
574        let subscripts = Subscripts::from(subscripts);
575        typed_view_einsum_subscripts(session, self, &subscripts, TYPED_TENSOR_READ_EINSUM_OP)
576    }
577}
578
579impl<'a, T: TensorScalar, const N: usize> TypedTensorReadEinsumExt<T>
580    for [TypedTensorView<'a, T>; N]
581{
582    fn einsum_read(
583        &self,
584        subscripts: &str,
585        session: &mut dyn BackendSession,
586    ) -> Result<TypedTensor<T>> {
587        self.as_slice().einsum_read(subscripts, session)
588    }
589
590    fn einsum_read_notation(
591        &self,
592        notation: &EinsumNotation,
593        session: &mut dyn BackendSession,
594    ) -> Result<TypedTensor<T>> {
595        self.as_slice().einsum_read_notation(notation, session)
596    }
597
598    fn einsum_read_subscripts(
599        &self,
600        subscripts: &EinsumSubscripts,
601        session: &mut dyn BackendSession,
602    ) -> Result<TypedTensor<T>> {
603        self.as_slice().einsum_read_subscripts(subscripts, session)
604    }
605}
606
607/// Backend-explicit preallocated-output einsum methods for typed concrete tensors.
608pub trait TypedTensorEinsumIntoExt<T: TensorScalar> {
609    /// Execute an einsum from string notation into caller-provided typed output.
610    ///
611    /// # Errors
612    ///
613    /// Returns [`Error::InvalidSubscripts`] for malformed notation,
614    /// [`Error::Validation`] with a shape, rank, or dtype payload when inputs or
615    /// output do not match, or [`Error::Tensor`] for a typed backend failure.
616    fn einsum_into<'out, O>(
617        &self,
618        subscripts: &str,
619        session: &mut dyn BackendSession,
620        out: O,
621    ) -> Result<()>
622    where
623        O: Into<TypedTensorWrite<'out, T>>;
624
625    /// Execute an einsum from rank-unresolved notation into typed output.
626    ///
627    /// # Examples
628    ///
629    /// ```
630    /// use tenferro_einsum::{EinsumAxis, EinsumNotation};
631    /// let notation = EinsumNotation::new(&[&[EinsumAxis::Ellipsis]], &[]);
632    /// assert_eq!(notation.input_count(), 1);
633    /// ```
634    ///
635    /// # Errors
636    ///
637    /// Returns a typed validation or backend error when notation, shapes, or output are invalid.
638    fn einsum_into_notation<'out, O>(
639        &self,
640        notation: &EinsumNotation,
641        session: &mut dyn BackendSession,
642        out: O,
643    ) -> Result<()>
644    where
645        O: Into<TypedTensorWrite<'out, T>>;
646
647    /// Execute an einsum from parsed integer-label subscripts into caller-provided typed output.
648    ///
649    /// # Errors
650    ///
651    /// Returns [`Error::Validation`] with a shape, rank, or dtype payload when
652    /// inputs or output do not match, or [`Error::Tensor`] for a typed backend
653    /// failure.
654    fn einsum_into_subscripts<'out, O>(
655        &self,
656        subscripts: &EinsumSubscripts,
657        session: &mut dyn BackendSession,
658        out: O,
659    ) -> Result<()>
660    where
661        O: Into<TypedTensorWrite<'out, T>>;
662}
663
664impl<T: TensorScalar> TypedTensorEinsumIntoExt<T> for [&TypedTensor<T>] {
665    fn einsum_into<'out, O>(
666        &self,
667        subscripts: &str,
668        session: &mut dyn BackendSession,
669        out: O,
670    ) -> Result<()>
671    where
672        O: Into<TypedTensorWrite<'out, T>>,
673    {
674        let out = out.into().into_tensor_write();
675        if let ([lhs, rhs], Some((a, b, c))) = (self, parse_fast_ascii_binary_labels(subscripts)) {
676            let reads = [T::tensor_read(lhs), T::tensor_read(rhs)];
677            if let Some((order, config)) = read_binary_dot_config_for_labels(&reads, a, b, c, &out)
678            {
679                return execute_binary_dot_config_read_into(session, &reads, order, &config, out);
680            }
681        }
682        let notation = parse_einsum_notation(subscripts)?;
683        let subscripts = resolve_typed_notation(self, &notation)?;
684        typed_einsum_into_subscripts(session, self, &subscripts, out, TYPED_TENSOR_EINSUM_INTO_OP)
685    }
686
687    fn einsum_into_notation<'out, O>(
688        &self,
689        notation: &EinsumNotation,
690        session: &mut dyn BackendSession,
691        out: O,
692    ) -> Result<()>
693    where
694        O: Into<TypedTensorWrite<'out, T>>,
695    {
696        let subscripts = resolve_typed_notation(self, notation)?;
697        typed_einsum_into_subscripts(
698            session,
699            self,
700            &subscripts,
701            out.into().into_tensor_write(),
702            TYPED_TENSOR_EINSUM_INTO_OP,
703        )
704    }
705
706    fn einsum_into_subscripts<'out, O>(
707        &self,
708        subscripts: &EinsumSubscripts,
709        session: &mut dyn BackendSession,
710        out: O,
711    ) -> Result<()>
712    where
713        O: Into<TypedTensorWrite<'out, T>>,
714    {
715        let subscripts = Subscripts::from(subscripts);
716        typed_einsum_into_subscripts(
717            session,
718            self,
719            &subscripts,
720            out.into().into_tensor_write(),
721            TYPED_TENSOR_EINSUM_INTO_OP,
722        )
723    }
724}
725
726impl<T: TensorScalar, const N: usize> TypedTensorEinsumIntoExt<T> for [&TypedTensor<T>; N] {
727    fn einsum_into<'out, O>(
728        &self,
729        subscripts: &str,
730        session: &mut dyn BackendSession,
731        out: O,
732    ) -> Result<()>
733    where
734        O: Into<TypedTensorWrite<'out, T>>,
735    {
736        self.as_slice().einsum_into(subscripts, session, out)
737    }
738
739    fn einsum_into_notation<'out, O>(
740        &self,
741        notation: &EinsumNotation,
742        session: &mut dyn BackendSession,
743        out: O,
744    ) -> Result<()>
745    where
746        O: Into<TypedTensorWrite<'out, T>>,
747    {
748        self.as_slice().einsum_into_notation(notation, session, out)
749    }
750
751    fn einsum_into_subscripts<'out, O>(
752        &self,
753        subscripts: &EinsumSubscripts,
754        session: &mut dyn BackendSession,
755        out: O,
756    ) -> Result<()>
757    where
758        O: Into<TypedTensorWrite<'out, T>>,
759    {
760        self.as_slice()
761            .einsum_into_subscripts(subscripts, session, out)
762    }
763}
764
765/// Backend-explicit preallocated-output einsum methods for typed borrowed views.
766///
767/// # Examples
768///
769/// ```
770/// use tenferro_cpu::CpuBackend;
771/// use tenferro_einsum::TypedTensorReadEinsumIntoExt;
772/// use tenferro_tensor::{BackendSessionHost, TypedTensor};
773///
774/// let lhs = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![1.0, 2.0])?;
775/// let rhs = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![3.0, 4.0])?;
776/// let mut output = TypedTensor::<f64>::from_vec_col_major(vec![], vec![0.0])?;
777/// let mut backend = CpuBackend::new();
778/// backend.with_backend_session(|session| {
779///     [lhs.as_view(), rhs.as_view()].einsum_read_into("i,i->", session, &mut output)
780/// })??;
781/// assert_eq!(output.as_slice()?, &[11.0]);
782/// # Ok::<(), tenferro_einsum::Error>(())
783/// ```
784pub trait TypedTensorReadEinsumIntoExt<T: TensorScalar> {
785    /// Execute an einsum from string notation over typed borrowed views into a
786    /// caller-provided typed output.
787    ///
788    /// # Examples
789    ///
790    /// ```
791    /// use tenferro_cpu::CpuBackend;
792    /// use tenferro_einsum::TypedTensorReadEinsumIntoExt;
793    /// use tenferro_tensor::{BackendSessionHost, TypedTensor};
794    ///
795    /// let input = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![2.0, 3.0])?;
796    /// let mut output = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![0.0; 2])?;
797    /// let mut backend = CpuBackend::new();
798    /// backend.with_backend_session(|session| {
799    ///     [input.as_view()].einsum_read_into("i->i", session, &mut output)
800    /// })??;
801    /// assert_eq!(output.as_slice()?, &[2.0, 3.0]);
802    /// # Ok::<(), tenferro_einsum::Error>(())
803    /// ```
804    ///
805    /// # Errors
806    ///
807    /// Returns [`Error::InvalidSubscripts`] for malformed notation,
808    /// [`Error::Validation`] with a shape, rank, or dtype payload when inputs or
809    /// output do not match, or [`Error::Tensor`] for a typed backend failure.
810    fn einsum_read_into<'out, O>(
811        &self,
812        subscripts: &str,
813        session: &mut dyn BackendSession,
814        out: O,
815    ) -> Result<()>
816    where
817        O: Into<TypedTensorWrite<'out, T>>;
818
819    /// Execute an einsum from rank-unresolved notation over typed views into output.
820    ///
821    /// # Examples
822    ///
823    /// ```
824    /// use tenferro_einsum::{EinsumAxis, EinsumNotation};
825    /// let notation = EinsumNotation::new(&[&[EinsumAxis::Ellipsis]], &[]);
826    /// assert_eq!(notation.input_count(), 1);
827    /// ```
828    ///
829    /// # Errors
830    ///
831    /// Returns a typed validation or backend error when notation, views, or output are invalid.
832    fn einsum_read_into_notation<'out, O>(
833        &self,
834        notation: &EinsumNotation,
835        session: &mut dyn BackendSession,
836        out: O,
837    ) -> Result<()>
838    where
839        O: Into<TypedTensorWrite<'out, T>>;
840
841    /// Execute an einsum from parsed integer-label subscripts over typed
842    /// borrowed views into a caller-provided typed output.
843    ///
844    /// # Examples
845    ///
846    /// ```
847    /// use tenferro_cpu::CpuBackend;
848    /// use tenferro_einsum::{EinsumSubscripts, TypedTensorReadEinsumIntoExt};
849    /// use tenferro_tensor::{BackendSessionHost, TypedTensor};
850    ///
851    /// let input = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![2.0, 3.0])?;
852    /// let mut output = TypedTensor::<f64>::from_vec_col_major(vec![2], vec![0.0; 2])?;
853    /// let subscripts = EinsumSubscripts::new(&[&[0]], &[0]);
854    /// let mut backend = CpuBackend::new();
855    /// backend.with_backend_session(|session| {
856    ///     [input.as_view()].einsum_read_into_subscripts(&subscripts, session, &mut output)
857    /// })??;
858    /// assert_eq!(output.as_slice()?, &[2.0, 3.0]);
859    /// # Ok::<(), tenferro_einsum::Error>(())
860    /// ```
861    ///
862    /// # Errors
863    ///
864    /// Returns [`Error::Validation`] with a shape, rank, or dtype payload when
865    /// inputs or output do not match, or [`Error::Tensor`] for a typed backend
866    /// failure.
867    fn einsum_read_into_subscripts<'out, O>(
868        &self,
869        subscripts: &EinsumSubscripts,
870        session: &mut dyn BackendSession,
871        out: O,
872    ) -> Result<()>
873    where
874        O: Into<TypedTensorWrite<'out, T>>;
875}
876
877impl<'a, T: TensorScalar> TypedTensorReadEinsumIntoExt<T> for [TypedTensorView<'a, T>] {
878    fn einsum_read_into<'out, O>(
879        &self,
880        subscripts: &str,
881        session: &mut dyn BackendSession,
882        out: O,
883    ) -> Result<()>
884    where
885        O: Into<TypedTensorWrite<'out, T>>,
886    {
887        let out = out.into().into_tensor_write();
888        if let Some((lhs, rhs, output)) = parse_fast_ascii_binary_labels(subscripts) {
889            if let Some((order, config)) =
890                typed_view_binary_dot_config(self, lhs, rhs, output, &out)
891            {
892                return execute_typed_view_binary_dot_into(session, self, order, &config, out);
893            }
894        }
895        let notation = parse_einsum_notation(subscripts)?;
896        let subscripts = resolve_view_notation(self, &notation)?;
897        typed_view_einsum_into_subscripts(
898            session,
899            self,
900            &subscripts,
901            out,
902            TYPED_TENSOR_READ_EINSUM_INTO_OP,
903        )
904    }
905
906    fn einsum_read_into_notation<'out, O>(
907        &self,
908        notation: &EinsumNotation,
909        session: &mut dyn BackendSession,
910        out: O,
911    ) -> Result<()>
912    where
913        O: Into<TypedTensorWrite<'out, T>>,
914    {
915        let out = out.into().into_tensor_write();
916        if let Some([lhs, rhs, output]) = borrowed_notation_labels(notation) {
917            if let Some((order, config)) =
918                typed_view_binary_dot_config(self, &lhs, &rhs, &output, &out)
919            {
920                return execute_typed_view_binary_dot_into(session, self, order, &config, out);
921            }
922        }
923        let subscripts = resolve_view_notation(self, notation)?;
924        typed_view_einsum_into_subscripts(
925            session,
926            self,
927            &subscripts,
928            out,
929            TYPED_TENSOR_READ_EINSUM_INTO_OP,
930        )
931    }
932
933    fn einsum_read_into_subscripts<'out, O>(
934        &self,
935        subscripts: &EinsumSubscripts,
936        session: &mut dyn BackendSession,
937        out: O,
938    ) -> Result<()>
939    where
940        O: Into<TypedTensorWrite<'out, T>>,
941    {
942        let out = out.into().into_tensor_write();
943        if let [lhs, rhs] = subscripts.inputs.as_slice() {
944            if let Some((order, config)) =
945                typed_view_binary_dot_config(self, lhs, rhs, &subscripts.output, &out)
946            {
947                return execute_typed_view_binary_dot_into(session, self, order, &config, out);
948            }
949        }
950        let subscripts = Subscripts::from(subscripts);
951        typed_view_einsum_into_subscripts(
952            session,
953            self,
954            &subscripts,
955            out,
956            TYPED_TENSOR_READ_EINSUM_INTO_OP,
957        )
958    }
959}
960
961impl<'a, T: TensorScalar, const N: usize> TypedTensorReadEinsumIntoExt<T>
962    for [TypedTensorView<'a, T>; N]
963{
964    fn einsum_read_into<'out, O>(
965        &self,
966        subscripts: &str,
967        session: &mut dyn BackendSession,
968        out: O,
969    ) -> Result<()>
970    where
971        O: Into<TypedTensorWrite<'out, T>>,
972    {
973        self.as_slice().einsum_read_into(subscripts, session, out)
974    }
975
976    fn einsum_read_into_notation<'out, O>(
977        &self,
978        notation: &EinsumNotation,
979        session: &mut dyn BackendSession,
980        out: O,
981    ) -> Result<()>
982    where
983        O: Into<TypedTensorWrite<'out, T>>,
984    {
985        self.as_slice()
986            .einsum_read_into_notation(notation, session, out)
987    }
988
989    fn einsum_read_into_subscripts<'out, O>(
990        &self,
991        subscripts: &EinsumSubscripts,
992        session: &mut dyn BackendSession,
993        out: O,
994    ) -> Result<()>
995    where
996        O: Into<TypedTensorWrite<'out, T>>,
997    {
998        self.as_slice()
999            .einsum_read_into_subscripts(subscripts, session, out)
1000    }
1001}
1002
1003/// Backend-explicit einsum methods for [`TensorRead`] inputs.
1004///
1005/// Use this surface when an input is a borrowed tensor view rather than an
1006/// owned compact [`Tensor`]. The `_read` suffix follows the repository-wide
1007/// convention for APIs that explicitly accept read-oriented borrowed inputs.
1008///
1009/// # Examples
1010///
1011/// ```
1012/// use tenferro_cpu::CpuBackend;
1013/// use tenferro_einsum::TensorReadEinsumExt;
1014/// use tenferro_tensor::{BackendSessionHost, Tensor, TensorRead, TensorView};
1015///
1016/// let shape = [2, 3];
1017/// let data = [1.0_f64; 6];
1018/// let rhs = Tensor::from_vec_col_major(vec![3], vec![1.0_f64; 3]).unwrap();
1019/// let inputs = [
1020///     TensorRead::from_view(TensorView::f64(&shape, &data)?),
1021///     TensorRead::from_tensor(&rhs),
1022/// ];
1023/// let mut backend = CpuBackend::new();
1024///
1025/// let out = backend.with_backend_session(|session| inputs.einsum_read("ij,j->i", session))??;
1026/// assert_eq!(out.shape(), &[2]);
1027/// # Ok::<(), tenferro_einsum::Error>(())
1028/// ```
1029pub trait TensorReadEinsumExt {
1030    /// Execute an einsum from string notation over read-only tensor inputs.
1031    ///
1032    /// # Errors
1033    ///
1034    /// Returns [`Error::InvalidSubscripts`] for malformed notation,
1035    /// [`Error::Validation`] with shape, rank, or dtype payloads for an invalid
1036    /// contraction, or [`Error::Tensor`] for a typed backend failure.
1037    fn einsum_read(&self, subscripts: &str, session: &mut dyn BackendSession) -> Result<Tensor>;
1038
1039    /// Execute an einsum from rank-unresolved notation over read-only inputs.
1040    ///
1041    /// # Examples
1042    ///
1043    /// ```
1044    /// use tenferro_einsum::{EinsumAxis, EinsumNotation};
1045    /// let notation = EinsumNotation::new(&[&[EinsumAxis::Ellipsis]], &[]);
1046    /// assert_eq!(notation.input_count(), 1);
1047    /// ```
1048    ///
1049    /// # Errors
1050    ///
1051    /// Returns a typed validation or backend error when notation, shapes, or execution are invalid.
1052    fn einsum_read_notation(
1053        &self,
1054        notation: &EinsumNotation,
1055        session: &mut dyn BackendSession,
1056    ) -> Result<Tensor>;
1057
1058    /// Execute an einsum from parsed integer-label subscripts over read-only
1059    /// tensor inputs.
1060    ///
1061    /// # Errors
1062    ///
1063    /// Returns [`Error::Validation`] with shape, rank, or dtype payloads for an
1064    /// invalid contraction, or [`Error::Tensor`] for a typed backend failure.
1065    fn einsum_read_subscripts(
1066        &self,
1067        subscripts: &EinsumSubscripts,
1068        session: &mut dyn BackendSession,
1069    ) -> Result<Tensor>;
1070}
1071
1072impl<'a> TensorReadEinsumExt for [TensorRead<'a>] {
1073    fn einsum_read(&self, subscripts: &str, session: &mut dyn BackendSession) -> Result<Tensor> {
1074        let notation = parse_einsum_notation(subscripts)?;
1075        self.einsum_read_notation(&notation, session)
1076    }
1077
1078    fn einsum_read_notation(
1079        &self,
1080        notation: &EinsumNotation,
1081        session: &mut dyn BackendSession,
1082    ) -> Result<Tensor> {
1083        let subscripts = resolve_read_notation(self, notation)?;
1084        eager_einsum_read_subscripts_on_session(session, self, &subscripts).map_err(Error::from)
1085    }
1086
1087    fn einsum_read_subscripts(
1088        &self,
1089        subscripts: &EinsumSubscripts,
1090        session: &mut dyn BackendSession,
1091    ) -> Result<Tensor> {
1092        let subscripts = Subscripts::from(subscripts);
1093        eager_einsum_read_subscripts_on_session(session, self, &subscripts).map_err(Error::from)
1094    }
1095}
1096
1097impl<'a, const N: usize> TensorReadEinsumExt for [TensorRead<'a>; N] {
1098    fn einsum_read(&self, subscripts: &str, session: &mut dyn BackendSession) -> Result<Tensor> {
1099        self.as_slice().einsum_read(subscripts, session)
1100    }
1101
1102    fn einsum_read_notation(
1103        &self,
1104        notation: &EinsumNotation,
1105        session: &mut dyn BackendSession,
1106    ) -> Result<Tensor> {
1107        self.as_slice().einsum_read_notation(notation, session)
1108    }
1109
1110    fn einsum_read_subscripts(
1111        &self,
1112        subscripts: &EinsumSubscripts,
1113        session: &mut dyn BackendSession,
1114    ) -> Result<Tensor> {
1115        self.as_slice().einsum_read_subscripts(subscripts, session)
1116    }
1117}
1118
1119/// Backend-explicit preallocated-output einsum methods for [`TensorRead`] inputs.
1120pub trait TensorReadEinsumIntoExt {
1121    /// Execute an einsum from string notation over read-only inputs into caller-provided output.
1122    ///
1123    /// # Errors
1124    ///
1125    /// Returns [`Error::InvalidSubscripts`] for malformed notation,
1126    /// [`Error::Validation`] with a shape, rank, or dtype payload when inputs or
1127    /// output do not match, or [`Error::Tensor`] for a typed backend failure.
1128    fn einsum_read_into(
1129        &self,
1130        subscripts: &str,
1131        session: &mut dyn BackendSession,
1132        out: TensorWrite<'_>,
1133    ) -> Result<()>;
1134
1135    /// Execute an einsum from rank-unresolved notation over read-only inputs into output.
1136    ///
1137    /// # Examples
1138    ///
1139    /// ```
1140    /// use tenferro_einsum::{EinsumAxis, EinsumNotation};
1141    /// let notation = EinsumNotation::new(&[&[EinsumAxis::Ellipsis]], &[]);
1142    /// assert_eq!(notation.input_count(), 1);
1143    /// ```
1144    ///
1145    /// # Errors
1146    ///
1147    /// Returns a typed validation or backend error when notation, shapes, or output are invalid.
1148    fn einsum_read_into_notation(
1149        &self,
1150        notation: &EinsumNotation,
1151        session: &mut dyn BackendSession,
1152        out: TensorWrite<'_>,
1153    ) -> Result<()>;
1154
1155    /// Execute an einsum from parsed integer-label subscripts over read-only inputs into output.
1156    ///
1157    /// # Errors
1158    ///
1159    /// Returns [`Error::Validation`] with a shape, rank, or dtype payload when
1160    /// inputs or output do not match, or [`Error::Tensor`] for a typed backend
1161    /// failure.
1162    fn einsum_read_into_subscripts(
1163        &self,
1164        subscripts: &EinsumSubscripts,
1165        session: &mut dyn BackendSession,
1166        out: TensorWrite<'_>,
1167    ) -> Result<()>;
1168}
1169
1170impl<'a> TensorReadEinsumIntoExt for [TensorRead<'a>] {
1171    fn einsum_read_into(
1172        &self,
1173        subscripts: &str,
1174        session: &mut dyn BackendSession,
1175        out: TensorWrite<'_>,
1176    ) -> Result<()> {
1177        if let Some((lhs, rhs, output)) = parse_fast_ascii_binary_labels(subscripts) {
1178            if let Some((order, config)) =
1179                read_binary_dot_config_for_labels(self, lhs, rhs, output, &out)
1180            {
1181                return execute_binary_dot_config_read_into(session, self, order, &config, out);
1182            }
1183        }
1184        let notation = parse_einsum_notation(subscripts)?;
1185        self.einsum_read_into_notation(&notation, session, out)
1186    }
1187
1188    fn einsum_read_into_notation(
1189        &self,
1190        notation: &EinsumNotation,
1191        session: &mut dyn BackendSession,
1192        out: TensorWrite<'_>,
1193    ) -> Result<()> {
1194        if let Some([lhs, rhs, output]) = borrowed_notation_labels(notation) {
1195            if let Some((order, config)) =
1196                read_binary_dot_config_for_labels(self, &lhs, &rhs, &output, &out)
1197            {
1198                return execute_binary_dot_config_read_into(session, self, order, &config, out);
1199            }
1200        }
1201        let subscripts = resolve_read_notation(self, notation)?;
1202        tensor_read_einsum_into_subscripts(
1203            session,
1204            self,
1205            &subscripts,
1206            out,
1207            TENSOR_READ_EINSUM_INTO_OP,
1208        )
1209    }
1210
1211    fn einsum_read_into_subscripts(
1212        &self,
1213        subscripts: &EinsumSubscripts,
1214        session: &mut dyn BackendSession,
1215        out: TensorWrite<'_>,
1216    ) -> Result<()> {
1217        if let [lhs, rhs] = subscripts.inputs.as_slice() {
1218            if let Some((order, config)) =
1219                read_binary_dot_config_for_labels(self, lhs, rhs, &subscripts.output, &out)
1220            {
1221                return execute_binary_dot_config_read_into(session, self, order, &config, out);
1222            }
1223        }
1224        let subscripts = Subscripts::from(subscripts);
1225        tensor_read_einsum_into_subscripts(
1226            session,
1227            self,
1228            &subscripts,
1229            out,
1230            TENSOR_READ_EINSUM_INTO_OP,
1231        )
1232    }
1233}
1234
1235impl<'a, const N: usize> TensorReadEinsumIntoExt for [TensorRead<'a>; N] {
1236    fn einsum_read_into(
1237        &self,
1238        subscripts: &str,
1239        session: &mut dyn BackendSession,
1240        out: TensorWrite<'_>,
1241    ) -> Result<()> {
1242        self.as_slice().einsum_read_into(subscripts, session, out)
1243    }
1244
1245    fn einsum_read_into_notation(
1246        &self,
1247        notation: &EinsumNotation,
1248        session: &mut dyn BackendSession,
1249        out: TensorWrite<'_>,
1250    ) -> Result<()> {
1251        self.as_slice()
1252            .einsum_read_into_notation(notation, session, out)
1253    }
1254
1255    fn einsum_read_into_subscripts(
1256        &self,
1257        subscripts: &EinsumSubscripts,
1258        session: &mut dyn BackendSession,
1259        out: TensorWrite<'_>,
1260    ) -> Result<()> {
1261        self.as_slice()
1262            .einsum_read_into_subscripts(subscripts, session, out)
1263    }
1264}
1265
1266/// Prepared concrete einsum plan for repeated executions with fixed input
1267/// dtype and shape metadata.
1268///
1269/// Preparing a plan parses and optimizes the contraction tree once. Execution
1270/// validates the later inputs against the prepared dtype and shape contract,
1271/// then runs the stored tree without re-planning.
1272///
1273/// # Examples
1274///
1275/// ```
1276/// use tenferro_cpu::CpuBackend;
1277/// use tenferro_einsum::ConcreteEinsumPlan;
1278/// use tenferro_tensor::{BackendSessionHost, Tensor};
1279///
1280/// let lhs = Tensor::from_vec_col_major(vec![2, 3], vec![1.0_f64; 6]).unwrap();
1281/// let rhs = Tensor::from_vec_col_major(vec![3, 4], vec![1.0_f64; 12]).unwrap();
1282/// let plan = ConcreteEinsumPlan::prepare([&lhs, &rhs], "ij,jk->ik")?;
1283///
1284/// let mut backend = CpuBackend::new();
1285/// let out = backend
1286///     .with_backend_session(|session| plan.execute([&lhs, &rhs], session))??;
1287/// assert_eq!(out.shape(), &[2, 4]);
1288/// # Ok::<(), tenferro_einsum::Error>(())
1289/// ```
1290#[derive(Debug)]
1291pub struct ConcreteEinsumPlan {
1292    tree: ContractionTree,
1293    inputs: Vec<ConcreteEinsumInputSpec>,
1294    output_shape: Vec<usize>,
1295    binary_dot: Option<crate::binary_dot::BinaryDotPlan>,
1296}
1297
1298impl ConcreteEinsumPlan {
1299    /// Prepare a plan from dtype-erased concrete tensor inputs and string
1300    /// notation.
1301    ///
1302    /// # Errors
1303    ///
1304    /// Returns [`Error::InvalidSubscripts`] for malformed notation,
1305    /// [`Error::Validation`] for rank, shape, or dtype contract violations, or
1306    /// [`Error::Planning`] when no valid contraction tree can be built.
1307    pub fn prepare<'a, I>(inputs: I, subscripts: &str) -> Result<Self>
1308    where
1309        I: AsRef<[&'a Tensor]>,
1310    {
1311        let notation = parse_einsum_notation(subscripts)?;
1312        Self::prepare_notation(inputs, &notation)
1313    }
1314
1315    /// Prepare a plan from dtype-erased concrete tensor inputs and parsed
1316    /// integer-label subscripts.
1317    ///
1318    /// # Errors
1319    ///
1320    /// Returns [`Error::Validation`] for rank, shape, or dtype contract
1321    /// violations, or [`Error::Planning`] when no valid contraction tree can be
1322    /// built.
1323    pub fn prepare_subscripts<'a, I>(inputs: I, subscripts: &EinsumSubscripts) -> Result<Self>
1324    where
1325        I: AsRef<[&'a Tensor]>,
1326    {
1327        let subscripts = Subscripts::from(subscripts);
1328        Self::prepare_subscripts_internal(input_specs(inputs.as_ref()), &subscripts)
1329    }
1330
1331    /// Prepare a plan from rank-unresolved notation and concrete tensor inputs.
1332    ///
1333    /// # Errors
1334    ///
1335    /// Returns [`Error::InvalidSubscripts`] for malformed axis tokens,
1336    /// [`Error::Validation`] for rank, shape, or dtype violations, or
1337    /// [`Error::Planning`] when no contraction tree can be built.
1338    pub fn prepare_notation<'a, I>(inputs: I, notation: &EinsumNotation) -> Result<Self>
1339    where
1340        I: AsRef<[&'a Tensor]>,
1341    {
1342        let inputs = inputs.as_ref();
1343        let subscripts = resolve_tensor_notation(inputs, notation)?;
1344        Self::prepare_subscripts_internal(input_specs(inputs), &subscripts)
1345    }
1346
1347    /// Prepare a plan from typed concrete tensor inputs and string notation.
1348    ///
1349    /// # Errors
1350    ///
1351    /// Returns [`Error::InvalidSubscripts`] for malformed notation,
1352    /// [`Error::Validation`] for rank or shape contract violations, or
1353    /// [`Error::Planning`] when no valid contraction tree can be built.
1354    pub fn prepare_typed<'a, T, I>(inputs: I, subscripts: &str) -> Result<Self>
1355    where
1356        T: TensorScalar,
1357        I: AsRef<[&'a TypedTensor<T>]>,
1358    {
1359        let notation = parse_einsum_notation(subscripts)?;
1360        Self::prepare_typed_notation(inputs, &notation)
1361    }
1362
1363    /// Prepare a plan from typed concrete tensor inputs and parsed integer-label
1364    /// subscripts.
1365    ///
1366    /// # Errors
1367    ///
1368    /// Returns [`Error::Validation`] for rank or shape contract violations, or
1369    /// [`Error::Planning`] when no valid contraction tree can be built.
1370    pub fn prepare_typed_subscripts<'a, T, I>(
1371        inputs: I,
1372        subscripts: &EinsumSubscripts,
1373    ) -> Result<Self>
1374    where
1375        T: TensorScalar,
1376        I: AsRef<[&'a TypedTensor<T>]>,
1377    {
1378        let subscripts = Subscripts::from(subscripts);
1379        Self::prepare_subscripts_internal(typed_input_specs(inputs.as_ref()), &subscripts)
1380    }
1381
1382    /// Prepare a plan from rank-unresolved notation and typed concrete inputs.
1383    ///
1384    /// # Errors
1385    ///
1386    /// Returns [`Error::InvalidSubscripts`] for malformed axis tokens,
1387    /// [`Error::Validation`] for rank or shape violations, or [`Error::Planning`]
1388    /// when no contraction tree can be built.
1389    pub fn prepare_typed_notation<'a, T, I>(inputs: I, notation: &EinsumNotation) -> Result<Self>
1390    where
1391        T: TensorScalar,
1392        I: AsRef<[&'a TypedTensor<T>]>,
1393    {
1394        let inputs = inputs.as_ref();
1395        let subscripts = resolve_typed_notation(inputs, notation)?;
1396        Self::prepare_subscripts_internal(typed_input_specs(inputs), &subscripts)
1397    }
1398
1399    /// Prepare a plan from read-only tensor inputs and string notation.
1400    ///
1401    /// # Errors
1402    ///
1403    /// Returns [`Error::InvalidSubscripts`] for malformed notation,
1404    /// [`Error::Validation`] for rank, shape, or dtype contract violations, or
1405    /// [`Error::Planning`] when no valid contraction tree can be built.
1406    pub fn prepare_read<'a, I>(inputs: I, subscripts: &str) -> Result<Self>
1407    where
1408        I: AsRef<[TensorRead<'a>]>,
1409    {
1410        let notation = parse_einsum_notation(subscripts)?;
1411        Self::prepare_read_notation(inputs, &notation)
1412    }
1413
1414    /// Prepare a plan from read-only tensor inputs and parsed integer-label
1415    /// subscripts.
1416    ///
1417    /// # Errors
1418    ///
1419    /// Returns [`Error::Validation`] for rank, shape, or dtype contract
1420    /// violations, or [`Error::Planning`] when no valid contraction tree can be
1421    /// built.
1422    pub fn prepare_read_subscripts<'a, I>(inputs: I, subscripts: &EinsumSubscripts) -> Result<Self>
1423    where
1424        I: AsRef<[TensorRead<'a>]>,
1425    {
1426        let subscripts = Subscripts::from(subscripts);
1427        Self::prepare_subscripts_internal(read_input_specs(inputs.as_ref()), &subscripts)
1428    }
1429
1430    /// Prepare a plan from rank-unresolved notation and read-only inputs.
1431    ///
1432    /// # Errors
1433    ///
1434    /// Returns [`Error::InvalidSubscripts`] for malformed axis tokens,
1435    /// [`Error::Validation`] for rank, shape, or dtype violations, or
1436    /// [`Error::Planning`] when no contraction tree can be built.
1437    pub fn prepare_read_notation<'a, I>(inputs: I, notation: &EinsumNotation) -> Result<Self>
1438    where
1439        I: AsRef<[TensorRead<'a>]>,
1440    {
1441        let inputs = inputs.as_ref();
1442        let subscripts = resolve_read_notation(inputs, notation)?;
1443        Self::prepare_subscripts_internal(read_input_specs(inputs), &subscripts)
1444    }
1445
1446    /// Number of binary contraction steps in the prepared tree (diagnostics).
1447    pub(crate) fn step_count(&self) -> usize {
1448        self.tree.step_count()
1449    }
1450
1451    /// Execute this plan on dtype-erased concrete tensor inputs inside a
1452    /// borrowed backend session.
1453    ///
1454    /// Validation and the contraction itself run in the caller's `session`;
1455    /// this method never enters a new backend session.
1456    ///
1457    /// # Examples
1458    ///
1459    /// ```
1460    /// use tenferro_cpu::CpuBackend;
1461    /// use tenferro_einsum::ConcreteEinsumPlan;
1462    /// use tenferro_tensor::{BackendSessionHost, Tensor};
1463    ///
1464    /// let lhs = Tensor::from_vec_col_major(vec![2, 3], vec![1.0_f64; 6]).unwrap();
1465    /// let rhs = Tensor::from_vec_col_major(vec![3, 4], vec![1.0_f64; 12]).unwrap();
1466    /// let plan = ConcreteEinsumPlan::prepare([&lhs, &rhs], "ij,jk->ik")?;
1467    ///
1468    /// let mut backend = CpuBackend::new();
1469    /// let out = backend
1470    ///     .with_backend_session(|session| plan.execute([&lhs, &rhs], session))??;
1471    /// assert_eq!(out.shape(), &[2, 4]);
1472    /// # Ok::<(), tenferro_einsum::Error>(())
1473    /// ```
1474    ///
1475    /// # Errors
1476    ///
1477    /// Returns [`Error::Validation`] when inputs violate the prepared rank,
1478    /// shape, or input-count contract, [`Error::Tensor`] with a
1479    /// `tenferro_tensor::Error::Validation` `DTypeMismatch` payload when an
1480    /// input dtype differs from the prepared contract, or [`Error::Tensor`]
1481    /// for a typed backend failure.
1482    pub fn execute<'a, I>(&self, inputs: I, session: &mut dyn BackendSession) -> Result<Tensor>
1483    where
1484        I: AsRef<[&'a Tensor]>,
1485    {
1486        let inputs = inputs.as_ref();
1487        self.validate_tensor_inputs(inputs, PLAN_EXECUTE_OP)?;
1488        eager_einsum_exec(session, inputs, &self.tree).map_err(Error::from)
1489    }
1490
1491    /// Execute this plan on typed concrete tensor inputs inside a borrowed
1492    /// backend session.
1493    ///
1494    /// Validation and the contraction itself run in the caller's `session`;
1495    /// this method never enters a new backend session.
1496    ///
1497    /// # Examples
1498    ///
1499    /// ```
1500    /// use tenferro_cpu::CpuBackend;
1501    /// use tenferro_einsum::ConcreteEinsumPlan;
1502    /// use tenferro_tensor::{BackendSessionHost, TypedTensor};
1503    ///
1504    /// let lhs = TypedTensor::<f64>::from_vec_col_major(vec![2, 3], vec![1.0; 6]).unwrap();
1505    /// let rhs = TypedTensor::<f64>::from_vec_col_major(vec![3, 4], vec![1.0; 12]).unwrap();
1506    /// let plan = ConcreteEinsumPlan::prepare_typed([&lhs, &rhs], "ij,jk->ik")?;
1507    ///
1508    /// let mut backend = CpuBackend::new();
1509    /// let out = backend
1510    ///     .with_backend_session(|session| plan.execute_typed([&lhs, &rhs], session))??;
1511    /// assert_eq!(out.shape(), &[2, 4]);
1512    /// # Ok::<(), tenferro_einsum::Error>(())
1513    /// ```
1514    ///
1515    /// # Errors
1516    ///
1517    /// Returns [`Error::Validation`] when inputs violate the prepared rank,
1518    /// shape, or input-count contract, [`Error::Tensor`] with a
1519    /// `tenferro_tensor::Error::Validation` `DTypeMismatch` payload when the
1520    /// prepared dtype differs from `T` or the eager result dtype, or
1521    /// [`Error::Tensor`] for a typed backend failure.
1522    pub fn execute_typed<'a, T, I>(
1523        &self,
1524        inputs: I,
1525        session: &mut dyn BackendSession,
1526    ) -> Result<TypedTensor<T>>
1527    where
1528        T: TensorScalar,
1529        I: AsRef<[&'a TypedTensor<T>]>,
1530    {
1531        let inputs = inputs.as_ref();
1532        self.validate_typed_inputs(inputs, PLAN_EXECUTE_OP)?;
1533        let reads: Vec<_> = inputs.iter().map(|tensor| T::tensor_read(tensor)).collect();
1534        let result = eager_einsum_exec_read(session, &reads, &self.tree)?;
1535        into_typed_result(result, PLAN_EXECUTE_OP)
1536    }
1537
1538    /// Execute this plan on read-only tensor inputs inside a borrowed backend
1539    /// session.
1540    ///
1541    /// Validation and the contraction itself run in the caller's `session`;
1542    /// this method never enters a new backend session.
1543    ///
1544    /// # Examples
1545    ///
1546    /// ```
1547    /// use tenferro_cpu::CpuBackend;
1548    /// use tenferro_einsum::ConcreteEinsumPlan;
1549    /// use tenferro_tensor::{BackendSessionHost, Tensor, TensorRead};
1550    ///
1551    /// let lhs = Tensor::from_vec_col_major(vec![2, 3], vec![1.0_f64; 6]).unwrap();
1552    /// let rhs = Tensor::from_vec_col_major(vec![3, 4], vec![1.0_f64; 12]).unwrap();
1553    /// let plan = ConcreteEinsumPlan::prepare_read(
1554    ///     [TensorRead::from_tensor(&lhs), TensorRead::from_tensor(&rhs)],
1555    ///     "ij,jk->ik",
1556    /// )?;
1557    ///
1558    /// let mut backend = CpuBackend::new();
1559    /// let reads = [TensorRead::from_tensor(&lhs), TensorRead::from_tensor(&rhs)];
1560    /// let out = backend
1561    ///     .with_backend_session(|session| plan.execute_read(reads, session))??;
1562    /// assert_eq!(out.shape(), &[2, 4]);
1563    /// # Ok::<(), tenferro_einsum::Error>(())
1564    /// ```
1565    ///
1566    /// # Errors
1567    ///
1568    /// Returns [`Error::Validation`] when inputs violate the prepared rank,
1569    /// shape, or input-count contract, [`Error::Tensor`] with a
1570    /// `tenferro_tensor::Error::Validation` `DTypeMismatch` payload when an
1571    /// input dtype differs from the prepared contract, or [`Error::Tensor`]
1572    /// for a typed backend failure.
1573    pub fn execute_read<'a, I>(&self, inputs: I, session: &mut dyn BackendSession) -> Result<Tensor>
1574    where
1575        I: AsRef<[TensorRead<'a>]>,
1576    {
1577        let inputs = inputs.as_ref();
1578        self.validate_read_inputs(inputs, PLAN_EXECUTE_OP)?;
1579        eager_einsum_exec_read(session, inputs, &self.tree).map_err(Error::from)
1580    }
1581
1582    /// Execute this plan on dtype-erased concrete tensor inputs into
1583    /// caller-provided output inside a borrowed backend session.
1584    ///
1585    /// Validation and the contraction itself run in the caller's `session`;
1586    /// this method never enters a new backend session.
1587    ///
1588    /// # Examples
1589    ///
1590    /// ```
1591    /// use tenferro_cpu::CpuBackend;
1592    /// use tenferro_einsum::ConcreteEinsumPlan;
1593    /// use tenferro_tensor::{BackendSessionHost, Tensor, TensorWrite};
1594    ///
1595    /// let lhs = Tensor::from_vec_col_major(vec![2, 3], vec![1.0_f64; 6]).unwrap();
1596    /// let rhs = Tensor::from_vec_col_major(vec![3, 4], vec![1.0_f64; 12]).unwrap();
1597    /// let plan = ConcreteEinsumPlan::prepare([&lhs, &rhs], "ij,jk->ik")?;
1598    ///
1599    /// let mut backend = CpuBackend::new();
1600    /// let mut out = Tensor::from_vec_col_major(vec![2, 4], vec![0.0_f64; 8]).unwrap();
1601    /// backend.with_backend_session(|session| {
1602    ///     plan.execute_into(
1603    ///         [&lhs, &rhs],
1604    ///         session,
1605    ///         TensorWrite::from_tensor(&mut out),
1606    ///     )
1607    /// })??;
1608    /// assert_eq!(out.as_slice::<f64>()?, vec![3.0_f64; 8].as_slice());
1609    /// # Ok::<(), tenferro_einsum::Error>(())
1610    /// ```
1611    ///
1612    /// # Errors
1613    ///
1614    /// Returns [`Error::Validation`] for input or output rank, shape, or
1615    /// input-count contract violations, [`Error::Tensor`] with a
1616    /// `tenferro_tensor::Error::Validation` `DTypeMismatch` payload for dtype
1617    /// mismatches, or [`Error::Tensor`] for a typed backend failure.
1618    pub fn execute_into<'a, I>(
1619        &self,
1620        inputs: I,
1621        session: &mut dyn BackendSession,
1622        out: TensorWrite<'_>,
1623    ) -> Result<()>
1624    where
1625        I: AsRef<[&'a Tensor]>,
1626    {
1627        let inputs = inputs.as_ref();
1628        self.validate_tensor_inputs(inputs, PLAN_EXECUTE_OP)?;
1629        self.validate_cached_output(&out, PLAN_EXECUTE_OP)?;
1630        if let Some(binary_dot) = &self.binary_dot {
1631            if let [lhs, rhs] = inputs {
1632                let reads = [TensorRead::from_tensor(lhs), TensorRead::from_tensor(rhs)];
1633                return execute_binary_dot_read_into(session, &reads, binary_dot, out)
1634                    .map_err(Error::from);
1635            }
1636        }
1637        let reads: Vec<_> = inputs
1638            .iter()
1639            .map(|tensor| TensorRead::from_tensor(tensor))
1640            .collect();
1641        eager_einsum_exec_read_into(session, &reads, &self.tree, out).map_err(Error::from)
1642    }
1643
1644    /// Execute this plan on typed concrete tensor inputs into caller-provided
1645    /// output inside a borrowed backend session.
1646    ///
1647    /// Validation and the contraction itself run in the caller's `session`;
1648    /// this method never enters a new backend session.
1649    ///
1650    /// # Examples
1651    ///
1652    /// ```
1653    /// use tenferro_cpu::CpuBackend;
1654    /// use tenferro_einsum::ConcreteEinsumPlan;
1655    /// use tenferro_tensor::{BackendSessionHost, TypedTensor};
1656    ///
1657    /// let lhs = TypedTensor::<f64>::from_vec_col_major(vec![2, 3], vec![1.0; 6]).unwrap();
1658    /// let rhs = TypedTensor::<f64>::from_vec_col_major(vec![3, 4], vec![1.0; 12]).unwrap();
1659    /// let plan = ConcreteEinsumPlan::prepare_typed([&lhs, &rhs], "ij,jk->ik")?;
1660    ///
1661    /// let mut backend = CpuBackend::new();
1662    /// let mut out = TypedTensor::<f64>::from_vec_col_major(vec![2, 4], vec![0.0; 8]).unwrap();
1663    /// backend.with_backend_session(|session| {
1664    ///     plan.execute_typed_into([&lhs, &rhs], session, &mut out)
1665    /// })??;
1666    /// assert_eq!(out.as_slice()?, vec![3.0_f64; 8].as_slice());
1667    /// # Ok::<(), tenferro_einsum::Error>(())
1668    /// ```
1669    ///
1670    /// # Errors
1671    ///
1672    /// Returns [`Error::Validation`] for input or output rank, shape, or
1673    /// input-count contract violations, [`Error::Tensor`] with a
1674    /// `tenferro_tensor::Error::Validation` `DTypeMismatch` payload when the
1675    /// prepared dtype differs from `T` or the output dtype, or
1676    /// [`Error::Tensor`] for a typed backend failure.
1677    pub fn execute_typed_into<'a, 'out, T, I, O>(
1678        &self,
1679        inputs: I,
1680        session: &mut dyn BackendSession,
1681        out: O,
1682    ) -> Result<()>
1683    where
1684        T: TensorScalar,
1685        I: AsRef<[&'a TypedTensor<T>]>,
1686        O: Into<TypedTensorWrite<'out, T>>,
1687    {
1688        let inputs = inputs.as_ref();
1689        self.validate_typed_inputs(inputs, PLAN_EXECUTE_OP)?;
1690        let out = out.into().into_tensor_write();
1691        self.validate_cached_output(&out, PLAN_EXECUTE_OP)?;
1692        if let Some(binary_dot) = &self.binary_dot {
1693            if let [lhs, rhs] = inputs {
1694                let reads = [T::tensor_read(lhs), T::tensor_read(rhs)];
1695                return execute_binary_dot_read_into(session, &reads, binary_dot, out)
1696                    .map_err(Error::from);
1697            }
1698        }
1699        let reads: Vec<_> = inputs.iter().map(|tensor| T::tensor_read(tensor)).collect();
1700        eager_einsum_exec_read_into(session, &reads, &self.tree, out).map_err(Error::from)
1701    }
1702
1703    /// Execute this plan on read-only tensor inputs into caller-provided output
1704    /// inside a borrowed backend session.
1705    ///
1706    /// Validation and the contraction itself run in the caller's `session`;
1707    /// this method never enters a new backend session.
1708    ///
1709    /// # Examples
1710    ///
1711    /// ```
1712    /// use tenferro_cpu::CpuBackend;
1713    /// use tenferro_einsum::ConcreteEinsumPlan;
1714    /// use tenferro_tensor::{BackendSessionHost, Tensor, TensorRead, TensorWrite};
1715    ///
1716    /// let lhs = Tensor::from_vec_col_major(vec![2, 3], vec![1.0_f64; 6]).unwrap();
1717    /// let rhs = Tensor::from_vec_col_major(vec![3, 4], vec![1.0_f64; 12]).unwrap();
1718    /// let plan = ConcreteEinsumPlan::prepare_read(
1719    ///     [TensorRead::from_tensor(&lhs), TensorRead::from_tensor(&rhs)],
1720    ///     "ij,jk->ik",
1721    /// )?;
1722    ///
1723    /// let mut backend = CpuBackend::new();
1724    /// let mut out = Tensor::from_vec_col_major(vec![2, 4], vec![0.0_f64; 8]).unwrap();
1725    /// let reads = [TensorRead::from_tensor(&lhs), TensorRead::from_tensor(&rhs)];
1726    /// backend.with_backend_session(|session| {
1727    ///     plan.execute_read_into(reads, session, TensorWrite::from_tensor(&mut out))
1728    /// })??;
1729    /// assert_eq!(out.as_slice::<f64>()?, vec![3.0_f64; 8].as_slice());
1730    /// # Ok::<(), tenferro_einsum::Error>(())
1731    /// ```
1732    ///
1733    /// # Errors
1734    ///
1735    /// Returns [`Error::Validation`] for input or output rank, shape, or
1736    /// input-count contract violations, [`Error::Tensor`] with a
1737    /// `tenferro_tensor::Error::Validation` `DTypeMismatch` payload for dtype
1738    /// mismatches, or [`Error::Tensor`] for a typed backend failure.
1739    pub fn execute_read_into<'a, I>(
1740        &self,
1741        inputs: I,
1742        session: &mut dyn BackendSession,
1743        out: TensorWrite<'_>,
1744    ) -> Result<()>
1745    where
1746        I: AsRef<[TensorRead<'a>]>,
1747    {
1748        let inputs = inputs.as_ref();
1749        self.validate_read_inputs(inputs, PLAN_EXECUTE_OP)?;
1750        self.validate_cached_output(&out, PLAN_EXECUTE_OP)?;
1751        if let Some(binary_dot) = &self.binary_dot {
1752            return execute_binary_dot_read_into(session, inputs, binary_dot, out)
1753                .map_err(Error::from);
1754        }
1755        eager_einsum_exec_read_into(session, inputs, &self.tree, out).map_err(Error::from)
1756    }
1757
1758    /// Execute this plan on read-only inputs with scaled output accumulation
1759    /// inside a borrowed backend session.
1760    ///
1761    /// `accumulation` follows the dot-general contract:
1762    /// `out = alpha * einsum(inputs) + beta * out`.
1763    ///
1764    /// # Examples
1765    ///
1766    /// ```
1767    /// use tenferro_cpu::CpuBackend;
1768    /// use tenferro_einsum::ConcreteEinsumPlan;
1769    /// use tenferro_tensor::{
1770    ///     BackendSessionHost, DotGeneralAccumulation, DType, Tensor, TensorRead, TensorWrite,
1771    /// };
1772    ///
1773    /// let lhs = Tensor::from_vec_col_major(vec![1], vec![2.0_f64])?;
1774    /// let rhs = Tensor::from_vec_col_major(vec![1], vec![3.0_f64])?;
1775    /// let mut out = Tensor::from_vec_col_major(vec![], vec![1.0_f64])?;
1776    /// let plan = ConcreteEinsumPlan::prepare([&lhs, &rhs], "i,i->")?;
1777    /// let mut backend = CpuBackend::new();
1778    /// backend.with_backend_session(|session| {
1779    ///     plan.execute_read_into_accum(
1780    ///         [TensorRead::from_tensor(&lhs), TensorRead::from_tensor(&rhs)],
1781    ///         session,
1782    ///         DotGeneralAccumulation::add_to(DType::F64)?,
1783    ///         TensorWrite::from_tensor(&mut out),
1784    ///     )
1785    /// })??;
1786    /// assert_eq!(out.as_slice::<f64>()?, &[7.0]);
1787    /// # Ok::<(), tenferro_einsum::Error>(())
1788    /// ```
1789    ///
1790    /// # Errors
1791    ///
1792    /// Returns [`Error::Validation`] for input or output rank, shape, or
1793    /// input-count contract violations, [`Error::Tensor`] with a
1794    /// `tenferro_tensor::Error::Validation` `DTypeMismatch` payload for dtype
1795    /// mismatches, [`Error::Numerical`] for an invalid accumulation, or
1796    /// [`Error::Tensor`] for a typed backend failure.
1797    pub fn execute_read_into_accum<'a, I>(
1798        &self,
1799        inputs: I,
1800        session: &mut dyn BackendSession,
1801        accumulation: DotGeneralAccumulation,
1802        out: TensorWrite<'_>,
1803    ) -> Result<()>
1804    where
1805        I: AsRef<[TensorRead<'a>]>,
1806    {
1807        let inputs = inputs.as_ref();
1808        self.validate_read_inputs(inputs, PLAN_EXECUTE_OP)?;
1809        self.validate_cached_output(&out, PLAN_EXECUTE_OP)?;
1810        if let Some(binary_dot) = &self.binary_dot {
1811            return execute_binary_dot_read_into_accum(
1812                session,
1813                inputs,
1814                binary_dot,
1815                accumulation,
1816                out,
1817            )
1818            .map_err(Error::from);
1819        }
1820        eager_einsum_exec_read_into_accum(session, inputs, &self.tree, accumulation, out)
1821            .map_err(Error::from)
1822    }
1823
1824    fn prepare_subscripts_internal(
1825        inputs: Vec<ConcreteEinsumInputSpec>,
1826        subscripts: &Subscripts,
1827    ) -> Result<Self> {
1828        let shapes: Vec<&[usize]> = inputs.iter().map(|input| input.shape.as_slice()).collect();
1829        let binary_dot = binary_dot_plan_for_shapes(&shapes, subscripts);
1830        let tree = plan_subscripts(subscripts, &shapes)?;
1831        let output_shape = tree.output_shape();
1832        Ok(Self {
1833            tree,
1834            inputs,
1835            output_shape,
1836            binary_dot,
1837        })
1838    }
1839
1840    fn validate_tensor_inputs(&self, actual: &[&Tensor], op: &'static str) -> Result<()> {
1841        self.validate_input_metadata(
1842            actual.iter().map(|tensor| (tensor.dtype(), tensor.shape())),
1843            op,
1844        )
1845    }
1846
1847    fn validate_typed_inputs<T: TensorScalar>(
1848        &self,
1849        actual: &[&TypedTensor<T>],
1850        op: &'static str,
1851    ) -> Result<()> {
1852        self.validate_input_metadata(actual.iter().map(|tensor| (T::dtype(), tensor.shape())), op)
1853    }
1854
1855    fn validate_read_inputs(&self, actual: &[TensorRead<'_>], op: &'static str) -> Result<()> {
1856        self.validate_input_metadata(
1857            actual.iter().map(|tensor| (tensor.dtype(), tensor.shape())),
1858            op,
1859        )
1860    }
1861
1862    fn validate_input_metadata<'a>(
1863        &self,
1864        actual: impl ExactSizeIterator<Item = (DType, &'a [usize])>,
1865        op: &'static str,
1866    ) -> Result<()> {
1867        if actual.len() != self.inputs.len() {
1868            return Err(Error::invalid_argument(
1869                op,
1870                "inputs",
1871                format!(
1872                    "prepared einsum expects {} inputs, got {}",
1873                    self.inputs.len(),
1874                    actual.len()
1875                ),
1876            ));
1877        }
1878        for (expected, (actual_dtype, actual_shape)) in self.inputs.iter().zip(actual) {
1879            if expected.dtype != actual_dtype {
1880                return Err(Error::dtype_mismatch(op, expected.dtype, actual_dtype));
1881            }
1882            if expected.shape != actual_shape {
1883                return Err(Error::shape_mismatch(
1884                    op,
1885                    expected.shape.clone(),
1886                    actual_shape.to_vec(),
1887                ));
1888            }
1889        }
1890        Ok(())
1891    }
1892
1893    fn validate_cached_output(&self, out: &TensorWrite<'_>, op: &'static str) -> Result<()> {
1894        let dtype = self
1895            .inputs
1896            .first()
1897            .map(|input| input.dtype)
1898            .ok_or_else(|| {
1899                Error::invalid_argument(op, "inputs", "einsum requires at least one input tensor")
1900            })?;
1901        for input in &self.inputs[1..] {
1902            if input.dtype != dtype {
1903                return Err(Error::dtype_mismatch(op, dtype, input.dtype));
1904            }
1905        }
1906        if out.dtype() != dtype {
1907            return Err(Error::dtype_mismatch(op, dtype, out.dtype()));
1908        }
1909        if out.shape() != self.output_shape {
1910            return Err(Error::shape_mismatch(
1911                op,
1912                out.shape().to_vec(),
1913                self.output_shape.clone(),
1914            ));
1915        }
1916        Ok(())
1917    }
1918}
1919
1920#[derive(Clone, Debug)]
1921struct ConcreteEinsumInputSpec {
1922    dtype: DType,
1923    shape: Vec<usize>,
1924}
1925
1926fn resolve_shapes(notation: &EinsumNotation, shapes: Vec<&[usize]>) -> Result<Subscripts> {
1927    resolve_einsum_notation(notation, &shapes)
1928}
1929
1930fn resolve_tensor_notation(inputs: &[&Tensor], notation: &EinsumNotation) -> Result<Subscripts> {
1931    resolve_shapes(
1932        notation,
1933        inputs.iter().map(|tensor| tensor.shape()).collect(),
1934    )
1935}
1936
1937fn resolve_typed_notation<T: TensorScalar>(
1938    inputs: &[&TypedTensor<T>],
1939    notation: &EinsumNotation,
1940) -> Result<Subscripts> {
1941    resolve_shapes(
1942        notation,
1943        inputs.iter().map(|tensor| tensor.shape()).collect(),
1944    )
1945}
1946
1947fn resolve_view_notation<'a, T: TensorScalar>(
1948    inputs: &[TypedTensorView<'a, T>],
1949    notation: &EinsumNotation,
1950) -> Result<Subscripts> {
1951    resolve_shapes(notation, inputs.iter().map(|view| view.shape()).collect())
1952}
1953
1954fn resolve_read_notation<'a>(
1955    inputs: &[TensorRead<'a>],
1956    notation: &EinsumNotation,
1957) -> Result<Subscripts> {
1958    resolve_shapes(notation, inputs.iter().map(|input| input.shape()).collect())
1959}
1960
1961fn input_specs(inputs: &[&Tensor]) -> Vec<ConcreteEinsumInputSpec> {
1962    inputs
1963        .iter()
1964        .map(|tensor| ConcreteEinsumInputSpec {
1965            dtype: tensor.dtype(),
1966            shape: tensor.shape().to_vec(),
1967        })
1968        .collect()
1969}
1970
1971fn typed_input_specs<T: TensorScalar>(inputs: &[&TypedTensor<T>]) -> Vec<ConcreteEinsumInputSpec> {
1972    inputs
1973        .iter()
1974        .map(|tensor| ConcreteEinsumInputSpec {
1975            dtype: T::dtype(),
1976            shape: tensor.shape().to_vec(),
1977        })
1978        .collect()
1979}
1980
1981fn read_input_specs(inputs: &[TensorRead<'_>]) -> Vec<ConcreteEinsumInputSpec> {
1982    inputs
1983        .iter()
1984        .map(|tensor| ConcreteEinsumInputSpec {
1985            dtype: tensor.dtype(),
1986            shape: tensor.shape().to_vec(),
1987        })
1988        .collect()
1989}
1990
1991fn typed_view_einsum_subscripts<T: TensorScalar>(
1992    session: &mut dyn BackendSession,
1993    inputs: &[TypedTensorView<'_, T>],
1994    subscripts: &Subscripts,
1995    op: &'static str,
1996) -> Result<TypedTensor<T>> {
1997    let reads: Vec<_> = inputs
1998        .iter()
1999        .cloned()
2000        .map(|view| TensorRead::from_view(T::tensor_view(view)))
2001        .collect();
2002    let plan =
2003        ConcreteEinsumPlan::prepare_subscripts_internal(read_input_specs(&reads), subscripts)?;
2004    let result = plan.execute_read(&reads, session)?;
2005    into_typed_result(result, op)
2006}
2007
2008fn read_binary_dot_config_for_labels<L: Copy + PartialEq>(
2009    inputs: &[TensorRead<'_>],
2010    lhs_labels: &[L],
2011    rhs_labels: &[L],
2012    output_labels: &[L],
2013    out: &TensorWrite<'_>,
2014) -> Option<(BinaryDotOperandOrder, tenferro_tensor::DotGeneralConfig)> {
2015    if inputs.len() != 2
2016        || inputs[0].dtype() != inputs[1].dtype()
2017        || out.dtype() != inputs[0].dtype()
2018    {
2019        return None;
2020    }
2021    binary_dot_config_for_into(
2022        inputs[0].shape(),
2023        inputs[1].shape(),
2024        lhs_labels,
2025        rhs_labels,
2026        output_labels,
2027        out.shape(),
2028    )
2029}
2030
2031fn read_binary_dot_config(
2032    inputs: &[TensorRead<'_>],
2033    subscripts: &Subscripts,
2034    out: &TensorWrite<'_>,
2035) -> Option<(BinaryDotOperandOrder, tenferro_tensor::DotGeneralConfig)> {
2036    let [lhs, rhs] = subscripts.inputs.as_slice() else {
2037        return None;
2038    };
2039    read_binary_dot_config_for_labels(inputs, lhs, rhs, &subscripts.output, out)
2040}
2041
2042fn execute_binary_dot_config_read_into(
2043    session: &mut dyn BackendSession,
2044    inputs: &[TensorRead<'_>],
2045    order: BinaryDotOperandOrder,
2046    config: &tenferro_tensor::DotGeneralConfig,
2047    out: TensorWrite<'_>,
2048) -> Result<()> {
2049    let (lhs, rhs) = match order {
2050        BinaryDotOperandOrder::Original => (0, 1),
2051        BinaryDotOperandOrder::Swapped => (1, 0),
2052    };
2053    session
2054        .dot_general_read_into(inputs[lhs].clone(), inputs[rhs].clone(), config, out)
2055        .map_err(Error::from)
2056}
2057
2058fn tensor_einsum_into_subscripts(
2059    session: &mut dyn BackendSession,
2060    inputs: &[&Tensor],
2061    subscripts: &Subscripts,
2062    out: TensorWrite<'_>,
2063    op: &'static str,
2064) -> Result<()> {
2065    if inputs.len() == 2 {
2066        let reads = [
2067            TensorRead::from_tensor(inputs[0]),
2068            TensorRead::from_tensor(inputs[1]),
2069        ];
2070        if let Some((order, config)) = read_binary_dot_config(&reads, subscripts, &out) {
2071            return execute_binary_dot_config_read_into(session, &reads, order, &config, out);
2072        }
2073    }
2074    let plan = ConcreteEinsumPlan::prepare_subscripts_internal(input_specs(inputs), subscripts)?;
2075    validate_output(&plan.inputs, &plan.tree, &out, op)?;
2076    plan.execute_into(inputs, session, out)
2077}
2078
2079fn parse_fast_ascii_binary_labels(notation: &str) -> Option<(&[u8], &[u8], &[u8])> {
2080    let (terms, output) = notation.split_once("->")?;
2081    let (lhs, rhs) = terms.split_once(',')?;
2082    // Borrow labels directly. Separators or unsupported labels in any term send
2083    // the entire expression through the canonical parser, including its errors.
2084    [lhs, rhs, output]
2085        .iter()
2086        .all(|term| term.bytes().all(|byte| byte.is_ascii_alphabetic()))
2087        .then_some((lhs.as_bytes(), rhs.as_bytes(), output.as_bytes()))
2088}
2089
2090fn borrowed_notation_labels(notation: &EinsumNotation) -> Option<[SmallVec<[u32; 8]>; 3]> {
2091    let [lhs, rhs] = notation.inputs.as_slice() else {
2092        return None;
2093    };
2094    let labels = |axes: &[crate::EinsumAxis]| {
2095        axes.iter()
2096            .map(|axis| match axis {
2097                crate::EinsumAxis::Label(label) => Some(*label),
2098                crate::EinsumAxis::Ellipsis => None,
2099            })
2100            .collect::<Option<SmallVec<[u32; 8]>>>()
2101    };
2102    Some([labels(lhs)?, labels(rhs)?, labels(&notation.output)?])
2103}
2104
2105fn typed_view_binary_dot_config<T: TensorScalar, L: Copy + PartialEq>(
2106    inputs: &[TypedTensorView<'_, T>],
2107    lhs_labels: &[L],
2108    rhs_labels: &[L],
2109    output_labels: &[L],
2110    out: &TensorWrite<'_>,
2111) -> Option<(BinaryDotOperandOrder, tenferro_tensor::DotGeneralConfig)> {
2112    if inputs.len() != 2 || out.dtype() != T::dtype() {
2113        return None;
2114    }
2115    binary_dot_config_for_into(
2116        inputs[0].shape(),
2117        inputs[1].shape(),
2118        lhs_labels,
2119        rhs_labels,
2120        output_labels,
2121        out.shape(),
2122    )
2123}
2124
2125fn execute_typed_view_binary_dot_into<T: TensorScalar>(
2126    session: &mut dyn BackendSession,
2127    inputs: &[TypedTensorView<'_, T>],
2128    order: BinaryDotOperandOrder,
2129    config: &tenferro_tensor::DotGeneralConfig,
2130    out: TensorWrite<'_>,
2131) -> Result<()> {
2132    let (lhs, rhs) = match order {
2133        BinaryDotOperandOrder::Original => (0, 1),
2134        BinaryDotOperandOrder::Swapped => (1, 0),
2135    };
2136    let lhs = TensorRead::from_view(T::tensor_view(inputs[lhs].clone()));
2137    let rhs = TensorRead::from_view(T::tensor_view(inputs[rhs].clone()));
2138    session
2139        .dot_general_read_into(lhs, rhs, config, out)
2140        .map_err(Error::from)
2141}
2142
2143fn typed_view_einsum_into_subscripts<T: TensorScalar>(
2144    session: &mut dyn BackendSession,
2145    inputs: &[TypedTensorView<'_, T>],
2146    subscripts: &Subscripts,
2147    out: TensorWrite<'_>,
2148    op: &'static str,
2149) -> Result<()> {
2150    let reads: Vec<_> = inputs
2151        .iter()
2152        .cloned()
2153        .map(|view| TensorRead::from_view(T::tensor_view(view)))
2154        .collect();
2155    let plan =
2156        ConcreteEinsumPlan::prepare_subscripts_internal(read_input_specs(&reads), subscripts)?;
2157    validate_output(&plan.inputs, &plan.tree, &out, op)?;
2158    plan.execute_read_into(&reads, session, out)
2159}
2160
2161fn typed_einsum_into_subscripts<T: TensorScalar>(
2162    session: &mut dyn BackendSession,
2163    inputs: &[&TypedTensor<T>],
2164    subscripts: &Subscripts,
2165    out: TensorWrite<'_>,
2166    op: &'static str,
2167) -> Result<()> {
2168    if inputs.len() == 2 {
2169        if let [lhs, rhs] = subscripts.inputs.as_slice() {
2170            if let Some((order, config)) = binary_dot_config_for_into(
2171                inputs[0].shape(),
2172                inputs[1].shape(),
2173                lhs,
2174                rhs,
2175                &subscripts.output,
2176                out.shape(),
2177            ) {
2178                if out.dtype() == T::dtype() {
2179                    let reads = [T::tensor_read(inputs[0]), T::tensor_read(inputs[1])];
2180                    return execute_binary_dot_config_read_into(
2181                        session, &reads, order, &config, out,
2182                    );
2183                }
2184            }
2185        }
2186    }
2187    let reads: Vec<_> = inputs.iter().map(|tensor| T::tensor_read(tensor)).collect();
2188    let plan =
2189        ConcreteEinsumPlan::prepare_subscripts_internal(read_input_specs(&reads), subscripts)?;
2190    validate_output(&plan.inputs, &plan.tree, &out, op)?;
2191    plan.execute_read_into(&reads, session, out)
2192}
2193
2194fn tensor_read_einsum_into_subscripts(
2195    session: &mut dyn BackendSession,
2196    inputs: &[TensorRead<'_>],
2197    subscripts: &Subscripts,
2198    out: TensorWrite<'_>,
2199    op: &'static str,
2200) -> Result<()> {
2201    if let Some((order, config)) = read_binary_dot_config(inputs, subscripts, &out) {
2202        return execute_binary_dot_config_read_into(session, inputs, order, &config, out);
2203    }
2204    let plan =
2205        ConcreteEinsumPlan::prepare_subscripts_internal(read_input_specs(inputs), subscripts)?;
2206    validate_output(&plan.inputs, &plan.tree, &out, op)?;
2207    plan.execute_read_into(inputs, session, out)
2208}
2209
2210fn validate_output(
2211    inputs: &[ConcreteEinsumInputSpec],
2212    tree: &ContractionTree,
2213    out: &TensorWrite<'_>,
2214    op: &'static str,
2215) -> Result<()> {
2216    let expected = output_spec(inputs, tree, op)?;
2217    if out.dtype() != expected.dtype {
2218        return Err(Error::dtype_mismatch(op, expected.dtype, out.dtype()));
2219    }
2220    if out.shape() != expected.shape.as_slice() {
2221        return Err(Error::shape_mismatch(
2222            op,
2223            out.shape().to_vec(),
2224            expected.shape.clone(),
2225        ));
2226    }
2227    Ok(())
2228}
2229
2230fn output_spec(
2231    inputs: &[ConcreteEinsumInputSpec],
2232    tree: &ContractionTree,
2233    op: &'static str,
2234) -> Result<ConcreteEinsumInputSpec> {
2235    let dtype = inputs
2236        .first()
2237        .ok_or_else(|| {
2238            Error::invalid_argument(op, "inputs", "einsum requires at least one input tensor")
2239        })?
2240        .dtype;
2241    for input in inputs {
2242        if input.dtype != dtype {
2243            return Err(Error::dtype_mismatch(op, dtype, input.dtype));
2244        }
2245    }
2246
2247    for (input, labels) in inputs.iter().zip(tree.subscripts.inputs.iter()) {
2248        if labels.len() != input.shape.len() {
2249            return Err(Error::rank_mismatch(op, labels.len(), input.shape.len()));
2250        }
2251    }
2252    let output_shape = tree.output_shape();
2253    if output_shape.len() != tree.subscripts.output.len() {
2254        return Err(Error::invalid_argument(
2255            op,
2256            "output labels",
2257            "an output label is missing from all inputs",
2258        ));
2259    }
2260    Ok(ConcreteEinsumInputSpec {
2261        dtype,
2262        shape: output_shape,
2263    })
2264}
2265
2266fn typed_einsum_subscripts<T: TensorScalar>(
2267    session: &mut dyn BackendSession,
2268    inputs: &[&TypedTensor<T>],
2269    subscripts: &Subscripts,
2270    op: &'static str,
2271) -> Result<TypedTensor<T>> {
2272    let reads: Vec<_> = inputs.iter().map(|tensor| T::tensor_read(tensor)).collect();
2273    let plan =
2274        ConcreteEinsumPlan::prepare_subscripts_internal(read_input_specs(&reads), subscripts)?;
2275    let result = plan.execute_read(&reads, session)?;
2276    into_typed_result(result, op)
2277}
2278
2279pub(crate) fn into_typed_result<T: TensorScalar>(
2280    result: Tensor,
2281    op: &'static str,
2282) -> Result<TypedTensor<T>> {
2283    let actual = result.dtype();
2284    T::into_typed(result).map_err(|_| Error::dtype_mismatch(op, T::dtype(), actual))
2285}