1use 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
32pub trait TensorTensordotExt {
34 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
70pub trait TypedTensorTensordotExt<T: TensorScalar> {
72 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
105pub trait TensorEinsumExt {
129 fn einsum(&self, subscripts: &str, session: &mut dyn BackendSession) -> Result<Tensor>;
137
138 fn einsum_notation(
152 &self,
153 notation: &EinsumNotation,
154 session: &mut dyn BackendSession,
155 ) -> Result<Tensor>;
156
157 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(¬ation, 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
217pub trait TensorEinsumIntoExt {
219 fn einsum_into(
227 &self,
228 subscripts: &str,
229 session: &mut dyn BackendSession,
230 out: TensorWrite<'_>,
231 ) -> Result<()>;
232
233 fn einsum_into_notation(
247 &self,
248 notation: &EinsumNotation,
249 session: &mut dyn BackendSession,
250 out: TensorWrite<'_>,
251 ) -> Result<()>;
252
253 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(¬ation, 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
337pub trait TypedTensorEinsumExt<T: TensorScalar> {
360 fn einsum(&self, subscripts: &str, session: &mut dyn BackendSession) -> Result<TypedTensor<T>>;
368
369 fn einsum_notation(
383 &self,
384 notation: &EinsumNotation,
385 session: &mut dyn BackendSession,
386 ) -> Result<TypedTensor<T>>;
387
388 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(¬ation, 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
448pub trait TypedTensorReadEinsumExt<T: TensorScalar> {
471 fn einsum_read(
495 &self,
496 subscripts: &str,
497 session: &mut dyn BackendSession,
498 ) -> Result<TypedTensor<T>>;
499
500 fn einsum_read_notation(
514 &self,
515 notation: &EinsumNotation,
516 session: &mut dyn BackendSession,
517 ) -> Result<TypedTensor<T>>;
518
519 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(¬ation, 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
607pub trait TypedTensorEinsumIntoExt<T: TensorScalar> {
609 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 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 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, ¬ation)?;
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
765pub trait TypedTensorReadEinsumIntoExt<T: TensorScalar> {
785 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 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 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, ¬ation)?;
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
1003pub trait TensorReadEinsumExt {
1030 fn einsum_read(&self, subscripts: &str, session: &mut dyn BackendSession) -> Result<Tensor>;
1038
1039 fn einsum_read_notation(
1053 &self,
1054 notation: &EinsumNotation,
1055 session: &mut dyn BackendSession,
1056 ) -> Result<Tensor>;
1057
1058 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(¬ation, 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
1119pub trait TensorReadEinsumIntoExt {
1121 fn einsum_read_into(
1129 &self,
1130 subscripts: &str,
1131 session: &mut dyn BackendSession,
1132 out: TensorWrite<'_>,
1133 ) -> Result<()>;
1134
1135 fn einsum_read_into_notation(
1149 &self,
1150 notation: &EinsumNotation,
1151 session: &mut dyn BackendSession,
1152 out: TensorWrite<'_>,
1153 ) -> Result<()>;
1154
1155 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(¬ation, 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#[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 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, ¬ation)
1313 }
1314
1315 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 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 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, ¬ation)
1361 }
1362
1363 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 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 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, ¬ation)
1412 }
1413
1414 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 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 pub(crate) fn step_count(&self) -> usize {
1448 self.tree.step_count()
1449 }
1450
1451 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 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 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 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 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 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 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 [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(¬ation.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}