1use std::any::Any;
9use std::hash::Hasher;
10use std::sync::Arc;
11
12use tenferro_ad::extension::{apply_eager_with_extension_session, ExtensionOp};
13use tenferro_ad::EagerTensor;
14use tenferro_cpu::{scalar_fold, CpuBackend};
15use tenferro_ops::{ExtensionShapeContext, SymDim};
16use tenferro_runtime::{
17 EngineId, ErasedExecutionContext, ExecutionContextIdentity, ExtensionCacheKey,
18 ExtensionCacheStore, ExtensionEngine, ExtensionModule, ExtensionModuleError, ExtensionModuleId,
19 ExtensionModuleRegistrar, ExtensionPlanningConfig, ExtensionPrepareRequest, PrepareCapability,
20 PrepareError, PreparedOperation, PreparedOperationBinding, PreparedOperationExecutor,
21 PreparedOperationExecutorHandle, PreparedOperationHandle, PreparedOperationPlan,
22 SpecializationProjection,
23};
24use tenferro_tensor::{DType, Tensor, TensorRead, TensorView};
25use tenferro_tensor::{DynRank, ErasedHostTensor, Host, TypedTensor};
26use tenferro_tensor_core::Scalar;
27
28use crate::{Df64, Df64Add};
29
30pub const DF64_OPS_FAMILY: &str = "tenferro-df64-proof.df64_ops.v1";
35
36pub const DF64_SCALAR_IDENTITY: &str = "tenferro-df64-proof.df64.v1";
42
43macro_rules! df64_operation {
64 ($operation:ty, inputs = $inputs:expr, outputs = $outputs:expr, infer = |$ctx:ident| $infer:block) => {
65 impl ExtensionOp for $operation {
66 fn family_id(&self) -> &'static str {
67 DF64_OPS_FAMILY
68 }
69
70 fn payload_hash(&self, _hasher: &mut dyn Hasher) {}
71
72 fn payload_eq(&self, other: &dyn ExtensionOp) -> bool {
73 other.as_any().downcast_ref::<Self>().is_some()
74 }
75
76 fn clone_arc(&self) -> Arc<dyn ExtensionOp> {
77 Arc::new(self.clone())
78 }
79
80 fn as_any(&self) -> &dyn Any {
81 self
82 }
83
84 fn input_count(&self) -> usize {
85 $inputs
86 }
87
88 fn output_count(&self) -> usize {
89 $outputs
90 }
91
92 fn semantic_effects(&self) -> tenferro_ops::ext_op::ExtensionEffectDeclaration<'_> {
93 tenferro_ops::ext_op::ExtensionEffectDeclaration::Declared(&[])
94 }
95
96 fn semantic_aliases(&self) -> tenferro_ops::ext_op::ExtensionAliasDeclaration<'_> {
97 tenferro_ops::ext_op::ExtensionAliasDeclaration::AllFresh
98 }
99
100 fn scalar_identity(&self) -> Option<&'static str> {
101 Some(DF64_SCALAR_IDENTITY)
102 }
103
104 fn infer_output_meta(
105 &self,
106 $ctx: &mut ExtensionShapeContext<'_>,
107 ) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
108 $infer
109 }
110 }
111 };
112}
113
114#[derive(Clone, Copy, Debug, Default)]
128pub struct Df64Total;
129
130df64_operation!(
131 Df64Total,
132 inputs = 1,
133 outputs = 1,
134 infer = |ctx| {
135 let dtype = ctx.input_dtype(0)?;
136 if !matches!(dtype, DType::External(_)) {
137 return Err(tenferro_tensor::Error::unsupported_dtype(
140 "df64_total",
141 dtype,
142 "df64_total takes an externally defined scalar",
143 ));
144 }
145 Ok(vec![(dtype, Vec::new())])
147 }
148);
149
150#[derive(Clone, Debug)]
167pub struct Df64Expand {
168 pub shape: Box<[usize]>,
170}
171
172impl Df64Expand {
173 #[must_use]
184 pub fn new(shape: Vec<usize>) -> Self {
185 Self {
186 shape: shape.into_boxed_slice(),
187 }
188 }
189}
190
191impl ExtensionOp for Df64Expand {
192 fn family_id(&self) -> &'static str {
193 DF64_OPS_FAMILY
194 }
195
196 fn payload_hash(&self, hasher: &mut dyn Hasher) {
197 for extent in self.shape.iter() {
198 hasher.write_usize(*extent);
199 }
200 }
201
202 fn payload_eq(&self, other: &dyn ExtensionOp) -> bool {
203 other
204 .as_any()
205 .downcast_ref::<Self>()
206 .is_some_and(|other| other.shape == self.shape)
207 }
208
209 fn clone_arc(&self) -> Arc<dyn ExtensionOp> {
210 Arc::new(self.clone())
211 }
212
213 fn as_any(&self) -> &dyn Any {
214 self
215 }
216
217 fn input_count(&self) -> usize {
218 1
219 }
220
221 fn output_count(&self) -> usize {
222 1
223 }
224
225 fn semantic_effects(&self) -> tenferro_ops::ext_op::ExtensionEffectDeclaration<'_> {
226 tenferro_ops::ext_op::ExtensionEffectDeclaration::Declared(&[])
227 }
228
229 fn semantic_aliases(&self) -> tenferro_ops::ext_op::ExtensionAliasDeclaration<'_> {
230 tenferro_ops::ext_op::ExtensionAliasDeclaration::AllFresh
231 }
232
233 fn scalar_identity(&self) -> Option<&'static str> {
234 Some(DF64_SCALAR_IDENTITY)
235 }
236
237 fn infer_output_meta(
238 &self,
239 ctx: &mut ExtensionShapeContext<'_>,
240 ) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
241 let dtype = ctx.input_dtype(0)?;
242 if !matches!(dtype, DType::External(_)) {
243 return Err(tenferro_tensor::Error::unsupported_dtype(
244 "df64_expand",
245 dtype,
246 "df64_expand takes an externally defined scalar",
247 ));
248 }
249 Ok(vec![(
250 dtype,
251 self.shape
252 .iter()
253 .map(|extent| SymDim::from(*extent))
254 .collect(),
255 )])
256 }
257}
258
259#[derive(Clone, Copy, Debug, Default)]
274pub struct Df64FromF64;
275
276df64_operation!(
277 Df64FromF64,
278 inputs = 1,
279 outputs = 1,
280 infer = |ctx| {
281 let dtype = ctx.input_dtype(0)?;
282 if dtype != DType::F64 {
283 return Err(tenferro_tensor::Error::unsupported_dtype(
284 "df64_from_f64",
285 dtype,
286 "df64_from_f64 takes a preset f64 tensor",
287 ));
288 }
289 Ok(vec![(
290 DType::External(DF64_SCALAR),
291 ctx.input_shape(0)?.to_vec(),
292 )])
293 }
294);
295
296#[derive(Clone, Copy, Debug, Default)]
311pub struct Df64ToF64;
312
313df64_operation!(
314 Df64ToF64,
315 inputs = 1,
316 outputs = 1,
317 infer = |ctx| {
318 let dtype = ctx.input_dtype(0)?;
319 if !matches!(dtype, DType::External(_)) {
320 return Err(tenferro_tensor::Error::unsupported_dtype(
321 "df64_to_f64",
322 dtype,
323 "df64_to_f64 takes an externally defined scalar",
324 ));
325 }
326 Ok(vec![(DType::F64, ctx.input_shape(0)?.to_vec())])
327 }
328);
329
330#[derive(Clone, Copy, Debug, Default)]
352pub struct Df64QrVjp {
353 pub has_q: bool,
355 pub has_r: bool,
357}
358
359impl Df64QrVjp {
360 #[must_use]
373 pub const fn of(has_q: bool, has_r: bool) -> Self {
374 Self { has_q, has_r }
375 }
376}
377
378impl ExtensionOp for Df64QrVjp {
379 fn family_id(&self) -> &'static str {
380 DF64_OPS_FAMILY
381 }
382
383 fn payload_hash(&self, hasher: &mut dyn Hasher) {
384 hasher.write_u8(u8::from(self.has_q) | (u8::from(self.has_r) << 1));
385 }
386
387 fn payload_eq(&self, other: &dyn ExtensionOp) -> bool {
388 other
389 .as_any()
390 .downcast_ref::<Self>()
391 .is_some_and(|other| other.has_q == self.has_q && other.has_r == self.has_r)
392 }
393
394 fn clone_arc(&self) -> Arc<dyn ExtensionOp> {
395 Arc::new(*self)
396 }
397
398 fn as_any(&self) -> &dyn Any {
399 self
400 }
401
402 fn input_count(&self) -> usize {
403 2 + usize::from(self.has_q) + usize::from(self.has_r)
405 }
406
407 fn output_count(&self) -> usize {
408 1
409 }
410
411 fn semantic_effects(&self) -> tenferro_ops::ext_op::ExtensionEffectDeclaration<'_> {
412 tenferro_ops::ext_op::ExtensionEffectDeclaration::Declared(&[])
413 }
414
415 fn semantic_aliases(&self) -> tenferro_ops::ext_op::ExtensionAliasDeclaration<'_> {
416 tenferro_ops::ext_op::ExtensionAliasDeclaration::AllFresh
417 }
418
419 fn scalar_identity(&self) -> Option<&'static str> {
420 Some(DF64_SCALAR_IDENTITY)
421 }
422
423 fn infer_output_meta(
424 &self,
425 ctx: &mut ExtensionShapeContext<'_>,
426 ) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
427 let dtype = ctx.input_dtype(0)?;
428 if !matches!(dtype, DType::External(_)) {
429 return Err(tenferro_tensor::Error::unsupported_dtype(
430 "df64_qr_vjp",
431 dtype,
432 "df64_qr_vjp takes an externally defined scalar",
433 ));
434 }
435 Ok(vec![(dtype, ctx.input_shape(0)?.to_vec())])
436 }
437}
438
439#[derive(Clone, Debug, PartialEq, Eq)]
459pub struct Df64Einsum {
460 inputs: Vec<Vec<u32>>,
461 out: Vec<u32>,
462}
463
464impl Df64Einsum {
465 pub fn new(lhs: &[u32], rhs: &[u32], out: &[u32]) -> tenferro_runtime::Result<Self> {
487 Self::new_nary(&[lhs, rhs], out)
488 }
489
490 pub fn new_nary(inputs: &[&[u32]], out: &[u32]) -> tenferro_runtime::Result<Self> {
511 let invalid = |message: &str| {
512 tenferro_runtime::Error::from(tenferro_tensor::Error::invalid_argument(
513 "df64_einsum",
514 "pattern",
515 message,
516 ))
517 };
518 if inputs.len() < 2 {
519 return Err(invalid("a contraction takes at least two operands"));
520 }
521 if inputs.iter().any(|labels| labels.is_empty()) {
522 return Err(invalid("an operand must carry at least one label"));
523 }
524 for label in out {
525 if !inputs.iter().any(|labels| labels.contains(label)) {
526 return Err(invalid(
527 "an output label must appear in at least one operand",
528 ));
529 }
530 }
531 let mut seen = out.to_vec();
532 seen.sort_unstable();
533 seen.dedup();
534 if seen.len() != out.len() {
535 return Err(invalid("an output label repeats"));
536 }
537 Ok(Self {
538 inputs: inputs.iter().map(|labels| labels.to_vec()).collect(),
539 out: out.to_vec(),
540 })
541 }
542
543 #[must_use]
559 pub fn labels(&self) -> Option<(&[u32], &[u32], &[u32])> {
560 match self.inputs.as_slice() {
561 [lhs, rhs] => Some((lhs, rhs, &self.out)),
562 _ => None,
563 }
564 }
565
566 #[must_use]
577 pub fn input_labels(&self) -> &[Vec<u32>] {
578 &self.inputs
579 }
580
581 #[must_use]
592 pub fn out_labels(&self) -> &[u32] {
593 &self.out
594 }
595}
596
597impl ExtensionOp for Df64Einsum {
598 fn family_id(&self) -> &'static str {
599 DF64_OPS_FAMILY
600 }
601
602 fn payload_hash(&self, hasher: &mut dyn Hasher) {
603 hasher.write_usize(self.inputs.len());
604 for labels in self.inputs.iter().chain(core::iter::once(&self.out)) {
605 hasher.write_usize(labels.len());
606 for label in labels {
607 hasher.write_u32(*label);
608 }
609 }
610 }
611
612 fn payload_eq(&self, other: &dyn ExtensionOp) -> bool {
613 other
614 .as_any()
615 .downcast_ref::<Self>()
616 .is_some_and(|other| other == self)
617 }
618
619 fn clone_arc(&self) -> Arc<dyn ExtensionOp> {
620 Arc::new(self.clone())
621 }
622
623 fn as_any(&self) -> &dyn Any {
624 self
625 }
626
627 fn input_count(&self) -> usize {
628 self.inputs.len()
629 }
630
631 fn output_count(&self) -> usize {
632 1
633 }
634
635 fn semantic_effects(&self) -> tenferro_ops::ext_op::ExtensionEffectDeclaration<'_> {
636 tenferro_ops::ext_op::ExtensionEffectDeclaration::Declared(&[])
637 }
638
639 fn semantic_aliases(&self) -> tenferro_ops::ext_op::ExtensionAliasDeclaration<'_> {
640 tenferro_ops::ext_op::ExtensionAliasDeclaration::AllFresh
641 }
642
643 fn scalar_identity(&self) -> Option<&'static str> {
644 Some(DF64_SCALAR_IDENTITY)
645 }
646
647 fn infer_output_meta(
648 &self,
649 ctx: &mut ExtensionShapeContext<'_>,
650 ) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
651 let dtype = ctx.input_dtype(0)?;
652 if !matches!(dtype, DType::External(_)) {
653 return Err(tenferro_tensor::Error::unsupported_dtype(
654 "df64_einsum",
655 dtype,
656 "df64_einsum takes an externally defined scalar",
657 ));
658 }
659 let mut out_shape = Vec::with_capacity(self.out.len());
663 for label in &self.out {
664 let mut extent = None;
665 for (operand, labels) in self.inputs.iter().enumerate() {
666 if ctx.input_dtype(operand)? != dtype {
667 return Err(tenferro_tensor::Error::invalid_argument(
668 "df64_einsum",
669 "inputs",
670 "every operand must carry the same scalar",
671 ));
672 }
673 if let Some(axis) = labels.iter().position(|candidate| candidate == label) {
674 let shape = ctx.input_shape(operand)?;
675 if shape.len() != labels.len() {
676 return Err(tenferro_tensor::Error::rank_mismatch(
677 "df64_einsum",
678 labels.len(),
679 shape.len(),
680 ));
681 }
682 extent = Some(shape[axis].clone());
683 break;
684 }
685 }
686 out_shape.push(extent.ok_or_else(|| {
687 tenferro_tensor::Error::invalid_argument(
688 "df64_einsum",
689 "pattern",
690 "an output label must appear in at least one operand",
691 )
692 })?);
693 }
694 Ok(vec![(dtype, out_shape)])
695 }
696}
697
698#[derive(Clone, Debug, PartialEq, Eq)]
716pub struct Df64EinsumVjp {
717 inputs: Vec<Vec<u32>>,
718 out: Vec<u32>,
719}
720
721impl Df64EinsumVjp {
722 pub fn of(inputs: &[&[u32]], out: &[u32]) -> tenferro_runtime::Result<Self> {
739 Df64Einsum::new_nary(inputs, out)?;
740 Ok(Self {
741 inputs: inputs.iter().map(|labels| labels.to_vec()).collect(),
742 out: out.to_vec(),
743 })
744 }
745
746 #[must_use]
759 pub fn input_labels(&self) -> &[Vec<u32>] {
760 &self.inputs
761 }
762
763 #[must_use]
774 pub fn out_labels(&self) -> &[u32] {
775 &self.out
776 }
777}
778
779impl ExtensionOp for Df64EinsumVjp {
780 fn family_id(&self) -> &'static str {
781 DF64_OPS_FAMILY
782 }
783
784 fn payload_hash(&self, hasher: &mut dyn Hasher) {
785 hasher.write_usize(self.inputs.len());
786 for labels in self.inputs.iter().chain(core::iter::once(&self.out)) {
787 hasher.write_usize(labels.len());
788 for label in labels {
789 hasher.write_u32(*label);
790 }
791 }
792 }
793
794 fn payload_eq(&self, other: &dyn ExtensionOp) -> bool {
795 other
796 .as_any()
797 .downcast_ref::<Self>()
798 .is_some_and(|other| other == self)
799 }
800
801 fn clone_arc(&self) -> Arc<dyn ExtensionOp> {
802 Arc::new(self.clone())
803 }
804
805 fn as_any(&self) -> &dyn Any {
806 self
807 }
808
809 fn input_count(&self) -> usize {
810 self.inputs.len() + 1
812 }
813
814 fn output_count(&self) -> usize {
815 self.inputs.len()
817 }
818
819 fn semantic_effects(&self) -> tenferro_ops::ext_op::ExtensionEffectDeclaration<'_> {
820 tenferro_ops::ext_op::ExtensionEffectDeclaration::Declared(&[])
821 }
822
823 fn semantic_aliases(&self) -> tenferro_ops::ext_op::ExtensionAliasDeclaration<'_> {
824 tenferro_ops::ext_op::ExtensionAliasDeclaration::AllFresh
825 }
826
827 fn scalar_identity(&self) -> Option<&'static str> {
828 Some(DF64_SCALAR_IDENTITY)
829 }
830
831 fn infer_output_meta(
832 &self,
833 ctx: &mut ExtensionShapeContext<'_>,
834 ) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
835 let dtype = ctx.input_dtype(0)?;
836 if !matches!(dtype, DType::External(_)) {
837 return Err(tenferro_tensor::Error::unsupported_dtype(
838 "df64_einsum_vjp",
839 dtype,
840 "df64_einsum_vjp takes an externally defined scalar",
841 ));
842 }
843 let mut shapes = Vec::with_capacity(self.inputs.len());
844 for operand in 0..self.inputs.len() {
845 shapes.push((dtype, ctx.input_shape(operand)?.to_vec()));
846 }
847 Ok(shapes)
848 }
849}
850
851#[derive(Clone, Debug, PartialEq, Eq)]
870pub struct Df64EinsumJvp {
871 inputs: Vec<Vec<u32>>,
872 out: Vec<u32>,
873 tangents: Vec<bool>,
874}
875
876impl Df64EinsumJvp {
877 pub fn of(inputs: &[&[u32]], out: &[u32], tangents: &[bool]) -> tenferro_runtime::Result<Self> {
894 Df64Einsum::new_nary(inputs, out)?;
895 if tangents.len() != inputs.len() {
896 return Err(tenferro_runtime::Error::from(
897 tenferro_tensor::Error::invalid_argument(
898 "df64_einsum_jvp",
899 "tangents",
900 "the tangent mask needs one entry per operand",
901 ),
902 ));
903 }
904 if !tangents.iter().any(|present| *present) {
905 return Err(tenferro_runtime::Error::from(
906 tenferro_tensor::Error::invalid_argument(
907 "df64_einsum_jvp",
908 "tangents",
909 "at least one operand must carry a tangent",
910 ),
911 ));
912 }
913 Ok(Self {
914 inputs: inputs.iter().map(|labels| labels.to_vec()).collect(),
915 out: out.to_vec(),
916 tangents: tangents.to_vec(),
917 })
918 }
919
920 #[must_use]
931 pub fn input_labels(&self) -> &[Vec<u32>] {
932 &self.inputs
933 }
934
935 #[must_use]
947 pub fn out_labels(&self) -> &[u32] {
948 &self.out
949 }
950
951 #[must_use]
963 pub fn tangents(&self) -> &[bool] {
964 &self.tangents
965 }
966}
967
968impl ExtensionOp for Df64EinsumJvp {
969 fn family_id(&self) -> &'static str {
970 DF64_OPS_FAMILY
971 }
972
973 fn payload_hash(&self, hasher: &mut dyn Hasher) {
974 hasher.write_usize(self.inputs.len());
975 for labels in self.inputs.iter().chain(core::iter::once(&self.out)) {
976 hasher.write_usize(labels.len());
977 for label in labels {
978 hasher.write_u32(*label);
979 }
980 }
981 hasher.write_usize(self.tangents.len());
982 for present in &self.tangents {
983 hasher.write_u8(u8::from(*present));
984 }
985 }
986
987 fn payload_eq(&self, other: &dyn ExtensionOp) -> bool {
988 other
989 .as_any()
990 .downcast_ref::<Self>()
991 .is_some_and(|other| other == self)
992 }
993
994 fn clone_arc(&self) -> Arc<dyn ExtensionOp> {
995 Arc::new(self.clone())
996 }
997
998 fn as_any(&self) -> &dyn Any {
999 self
1000 }
1001
1002 fn input_count(&self) -> usize {
1003 self.inputs.len() + self.tangents.iter().filter(|present| **present).count()
1005 }
1006
1007 fn output_count(&self) -> usize {
1008 1
1009 }
1010
1011 fn semantic_effects(&self) -> tenferro_ops::ext_op::ExtensionEffectDeclaration<'_> {
1012 tenferro_ops::ext_op::ExtensionEffectDeclaration::Declared(&[])
1013 }
1014
1015 fn semantic_aliases(&self) -> tenferro_ops::ext_op::ExtensionAliasDeclaration<'_> {
1016 tenferro_ops::ext_op::ExtensionAliasDeclaration::AllFresh
1017 }
1018
1019 fn scalar_identity(&self) -> Option<&'static str> {
1020 Some(DF64_SCALAR_IDENTITY)
1021 }
1022
1023 fn infer_output_meta(
1024 &self,
1025 ctx: &mut ExtensionShapeContext<'_>,
1026 ) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1027 let dtype = ctx.input_dtype(0)?;
1028 if !matches!(dtype, DType::External(_)) {
1029 return Err(tenferro_tensor::Error::unsupported_dtype(
1030 "df64_einsum_jvp",
1031 dtype,
1032 "df64_einsum_jvp takes an externally defined scalar",
1033 ));
1034 }
1035 let shapes: Vec<Vec<SymDim>> = (0..self.inputs.len())
1036 .map(|operand| {
1037 let shape = ctx.input_shape(operand)?;
1038 if shape.len() != self.inputs[operand].len() {
1039 return Err(tenferro_tensor::Error::rank_mismatch(
1040 "df64_einsum_jvp",
1041 self.inputs[operand].len(),
1042 shape.len(),
1043 ));
1044 }
1045 Ok(shape.to_vec())
1046 })
1047 .collect::<tenferro_tensor::Result<Vec<_>>>()?;
1048 let mut out_shape = Vec::with_capacity(self.out.len());
1049 for label in &self.out {
1050 let mut extent = None;
1051 for (operand, labels) in self.inputs.iter().enumerate() {
1052 if let Some(axis) = labels.iter().position(|candidate| candidate == label) {
1053 extent = Some(shapes[operand][axis].clone());
1054 break;
1055 }
1056 }
1057 out_shape.push(extent.ok_or_else(|| {
1058 tenferro_tensor::Error::invalid_argument(
1059 "df64_einsum_jvp",
1060 "pattern",
1061 "an output label must appear in at least one operand",
1062 )
1063 })?);
1064 }
1065 Ok(vec![(dtype, out_shape)])
1066 }
1067}
1068
1069#[derive(Clone, Copy, Debug, Default)]
1081pub struct Df64QrJvp;
1082
1083df64_operation!(
1084 Df64QrJvp,
1085 inputs = 3,
1086 outputs = 2,
1087 infer = |ctx| {
1088 let dtype = ctx.input_dtype(0)?;
1089 if !matches!(dtype, DType::External(_)) {
1090 return Err(tenferro_tensor::Error::unsupported_dtype(
1091 "df64_qr_jvp",
1092 dtype,
1093 "df64_qr_jvp takes an externally defined scalar",
1094 ));
1095 }
1096 Ok(vec![
1097 (dtype, ctx.input_shape(0)?.to_vec()),
1098 (dtype, ctx.input_shape(1)?.to_vec()),
1099 ])
1100 }
1101);
1102
1103pub const DF64_SCALAR: std::any::TypeId = std::any::TypeId::of::<Df64>();
1107
1108#[derive(Clone, Copy, Debug, Default)]
1125pub struct Df64Qr;
1126
1127df64_operation!(
1128 Df64Qr,
1129 inputs = 1,
1130 outputs = 2,
1131 infer = |ctx| {
1132 let dtype = ctx.input_dtype(0)?;
1133 if !matches!(dtype, DType::External(_)) {
1134 return Err(tenferro_tensor::Error::unsupported_dtype(
1135 "df64_qr",
1136 dtype,
1137 "df64_qr takes an externally defined scalar",
1138 ));
1139 }
1140 let shape = ctx.input_shape(0)?.to_vec();
1144 let [rows, columns] = match shape.as_slice() {
1145 [rows, columns] => [rows.clone(), columns.clone()],
1146 _ => {
1147 return Err(tenferro_tensor::Error::invalid_argument(
1148 "df64_qr",
1149 "input",
1150 "df64_qr takes a rank-2 matrix",
1151 ));
1152 }
1153 };
1154 Ok(vec![
1155 (dtype, vec![rows.clone(), columns.clone()]),
1156 (dtype, vec![columns.clone(), columns]),
1157 ])
1158 }
1159);
1160
1161fn qr_of(
1168 session: Option<&mut dyn tenferro_tensor::BackendSession>,
1169 inputs: &[TensorRead<'_>],
1170) -> tenferro_runtime::Result<Vec<Tensor>> {
1171 let input = sole_input("df64_qr", session, inputs)?;
1172 let tensor = input.tensor();
1173 let payload =
1174 external_payload::<Df64>("df64_qr", tensor).map_err(tenferro_runtime::Error::from)?;
1175 let shape = tensor.shape();
1176 let invalid = |message: &'static str| {
1177 tenferro_runtime::Error::from(tenferro_tensor::Error::invalid_argument(
1178 "df64_qr", "input", message,
1179 ))
1180 };
1181 let [rows, columns] = match shape {
1182 [rows, columns] if *rows >= *columns => [*rows, *columns],
1183 _ => {
1184 return Err(invalid(
1185 "df64_qr takes a rank-2 matrix with at least as many rows as columns",
1186 ));
1187 }
1188 };
1189 let values = payload.as_slice();
1190 if values.len() != rows * columns {
1191 return Err(invalid("df64_qr takes a dense column-major matrix"));
1192 }
1193
1194 let source = |row: usize, column: usize| values[row + column * rows];
1196 let mut q = vec![Df64::zero(); rows * columns];
1197 let mut r = vec![Df64::zero(); columns * columns];
1198 for column in 0..columns {
1199 for row in 0..rows {
1200 q[row + column * rows] = source(row, column);
1201 }
1202 for _pass in 0..2 {
1206 for previous in 0..column {
1207 let mut projection = Df64::zero();
1208 for row in 0..rows {
1209 projection = projection + q[row + previous * rows] * q[row + column * rows];
1210 }
1211 for row in 0..rows {
1212 q[row + column * rows] =
1213 q[row + column * rows] - projection * q[row + previous * rows];
1214 }
1215 r[previous + column * columns] = r[previous + column * columns] + projection;
1216 }
1217 }
1218 let mut squares = Df64::zero();
1219 for row in 0..rows {
1220 squares = squares + q[row + column * rows] * q[row + column * rows];
1221 }
1222 let norm = squares.sqrt();
1223 if norm.hi == 0.0 {
1224 return Err(invalid("df64_qr takes a matrix with no zero column"));
1225 }
1226 let sign = if norm.hi < 0.0 {
1229 Df64::from_f64(-1.0)
1230 } else {
1231 Df64::from_f64(1.0)
1232 };
1233 let scale = sign * Df64::from_f64(1.0).ratio(norm);
1234 for row in 0..rows {
1235 q[row + column * rows] = q[row + column * rows] * scale;
1236 }
1237 r[column + column * columns] = sign * norm;
1238 }
1239
1240 let q = TypedTensor::<_, DynRank, Host>::from_host_vec_col_major(vec![rows, columns], q)
1241 .map_err(tenferro_runtime::Error::from)?;
1242 let r = TypedTensor::<_, DynRank, Host>::from_host_vec_col_major(vec![columns, columns], r)
1243 .map_err(tenferro_runtime::Error::from)?;
1244 Ok(vec![
1245 Tensor::external(ErasedHostTensor::new(q)),
1246 Tensor::external(ErasedHostTensor::new(r)),
1247 ])
1248}
1249
1250fn to_f64_of(
1252 session: Option<&mut dyn tenferro_tensor::BackendSession>,
1253 inputs: &[TensorRead<'_>],
1254) -> tenferro_runtime::Result<Vec<Tensor>> {
1255 let input = sole_input("df64_to_f64", session, inputs)?;
1256 let tensor = input.tensor();
1257 crate::conversion::to_f64(tensor)
1258 .map(|tensor| vec![tensor])
1259 .map_err(tenferro_runtime::Error::from)
1260}
1261
1262fn from_f64_of(
1264 session: Option<&mut dyn tenferro_tensor::BackendSession>,
1265 inputs: &[TensorRead<'_>],
1266) -> tenferro_runtime::Result<Vec<Tensor>> {
1267 let input = match inputs.first() {
1268 Some(read @ TensorRead::View(_)) if session.is_none() => {
1271 Input::Materialized(Box::new(owned_f64_read(read)?))
1272 }
1273 _ => sole_input("df64_from_f64", session, inputs)?,
1274 };
1275 let tensor = input.tensor();
1276 crate::conversion::to_df64(tensor)
1277 .map(|tensor| vec![tensor])
1278 .map_err(tenferro_runtime::Error::from)
1279}
1280
1281fn external_payload<'a, T: Scalar>(
1282 op: &'static str,
1283 tensor: &'a Tensor,
1284) -> tenferro_tensor::Result<&'a TypedTensor<T, DynRank, Host>> {
1285 match tensor.external_payload() {
1286 Some(payload) => payload.downcast_ref::<T>().ok_or_else(|| {
1287 tenferro_tensor::Error::unsupported_dtype(
1288 op,
1289 tensor.dtype(),
1290 "the external payload does not hold the expected scalar",
1291 )
1292 }),
1293 None => Err(tenferro_tensor::Error::unsupported_dtype(
1294 op,
1295 tensor.dtype(),
1296 "the operation takes an externally defined payload",
1297 )),
1298 }
1299}
1300
1301#[derive(Debug, Default)]
1309struct Scratch {
1310 q_bar_t: Vec<Df64>,
1312 r_bar_t: Vec<Df64>,
1314 q_bar_t_q: Vec<Df64>,
1316 r_r_bar: Vec<Df64>,
1318 m: Vec<Df64>,
1320 s: Vec<Df64>,
1322 product: Vec<Df64>,
1324 b: Vec<Df64>,
1326}
1327
1328impl Scratch {
1329 const CACHE_NAME: &'static str = "scratch";
1331
1332 fn buffers(&self) -> [&Vec<Df64>; 8] {
1334 [
1335 &self.q_bar_t,
1336 &self.r_bar_t,
1337 &self.q_bar_t_q,
1338 &self.r_r_bar,
1339 &self.m,
1340 &self.s,
1341 &self.product,
1342 &self.b,
1343 ]
1344 }
1345
1346 fn slot(buffer: &mut Vec<Df64>, length: usize) -> &mut [Df64] {
1348 buffer.clear();
1349 buffer.resize(length, Df64::zero());
1350 buffer.as_mut_slice()
1351 }
1352
1353 fn retained_bytes(&self) -> usize {
1355 self.buffers()
1356 .iter()
1357 .map(|buffer| buffer.capacity() * std::mem::size_of::<Df64>())
1358 .sum()
1359 }
1360}
1361
1362fn acquire_scratch(caches: &mut ExtensionCacheStore, shape: usize) -> Scratch {
1364 let key = ExtensionCacheKey::new(DF64_OPS_FAMILY, Scratch::CACHE_NAME, shape as u64);
1365 caches
1366 .get_mut::<Scratch>(&key)
1367 .map_or_else(Scratch::default, |scratch| Scratch {
1368 q_bar_t: std::mem::take(&mut scratch.q_bar_t),
1369 r_bar_t: std::mem::take(&mut scratch.r_bar_t),
1370 q_bar_t_q: std::mem::take(&mut scratch.q_bar_t_q),
1371 r_r_bar: std::mem::take(&mut scratch.r_r_bar),
1372 m: std::mem::take(&mut scratch.m),
1373 s: std::mem::take(&mut scratch.s),
1374 product: std::mem::take(&mut scratch.product),
1375 b: std::mem::take(&mut scratch.b),
1376 })
1377}
1378
1379fn release_scratch(caches: &mut ExtensionCacheStore, shape: usize, scratch: Scratch) {
1381 let key = ExtensionCacheKey::new(DF64_OPS_FAMILY, Scratch::CACHE_NAME, shape as u64);
1382 let retained = scratch.retained_bytes();
1383 caches.put(key, scratch, retained);
1384}
1385
1386enum Input<'a> {
1388 Borrowed(&'a Tensor),
1389 Materialized(Box<Tensor>),
1392}
1393
1394impl Input<'_> {
1395 fn tensor(&self) -> &Tensor {
1396 match self {
1397 Self::Borrowed(tensor) => tensor,
1398 Self::Materialized(tensor) => tensor,
1399 }
1400 }
1401}
1402
1403fn owned_f64_read(read: &TensorRead<'_>) -> tenferro_runtime::Result<Tensor> {
1414 let invalid = |message: &'static str| {
1415 tenferro_runtime::Error::from(tenferro_tensor::Error::invalid_argument(
1416 "df64_from_f64",
1417 "input",
1418 message,
1419 ))
1420 };
1421 let TensorView::F64(view) = read.clone().tensor_view() else {
1422 return Err(invalid("df64_from_f64 takes a preset f64 tensor"));
1423 };
1424 let storage = view.host_storage().map_err(|source| {
1425 tenferro_runtime::Error::from(tenferro_tensor::Error::runtime_state_source(
1426 "df64_from_f64",
1427 source,
1428 ))
1429 })?;
1430 let shape = view.shape().to_vec();
1431 let mut values = Vec::with_capacity(storage.len().min(view.n_elements()));
1432 let mut index = vec![0usize; shape.len()];
1433 for _ in 0..view.n_elements() {
1434 let offset = view
1435 .linear_offset(&index)
1436 .ok_or_else(|| invalid("the borrowed f64 view is outside its buffer"))?;
1437 let value = storage
1438 .get(offset)
1439 .ok_or_else(|| invalid("the borrowed f64 view is outside its buffer"))?;
1440 values.push(*value);
1441 for (position, current) in index.iter_mut().enumerate() {
1442 *current += 1;
1443 if *current < shape[position] {
1444 break;
1445 }
1446 *current = 0;
1447 }
1448 }
1449 Tensor::from_vec_col_major(shape, values).map_err(|source| {
1450 tenferro_runtime::Error::from(tenferro_tensor::Error::runtime_state_source(
1451 "df64_from_f64",
1452 source,
1453 ))
1454 })
1455}
1456
1457fn inputs_of<'a>(
1462 op: &'static str,
1463 session: Option<&mut dyn tenferro_tensor::BackendSession>,
1464 inputs: &'a [TensorRead<'a>],
1465) -> tenferro_runtime::Result<Vec<Input<'a>>> {
1466 let borrowed = inputs
1467 .iter()
1468 .any(|read| matches!(read, TensorRead::View(_)));
1469 if borrowed && session.is_none() {
1470 return Err(tenferro_runtime::Error::from(
1471 tenferro_tensor::Error::invalid_argument(
1472 op,
1473 "input",
1474 "the operation needs a session to read a borrowed input",
1475 ),
1476 ));
1477 }
1478 let mut materialized: Vec<Option<Tensor>> = (0..inputs.len()).map(|_| None).collect();
1479 if let Some(session) = session {
1480 for (index, read) in inputs.iter().enumerate() {
1481 if matches!(read, TensorRead::View(_)) {
1482 materialized[index] = Some(
1483 session
1484 .to_contiguous_read(read.clone())
1485 .map_err(tenferro_runtime::Error::from)?,
1486 );
1487 }
1488 }
1489 }
1490 let mut resolved = Vec::with_capacity(inputs.len());
1491 for (read, owned) in inputs.iter().zip(materialized) {
1492 resolved.push(match (read, owned) {
1493 (TensorRead::Tensor(tensor), _) => Input::Borrowed(tensor),
1494 (TensorRead::View(_), Some(tensor)) => Input::Materialized(Box::new(tensor)),
1495 (TensorRead::View(_), None) => {
1496 return Err(tenferro_runtime::Error::from(
1497 tenferro_tensor::Error::invalid_argument(
1498 op,
1499 "input",
1500 "the operation needs a session to read a borrowed input",
1501 ),
1502 ));
1503 }
1504 });
1505 }
1506 Ok(resolved)
1507}
1508
1509fn resolve_input<'a>(
1511 op: &'static str,
1512 session: Option<&mut dyn tenferro_tensor::BackendSession>,
1513 read: &TensorRead<'a>,
1514) -> tenferro_runtime::Result<Input<'a>> {
1515 match read {
1516 TensorRead::Tensor(tensor) => Ok(Input::Borrowed(tensor)),
1517 view @ TensorRead::View(_) => match session {
1518 Some(session) => Ok(Input::Materialized(Box::new(
1519 session
1520 .to_contiguous_read(view.clone())
1521 .map_err(tenferro_runtime::Error::from)?,
1522 ))),
1523 None => Err(tenferro_runtime::Error::from(
1524 tenferro_tensor::Error::invalid_argument(
1525 op,
1526 "input",
1527 "the operation needs a session to read a borrowed input",
1528 ),
1529 )),
1530 },
1531 }
1532}
1533
1534fn sole_input<'a>(
1540 op: &'static str,
1541 session: Option<&mut dyn tenferro_tensor::BackendSession>,
1542 inputs: &'a [TensorRead<'a>],
1543) -> tenferro_runtime::Result<Input<'a>> {
1544 match inputs.first() {
1545 Some(read) => resolve_input(op, session, read),
1546 None => Err(tenferro_runtime::Error::from(
1547 tenferro_tensor::Error::invalid_argument(op, "input", "the operation takes one input"),
1548 )),
1549 }
1550}
1551
1552fn matrix_of<'a>(
1554 op: &'static str,
1555 tensor: &'a Tensor,
1556) -> tenferro_runtime::Result<crate::dense::Matrix<'a>> {
1557 let invalid = |message: &'static str| {
1558 tenferro_runtime::Error::from(tenferro_tensor::Error::invalid_argument(
1559 op, "input", message,
1560 ))
1561 };
1562 let payload = external_payload::<Df64>(op, tensor).map_err(tenferro_runtime::Error::from)?;
1563 let [rows, columns] = match tensor.shape() {
1564 [rows, columns] => [*rows, *columns],
1565 _ => return Err(invalid("the operation takes a rank-2 matrix")),
1566 };
1567 if payload.as_slice().len() != rows * columns {
1568 return Err(invalid("the operation takes a dense column-major matrix"));
1569 }
1570 Ok(crate::dense::Matrix::borrowed(rows, payload.as_slice()))
1573}
1574
1575fn tensor_of(
1577 op: &'static str,
1578 matrix: crate::dense::Matrix<'_>,
1579) -> tenferro_runtime::Result<Tensor> {
1580 let columns = matrix.columns();
1581 let host = TypedTensor::<_, DynRank, Host>::from_host_vec_col_major(
1582 vec![matrix.rows, columns],
1583 matrix.data.into_owned(),
1584 )
1585 .map_err(|source| {
1586 tenferro_runtime::Error::from(tenferro_tensor::Error::runtime_state_source(op, source))
1587 })?;
1588 Ok(Tensor::external(ErasedHostTensor::new(host)))
1589}
1590
1591fn qr_vjp_of(
1597 mask: (bool, bool),
1598 session: Option<&mut dyn tenferro_tensor::BackendSession>,
1599 caches: &mut ExtensionCacheStore,
1600 inputs: &[TensorRead<'_>],
1601) -> tenferro_runtime::Result<Vec<Tensor>> {
1602 let op = "df64_qr_vjp";
1603 let (has_q, has_r) = mask;
1604 let resolved = inputs_of(op, session, inputs)?;
1605 if resolved.len() != 2 + usize::from(has_q) + usize::from(has_r) {
1606 return Err(tenferro_runtime::Error::from(
1607 tenferro_tensor::Error::invalid_argument(
1608 op,
1609 "input",
1610 "the adjoint takes Q, R, and both cotangents",
1611 ),
1612 ));
1613 }
1614 let q = matrix_of(op, resolved[0].tensor())?;
1615 let r = matrix_of(op, resolved[1].tensor())?;
1616 let mut next = 2;
1617 let q_bar = if has_q {
1620 let matrix = matrix_of(op, resolved[next].tensor())?;
1621 next += 1;
1622 matrix
1623 } else {
1624 crate::dense::zeros(q.rows, q.columns())
1625 };
1626 let r_bar = if has_r {
1627 matrix_of(op, resolved[next].tensor())?
1628 } else {
1629 crate::dense::zeros(r.rows, r.columns())
1630 };
1631
1632 let rows = q.rows;
1636 let columns = q.columns();
1637 let length = rows * columns;
1638 let square = r.rows * r.columns();
1639 let mut scratch = acquire_scratch(caches, square);
1640
1641 crate::dense::transpose_into(Scratch::slot(&mut scratch.q_bar_t, length), &q_bar);
1642 let q_bar_t = crate::dense::Matrix::borrowed(columns, scratch.q_bar_t.as_slice());
1643 crate::dense::transpose_into(Scratch::slot(&mut scratch.r_bar_t, square), &r_bar);
1644 {
1645 let r_bar_t = crate::dense::Matrix::borrowed(r.columns(), scratch.r_bar_t.as_slice());
1646 crate::dense::multiply_into(Scratch::slot(&mut scratch.r_r_bar, square), &r, &r_bar_t);
1647 }
1648 {
1649 let r = crate::dense::Matrix::borrowed(r.rows, scratch.r_r_bar.as_slice());
1650 crate::dense::multiply_into(Scratch::slot(&mut scratch.q_bar_t_q, square), &q_bar_t, &q);
1651 let q_bar_t_q = crate::dense::Matrix::borrowed(columns, scratch.q_bar_t_q.as_slice());
1652 crate::dense::subtract_into(Scratch::slot(&mut scratch.m, square), &r, &q_bar_t_q);
1654 }
1655 {
1656 let m = crate::dense::Matrix::borrowed(r.rows, scratch.m.as_slice());
1657 crate::dense::lower_triangle_into(Scratch::slot(&mut scratch.s, square), &m);
1659 crate::dense::add_strictly_lower_transposed_into(scratch.s.as_mut_slice(), &m);
1662 }
1663 {
1664 let s = crate::dense::Matrix::borrowed(r.rows, scratch.s.as_slice());
1665 crate::dense::multiply_into(Scratch::slot(&mut scratch.product, length), &q, &s);
1666 let product = crate::dense::Matrix::borrowed(rows, scratch.product.as_slice());
1667 let accumulated = Scratch::slot(&mut scratch.b, length);
1668 for (index, slot) in accumulated.iter_mut().enumerate() {
1669 let cotangent = q_bar.data.get(index).copied().unwrap_or_else(Df64::zero);
1670 *slot = cotangent + product.data[index];
1671 }
1672 }
1673 let b = crate::dense::Matrix::borrowed(rows, scratch.b.as_slice());
1674 let a_bar = crate::dense::solve_upper_from_the_right(&r, &b).ok_or_else(|| {
1675 tenferro_runtime::Error::from(tenferro_tensor::Error::invalid_argument(
1676 op,
1677 "input",
1678 "the adjoint needs an invertible triangular factor",
1679 ))
1680 })?;
1681 release_scratch(caches, square, scratch);
1682 Ok(vec![tensor_of(op, a_bar)?])
1683}
1684
1685fn qr_jvp_of(
1691 session: Option<&mut dyn tenferro_tensor::BackendSession>,
1692 inputs: &[TensorRead<'_>],
1693) -> tenferro_runtime::Result<Vec<Tensor>> {
1694 let op = "df64_qr_jvp";
1695 let resolved = inputs_of(op, session, inputs)?;
1696 if resolved.len() != 3 {
1697 return Err(tenferro_runtime::Error::from(
1698 tenferro_tensor::Error::invalid_argument(
1699 op,
1700 "input",
1701 "the tangent takes Q, R, and the input tangent",
1702 ),
1703 ));
1704 }
1705 let q = matrix_of(op, resolved[0].tensor())?;
1706 let r = matrix_of(op, resolved[1].tensor())?;
1707 let a_dot = matrix_of(op, resolved[2].tensor())?;
1708
1709 let w = crate::dense::multiply(&crate::dense::transpose(&q), &a_dot);
1715 let n = r.columns();
1716 let mut s = crate::dense::zeros(n, n);
1717 for column in 0..n {
1718 for row in (column + 1)..n {
1719 let mut value = w.at(row, column);
1720 for earlier in 0..column {
1721 value = value - s.at(row, earlier) * r.at(earlier, column);
1722 }
1723 let skew = value / r.at(column, column);
1724 s.set(row, column, skew);
1725 s.set(column, row, -skew);
1726 }
1727 }
1728 let r_dot = crate::dense::subtract(&w, &crate::dense::multiply(&s, &r));
1729 let residual = crate::dense::subtract(&a_dot, &crate::dense::multiply(&q, &r_dot));
1730 let q_dot = crate::dense::solve_upper_from_the_right(&r, &residual).ok_or_else(|| {
1731 tenferro_runtime::Error::from(tenferro_tensor::Error::invalid_argument(
1732 op,
1733 "input",
1734 "the tangent needs an invertible triangular factor",
1735 ))
1736 })?;
1737 Ok(vec![tensor_of(op, q_dot)?, tensor_of(op, r_dot)?])
1738}
1739
1740fn shape_of_labels(labels: &[Box<[u32]>], shapes: &[Vec<usize>], out_labels: &[u32]) -> Vec<usize> {
1742 let mut extents: Vec<(u32, usize)> = Vec::new();
1743 for (operand, operand_labels) in labels.iter().enumerate() {
1744 for (axis, label) in operand_labels.iter().enumerate() {
1745 if !extents.iter().any(|(existing, _)| existing == label) {
1746 extents.push((*label, shapes[operand][axis]));
1747 }
1748 }
1749 }
1750 out_labels
1751 .iter()
1752 .map(|label| {
1753 extents
1754 .iter()
1755 .find(|(existing, _)| existing == label)
1756 .map(|(_, extent)| *extent)
1757 .unwrap_or(1)
1758 })
1759 .collect()
1760}
1761
1762fn contract_in_scalar(
1769 op: &'static str,
1770 labels: &[Box<[u32]>],
1771 out_labels: &[u32],
1772 lhs_values: &[Df64],
1773 lhs_shape: &[usize],
1774 rhs_values: &[Df64],
1775 rhs_shape: &[usize],
1776) -> tenferro_runtime::Result<Vec<Df64>> {
1777 if element_count(lhs_shape) != lhs_values.len() || element_count(rhs_shape) != rhs_values.len()
1778 {
1779 return Err(tenferro_runtime::Error::from(
1780 tenferro_tensor::Error::invalid_argument(
1781 op,
1782 "inputs",
1783 "the payload length does not match the declared shape",
1784 ),
1785 ));
1786 }
1787
1788 let mut extents: Vec<(u32, usize)> = Vec::new();
1789 let label_extent = |label: u32, extent: usize, extents: &mut Vec<(u32, usize)>| match extents
1790 .iter()
1791 .find(|(existing, _)| *existing == label)
1792 {
1793 Some((_, existing)) => *existing == extent,
1794 None => {
1795 extents.push((label, extent));
1796 true
1797 }
1798 };
1799 let mut ok = true;
1800 for (axis, label) in labels[0].iter().enumerate() {
1801 ok &= label_extent(*label, lhs_shape[axis], &mut extents);
1802 }
1803 for (axis, label) in labels[1].iter().enumerate() {
1804 ok &= label_extent(*label, rhs_shape[axis], &mut extents);
1805 }
1806 if !ok {
1807 return Err(tenferro_runtime::Error::from(
1808 tenferro_tensor::Error::invalid_argument(
1809 op,
1810 "inputs",
1811 "the inputs disagree on the extent of a shared label",
1812 ),
1813 ));
1814 }
1815 let extent_of = |label: u32, extents: &[(u32, usize)]| {
1816 extents
1817 .iter()
1818 .find(|(existing, _)| *existing == label)
1819 .map(|(_, extent)| *extent)
1820 .unwrap_or(1)
1821 };
1822 let out_shape: Vec<usize> = out_labels
1823 .iter()
1824 .map(|label| extent_of(*label, &extents))
1825 .collect();
1826 let mut summed_labels: Vec<u32> = Vec::new();
1827 for label in labels[0].iter().chain(labels[1].iter()) {
1828 if !out_labels.contains(label) && !summed_labels.contains(label) {
1829 summed_labels.push(*label);
1830 }
1831 }
1832 let summed_shape: Vec<usize> = summed_labels
1833 .iter()
1834 .map(|label| extent_of(*label, &extents))
1835 .collect();
1836
1837 let out_count = element_count(&out_shape);
1838 let summed_count = element_count(&summed_shape);
1839 let mut result = vec![Df64::zero(); out_count];
1840 let mut out_index = vec![0usize; out_shape.len()];
1841 let mut summed_index = vec![0usize; summed_shape.len()];
1842 for slot in result.iter_mut() {
1843 for value in summed_index.iter_mut() {
1844 *value = 0;
1845 }
1846 let mut accumulator = Df64::zero();
1847 for _ in 0..summed_count {
1848 let lhs_offset = offset_for(
1849 &labels[0],
1850 lhs_shape,
1851 out_labels,
1852 &out_index,
1853 &summed_labels,
1854 &summed_index,
1855 );
1856 let rhs_offset = offset_for(
1857 &labels[1],
1858 rhs_shape,
1859 out_labels,
1860 &out_index,
1861 &summed_labels,
1862 &summed_index,
1863 );
1864 accumulator = accumulator + lhs_values[lhs_offset] * rhs_values[rhs_offset];
1865 advance(&mut summed_index, &summed_shape);
1866 }
1867 *slot = accumulator;
1868 advance(&mut out_index, &out_shape);
1869 }
1870 Ok(result)
1871}
1872
1873fn external_of(values: Vec<Df64>, shape: Vec<usize>) -> tenferro_runtime::Result<Tensor> {
1875 let tensor = TypedTensor::<_, DynRank, Host>::from_host_vec_col_major(shape, values)
1876 .map_err(tenferro_runtime::Error::from)?;
1877 Ok(Tensor::external(ErasedHostTensor::new(tensor)))
1878}
1879
1880fn fold_in_scalar(
1891 op: &'static str,
1892 labels: &[Box<[u32]>],
1893 out_labels: &[u32],
1894 operands: &[(Vec<Df64>, Vec<usize>)],
1895) -> tenferro_runtime::Result<(Vec<Df64>, Vec<usize>)> {
1896 let mut accumulator = operands[0].0.clone();
1897 let mut accumulator_shape = operands[0].1.clone();
1898 let mut accumulator_labels: Vec<u32> = labels[0].to_vec();
1899 for index in 1..labels.len() {
1900 let next_labels: Vec<u32> = labels[index].to_vec();
1901 let keep: Vec<u32> = if index + 1 == labels.len() {
1902 out_labels.to_vec()
1903 } else {
1904 let mut keep: Vec<u32> = Vec::new();
1905 for label in accumulator_labels.iter().chain(next_labels.iter()) {
1906 let needed_later = out_labels.contains(label)
1907 || labels[index + 1..].iter().any(|rest| rest.contains(label));
1908 if needed_later && !keep.contains(label) {
1909 keep.push(*label);
1910 }
1911 }
1912 keep
1913 };
1914 let pair = [
1915 accumulator_labels.clone().into_boxed_slice(),
1916 next_labels.clone().into_boxed_slice(),
1917 ];
1918 let contracted = contract_in_scalar(
1919 op,
1920 &pair,
1921 &keep,
1922 &accumulator,
1923 &accumulator_shape,
1924 &operands[index].0,
1925 &operands[index].1,
1926 )?;
1927 let contracted_shape = shape_of_labels(
1928 &pair,
1929 &[accumulator_shape.clone(), operands[index].1.clone()],
1930 &keep,
1931 );
1932 accumulator = contracted;
1933 accumulator_shape = contracted_shape;
1934 accumulator_labels = keep;
1935 }
1936 Ok((accumulator, accumulator_shape))
1937}
1938
1939fn einsum_of(
1941 labels: &[Box<[u32]>],
1942 out_labels: &[u32],
1943 session: Option<&mut dyn tenferro_tensor::BackendSession>,
1944 inputs: &[TensorRead<'_>],
1945) -> tenferro_runtime::Result<Vec<Tensor>> {
1946 let op = "df64_einsum";
1947 let resolved = inputs_of(op, session, inputs)?;
1948 if resolved.len() != labels.len() || labels.len() < 2 {
1949 return Err(tenferro_runtime::Error::from(
1950 tenferro_tensor::Error::invalid_argument(
1951 op,
1952 "input",
1953 "a contraction takes one input per operand",
1954 ),
1955 ));
1956 }
1957 let mut operands: Vec<(Vec<Df64>, Vec<usize>)> = Vec::with_capacity(resolved.len());
1958 for operand in &resolved {
1959 operands.push((
1960 payload_of::<Df64>(op, operand.tensor())?,
1961 operand.tensor().shape().to_vec(),
1962 ));
1963 }
1964 let (values, shape) = fold_in_scalar(op, labels, out_labels, &operands)?;
1965 Ok(vec![external_of(values, shape)?])
1966}
1967
1968fn einsum_vjp_of(
1973 labels: &[Box<[u32]>],
1974 out_labels: &[u32],
1975 session: Option<&mut dyn tenferro_tensor::BackendSession>,
1976 inputs: &[TensorRead<'_>],
1977) -> tenferro_runtime::Result<Vec<Tensor>> {
1978 let op = "df64_einsum_vjp";
1979 let resolved = inputs_of(op, session, inputs)?;
1980 if resolved.len() != labels.len() + 1 || labels.len() < 2 {
1981 return Err(tenferro_runtime::Error::from(
1982 tenferro_tensor::Error::invalid_argument(
1983 op,
1984 "input",
1985 "the adjoint takes one input per operand and the output cotangent",
1986 ),
1987 ));
1988 }
1989 let mut operands: Vec<(Vec<Df64>, Vec<usize>)> = Vec::with_capacity(labels.len());
1990 for operand in &resolved[..labels.len()] {
1991 operands.push((
1992 payload_of::<Df64>(op, operand.tensor())?,
1993 operand.tensor().shape().to_vec(),
1994 ));
1995 }
1996 let cotangent = payload_of::<Df64>(op, resolved[labels.len()].tensor())?;
1997 let cotangent_shape = resolved[labels.len()].tensor().shape().to_vec();
1998
1999 let mut outputs = Vec::with_capacity(labels.len());
2002 for position in 0..labels.len() {
2003 let mut substituted: Vec<(Vec<Df64>, Vec<usize>)> = Vec::with_capacity(labels.len());
2004 for (index, operand) in operands.iter().enumerate() {
2005 if index == position {
2006 substituted.push((cotangent.clone(), cotangent_shape.clone()));
2007 } else {
2008 substituted.push(operand.clone());
2009 }
2010 }
2011 let mut step_labels: Vec<Box<[u32]>> = labels.to_vec();
2012 step_labels[position] = out_labels.to_vec().into_boxed_slice();
2013 let (values, shape) = fold_in_scalar(op, &step_labels, &labels[position], &substituted)?;
2014 outputs.push(external_of(values, shape)?);
2015 }
2016 Ok(outputs)
2017}
2018
2019fn einsum_jvp_of(
2021 labels: &[Box<[u32]>],
2022 out_labels: &[u32],
2023 tangents: &[bool],
2024 session: Option<&mut dyn tenferro_tensor::BackendSession>,
2025 inputs: &[TensorRead<'_>],
2026) -> tenferro_runtime::Result<Vec<Tensor>> {
2027 let op = "df64_einsum_jvp";
2028 let resolved = inputs_of(op, session, inputs)?;
2029 let expected = labels.len() + tangents.iter().filter(|present| **present).count();
2030 if resolved.len() != expected || labels.len() != tangents.len() || labels.len() < 2 {
2031 return Err(tenferro_runtime::Error::from(
2032 tenferro_tensor::Error::invalid_argument(
2033 op,
2034 "input",
2035 "the tangent takes one input per operand and one per tangent",
2036 ),
2037 ));
2038 }
2039 let mut operands: Vec<(Vec<Df64>, Vec<usize>)> = Vec::with_capacity(labels.len());
2040 let mut shapes: Vec<Vec<usize>> = Vec::with_capacity(labels.len());
2041 for operand in &resolved[..labels.len()] {
2042 operands.push((
2043 payload_of::<Df64>(op, operand.tensor())?,
2044 operand.tensor().shape().to_vec(),
2045 ));
2046 shapes.push(operand.tensor().shape().to_vec());
2047 }
2048 let mut dots: Vec<Option<(Vec<Df64>, Vec<usize>)>> = vec![None; labels.len()];
2049 let mut slot = labels.len();
2050 for (index, present) in tangents.iter().enumerate() {
2051 if !present {
2052 continue;
2053 }
2054 dots[index] = Some((
2055 payload_of::<Df64>(op, resolved[slot].tensor())?,
2056 resolved[slot].tensor().shape().to_vec(),
2057 ));
2058 slot += 1;
2059 }
2060 let out_shape = shape_of_labels(labels, &shapes, out_labels);
2061
2062 let mut total: Option<Vec<Df64>> = None;
2064 for (position, dot) in dots.iter().enumerate() {
2065 let Some((values, shape)) = dot else {
2066 continue;
2067 };
2068 let mut substituted = operands.clone();
2069 substituted[position] = (values.clone(), shape.clone());
2070 let (part, _) = fold_in_scalar(op, labels, out_labels, &substituted)?;
2071 total = Some(match total {
2072 Some(existing) => existing
2073 .iter()
2074 .zip(&part)
2075 .map(|(left, right)| *left + *right)
2076 .collect(),
2077 None => part,
2078 });
2079 }
2080 let values = total.ok_or_else(|| {
2081 tenferro_runtime::Error::from(tenferro_tensor::Error::invalid_argument(
2082 op,
2083 "tangents",
2084 "at least one operand must carry a tangent",
2085 ))
2086 })?;
2087 Ok(vec![external_of(values, out_shape)?])
2088}
2089
2090fn offset_for(
2092 input_labels: &[u32],
2093 input_shape: &[usize],
2094 out_labels: &[u32],
2095 out_index: &[usize],
2096 summed_labels: &[u32],
2097 summed_index: &[usize],
2098) -> usize {
2099 let mut offset = 0usize;
2100 let mut stride = 1usize;
2101 for (axis, label) in input_labels.iter().enumerate() {
2102 let position = out_labels
2103 .iter()
2104 .position(|candidate| candidate == label)
2105 .map(|index| out_index[index])
2106 .or_else(|| {
2107 summed_labels
2108 .iter()
2109 .position(|candidate| candidate == label)
2110 .map(|index| summed_index[index])
2111 })
2112 .unwrap_or(0);
2113 offset += position * stride;
2114 stride *= input_shape[axis];
2115 }
2116 offset
2117}
2118
2119fn advance(index: &mut [usize], shape: &[usize]) {
2121 for axis in 0..shape.len() {
2122 index[axis] += 1;
2123 if index[axis] < shape[axis] {
2124 return;
2125 }
2126 index[axis] = 0;
2127 }
2128}
2129
2130fn element_count(shape: &[usize]) -> usize {
2132 shape.iter().product()
2133}
2134
2135fn payload_of<T: tenferro_tensor_core::Scalar>(
2137 op: &'static str,
2138 tensor: &Tensor,
2139) -> tenferro_runtime::Result<Vec<T>> {
2140 external_payload::<T>(op, tensor)
2141 .map(|payload| payload.as_slice().to_vec())
2142 .map_err(tenferro_runtime::Error::from)
2143}
2144
2145fn expand_of(
2147 shape: &[usize],
2148 session: Option<&mut dyn tenferro_tensor::BackendSession>,
2149 inputs: &[TensorRead<'_>],
2150) -> tenferro_runtime::Result<Vec<Tensor>> {
2151 let input = sole_input("df64_expand", session, inputs)?;
2152 let tensor = input.tensor();
2153 let payload =
2154 external_payload::<Df64>("df64_expand", tensor).map_err(tenferro_runtime::Error::from)?;
2155 let value = payload.as_slice().first().copied().ok_or_else(|| {
2156 tenferro_runtime::Error::from(tenferro_tensor::Error::invalid_argument(
2157 "df64_expand",
2158 "input",
2159 "df64_expand takes a scalar payload",
2160 ))
2161 })?;
2162 let count: usize = shape.iter().product();
2163 let output = TypedTensor::<_, DynRank, Host>::from_host_vec_col_major(
2164 shape.to_vec(),
2165 vec![value; count],
2166 )
2167 .map_err(tenferro_runtime::Error::from)?;
2168 Ok(vec![Tensor::external(ErasedHostTensor::new(output))])
2169}
2170
2171fn total_of(
2172 session: Option<&mut dyn tenferro_tensor::BackendSession>,
2173 inputs: &[TensorRead<'_>],
2174) -> tenferro_runtime::Result<Vec<Tensor>> {
2175 let input = sole_input("df64_total", session, inputs)?;
2176 let tensor = input.tensor();
2177 let payload =
2178 external_payload::<Df64>("df64_total", tensor).map_err(tenferro_runtime::Error::from)?;
2179 let total = scalar_fold::<Df64, Df64Add>("df64_total", payload, Df64::zero())
2180 .map_err(tenferro_runtime::Error::from)?;
2181 let output = TypedTensor::<_, DynRank, Host>::from_host_vec_col_major(vec![], vec![total])
2182 .map_err(tenferro_runtime::Error::from)?;
2183 Ok(vec![Tensor::external(ErasedHostTensor::new(output))])
2184}
2185
2186#[derive(Debug)]
2188enum Df64Body {
2189 Total,
2191 Expand(Box<[usize]>),
2193 Qr,
2195 ToF64,
2197 FromF64,
2199 QrVjp((bool, bool)),
2201 QrJvp,
2203 EinsumJvp {
2205 inputs: Box<[Box<[u32]>]>,
2207 out: Box<[u32]>,
2209 tangents: Box<[bool]>,
2211 },
2212 EinsumVjp {
2214 inputs: Box<[Box<[u32]>]>,
2216 out: Box<[u32]>,
2218 },
2219 Einsum {
2221 inputs: Box<[Box<[u32]>]>,
2223 out: Box<[u32]>,
2225 },
2226}
2227
2228impl Df64Body {
2229 fn execute(
2230 &self,
2231 session: Option<&mut dyn tenferro_tensor::BackendSession>,
2232 caches: &mut ExtensionCacheStore,
2233 inputs: &[TensorRead<'_>],
2234 ) -> tenferro_runtime::Result<Vec<Tensor>> {
2235 match self {
2236 Self::Total => total_of(session, inputs),
2237 Self::Expand(shape) => expand_of(shape, session, inputs),
2238 Self::Qr => qr_of(session, inputs),
2239 Self::ToF64 => to_f64_of(session, inputs),
2240 Self::FromF64 => from_f64_of(session, inputs),
2241 Self::QrVjp(mask) => qr_vjp_of(*mask, session, caches, inputs),
2242 Self::QrJvp => qr_jvp_of(session, inputs),
2243 Self::Einsum {
2244 inputs: labels,
2245 out,
2246 } => einsum_of(labels, out, session, inputs),
2247 Self::EinsumVjp {
2248 inputs: labels,
2249 out,
2250 } => einsum_vjp_of(labels, out, session, inputs),
2251 Self::EinsumJvp {
2252 inputs: labels,
2253 out,
2254 tangents,
2255 } => einsum_jvp_of(labels, out, tangents, session, inputs),
2256 }
2257 }
2258}
2259
2260#[derive(Debug)]
2261struct Df64Prepared {
2262 binding: PreparedOperationBinding,
2263 specialization: SpecializationProjection,
2264 body: Df64Body,
2265}
2266
2267impl PreparedOperation for Df64Prepared {
2268 fn binding(&self) -> &PreparedOperationBinding {
2269 &self.binding
2270 }
2271
2272 fn specialization(&self) -> &SpecializationProjection {
2273 &self.specialization
2274 }
2275
2276 fn retained_bytes(&self) -> usize {
2277 0
2278 }
2279}
2280
2281impl PreparedOperationExecutor for Df64Prepared {
2282 fn execute(
2283 &self,
2284 _context: &mut ErasedExecutionContext<'_>,
2285 caches: &mut ExtensionCacheStore,
2286 inputs: &[TensorRead<'_>],
2287 ) -> tenferro_runtime::Result<Vec<Tensor>> {
2288 self.body.execute(None, caches, inputs)
2291 }
2292
2293 fn supports_session(&self) -> bool {
2294 true
2295 }
2296
2297 fn execute_in_session(
2298 &self,
2299 session: &mut dyn tenferro_tensor::BackendSession,
2300 caches: &mut ExtensionCacheStore,
2301 inputs: &[TensorRead<'_>],
2302 ) -> tenferro_runtime::Result<Vec<Tensor>> {
2303 self.body.execute(Some(session), caches, inputs)
2304 }
2305}
2306
2307#[derive(Debug)]
2308struct Df64Engine {
2309 family_id: &'static str,
2310 engine_id: EngineId,
2311}
2312
2313impl ExtensionEngine for Df64Engine {
2314 fn family_id(&self) -> &'static str {
2315 self.family_id
2316 }
2317
2318 fn engine_id(&self) -> &EngineId {
2319 &self.engine_id
2320 }
2321
2322 fn context_identity(&self) -> ExecutionContextIdentity {
2323 ExecutionContextIdentity::of::<CpuBackend>()
2324 }
2325
2326 fn prepare(
2327 &self,
2328 request: ExtensionPrepareRequest<'_>,
2329 ) -> Result<PrepareCapability, PrepareError> {
2330 let operation = request.operation().as_any();
2331 let body = if let Some(expand) = operation.downcast_ref::<Df64Expand>() {
2332 Df64Body::Expand(expand.shape.clone())
2333 } else if operation.downcast_ref::<Df64Qr>().is_some() {
2334 Df64Body::Qr
2335 } else if operation.downcast_ref::<Df64ToF64>().is_some() {
2336 Df64Body::ToF64
2337 } else if operation.downcast_ref::<Df64FromF64>().is_some() {
2338 Df64Body::FromF64
2339 } else if let Some(adjoint) = operation.downcast_ref::<Df64QrVjp>() {
2340 Df64Body::QrVjp((adjoint.has_q, adjoint.has_r))
2341 } else if operation.downcast_ref::<Df64QrJvp>().is_some() {
2342 Df64Body::QrJvp
2343 } else if let Some(tangent) = operation.downcast_ref::<Df64EinsumJvp>() {
2344 Df64Body::EinsumJvp {
2345 inputs: tangent
2346 .input_labels()
2347 .iter()
2348 .map(|labels| labels.clone().into_boxed_slice())
2349 .collect(),
2350 out: tangent.out_labels().to_vec().into_boxed_slice(),
2351 tangents: tangent.tangents().to_vec().into_boxed_slice(),
2352 }
2353 } else if let Some(adjoint) = operation.downcast_ref::<Df64EinsumVjp>() {
2354 Df64Body::EinsumVjp {
2355 inputs: adjoint
2356 .input_labels()
2357 .iter()
2358 .map(|labels| labels.clone().into_boxed_slice())
2359 .collect(),
2360 out: adjoint.out_labels().to_vec().into_boxed_slice(),
2361 }
2362 } else if let Some(contraction) = operation.downcast_ref::<Df64Einsum>() {
2363 Df64Body::Einsum {
2364 inputs: contraction
2365 .input_labels()
2366 .iter()
2367 .map(|labels| labels.clone().into_boxed_slice())
2368 .collect(),
2369 out: contraction.out_labels().to_vec().into_boxed_slice(),
2370 }
2371 } else {
2372 Df64Body::Total
2373 };
2374 let prepared = Arc::new(Df64Prepared {
2375 binding: request.binding().clone(),
2376 specialization: request.specialization().clone(),
2377 body,
2378 });
2379 let operation: PreparedOperationHandle = Arc::clone(&prepared) as PreparedOperationHandle;
2380 let executor: PreparedOperationExecutorHandle = prepared as PreparedOperationExecutorHandle;
2381 Ok(PrepareCapability::Prepared(
2382 PreparedOperationPlan::executable(operation, executor),
2383 ))
2384 }
2385}
2386
2387#[derive(Debug)]
2388struct Df64Config {
2389 family_id: &'static str,
2390}
2391
2392impl ExtensionPlanningConfig for Df64Config {
2393 fn family_id(&self) -> &'static str {
2394 self.family_id
2395 }
2396
2397 fn as_any(&self) -> &dyn Any {
2398 self
2399 }
2400
2401 fn payload_hash(&self, state: &mut dyn Hasher) {
2402 state.write(self.family_id.as_bytes());
2403 }
2404
2405 fn payload_eq(&self, other: &dyn ExtensionPlanningConfig) -> bool {
2406 other
2407 .as_any()
2408 .downcast_ref::<Self>()
2409 .is_some_and(|other| self.family_id == other.family_id)
2410 }
2411
2412 fn retained_bytes(&self) -> usize {
2413 0
2414 }
2415}
2416
2417#[derive(Debug)]
2418struct Df64TotalModule {
2419 module_id: ExtensionModuleId,
2420 engine_id: EngineId,
2421}
2422
2423impl ExtensionModule for Df64TotalModule {
2424 fn module_id(&self) -> &ExtensionModuleId {
2425 &self.module_id
2426 }
2427
2428 fn configure(
2429 &self,
2430 registrar: &mut ExtensionModuleRegistrar<'_>,
2431 ) -> Result<(), ExtensionModuleError> {
2432 registrar.register_engine(Arc::new(Df64Engine {
2433 family_id: DF64_OPS_FAMILY,
2434 engine_id: self.engine_id.clone(),
2435 }))?;
2436 registrar.register_planning_config(
2437 self.engine_id.clone(),
2438 Arc::new(Df64Config {
2439 family_id: DF64_OPS_FAMILY,
2440 }),
2441 )
2442 }
2443}
2444
2445pub fn module() -> Result<Arc<dyn ExtensionModule>, tenferro_runtime::RuntimeConfigError> {
2460 module_for_engine(tenferro_cpu::runtime_engine_id()?)
2461}
2462
2463pub fn module_for_engine(
2486 engine_id: EngineId,
2487) -> Result<Arc<dyn ExtensionModule>, tenferro_runtime::RuntimeConfigError> {
2488 Ok(Arc::new(Df64TotalModule {
2489 module_id: ExtensionModuleId::new("tenferro-df64-proof.module")?,
2490 engine_id,
2491 }))
2492}
2493
2494pub fn apply_total(input: &EagerTensor) -> tenferro_runtime::Result<Vec<EagerTensor>> {
2512 let module = module().map_err(|source| {
2513 tenferro_runtime::Error::runtime_state_source(
2514 "df64_total",
2515 tenferro_runtime::ErrorPhase::Execution,
2516 source,
2517 )
2518 })?;
2519 apply_eager_with_extension_session(Arc::new(Df64Total), &[input], module)
2520}