1use std::any::Any;
2use std::hash::{Hash, Hasher};
3use std::sync::Arc;
4
5use num_complex::{Complex32, Complex64};
6use tenferro_cpu::with_cpu_exec_session;
7use tenferro_extension_macros::define_extension_runtime;
8use tenferro_ops::SymDim;
9use tenferro_runtime::extension::ExtensionOp;
10use tenferro_tensor::{BackendSession, DType, Error, ErrorKind, Tensor, TensorBackend, TensorRead};
11
12#[cfg(feature = "cuda")]
13use tenferro_gpu::cuda::with_cuda_exec_session;
14
15use crate::backend::LinalgBackend;
16use crate::RankRevealingQrOptions;
17
18#[cfg(all(test, feature = "cuda"))]
19#[path = "extension/cuda_tests.rs"]
20mod cuda_tests;
21mod gauge;
22#[cfg(all(test, not(feature = "cuda")))]
23mod tests;
24
25pub(crate) use gauge::{apply_eigh_gauge, apply_qr_gauge};
26
27pub const LINALG_EXTENSION_FAMILY_ID: &str = "tenferro-linalg.linalg.v1";
28
29pub const DEFAULT_DECOMPOSITION_DERIVATIVE_EPS: f64 = 1e-12;
43
44#[derive(Clone, Copy, Debug, PartialEq, Eq)]
55pub enum SvdGauge {
56 Raw,
58 CanonicalPivot,
61}
62
63#[derive(Clone, Copy, Debug, PartialEq, Eq)]
74pub enum EighGauge {
75 Raw,
77 CanonicalPivot,
79}
80
81#[derive(Clone, Copy, Debug, PartialEq, Eq)]
92pub enum QrGauge {
93 Raw,
95 PositiveDiagonal,
97}
98
99#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
117pub enum SvdDriver {
118 #[default]
122 Auto,
123 Gesvdj,
125 Gesvd,
127 Xgesvdp,
141}
142
143#[derive(Clone, Copy, Debug, PartialEq)]
159pub struct SvdOptions {
160 pub gauge: SvdGauge,
162 pub derivative_eps: f64,
164 pub driver: SvdDriver,
166}
167
168impl Default for SvdOptions {
169 fn default() -> Self {
170 Self {
171 gauge: SvdGauge::Raw,
172 derivative_eps: DEFAULT_DECOMPOSITION_DERIVATIVE_EPS,
173 driver: SvdDriver::Auto,
174 }
175 }
176}
177
178impl SvdOptions {
179 pub fn gauge(mut self, gauge: SvdGauge) -> Self {
190 self.gauge = gauge;
191 self
192 }
193
194 pub fn derivative_eps(mut self, derivative_eps: f64) -> Self {
205 self.derivative_eps = derivative_eps;
206 self
207 }
208
209 pub fn driver(mut self, driver: SvdDriver) -> Self {
220 self.driver = driver;
221 self
222 }
223}
224
225#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
243pub enum EighDriver {
244 #[default]
253 Auto,
254 Syevd,
257 Syevj,
265}
266
267#[derive(Clone, Copy, Debug, PartialEq)]
281pub struct EighOptions {
282 pub gauge: EighGauge,
284 pub derivative_eps: f64,
286 pub driver: EighDriver,
288}
289
290impl Default for EighOptions {
291 fn default() -> Self {
292 Self {
293 gauge: EighGauge::Raw,
294 derivative_eps: DEFAULT_DECOMPOSITION_DERIVATIVE_EPS,
295 driver: EighDriver::Auto,
296 }
297 }
298}
299
300impl EighOptions {
301 pub fn gauge(mut self, gauge: EighGauge) -> Self {
312 self.gauge = gauge;
313 self
314 }
315
316 pub fn derivative_eps(mut self, derivative_eps: f64) -> Self {
327 self.derivative_eps = derivative_eps;
328 self
329 }
330
331 pub fn driver(mut self, driver: EighDriver) -> Self {
342 self.driver = driver;
343 self
344 }
345}
346
347#[derive(Clone, Copy, Debug, PartialEq, Eq)]
358pub struct QrOptions {
359 pub gauge: QrGauge,
361}
362
363impl Default for QrOptions {
364 fn default() -> Self {
365 Self {
366 gauge: QrGauge::Raw,
367 }
368 }
369}
370
371impl QrOptions {
372 pub fn gauge(mut self, gauge: QrGauge) -> Self {
383 self.gauge = gauge;
384 self
385 }
386}
387
388pub(crate) fn validate_derivative_eps(
389 op: &'static str,
390 derivative_eps: f64,
391) -> tenferro_tensor::Result<()> {
392 if derivative_eps.is_finite() && derivative_eps > 0.0 {
393 Ok(())
394 } else {
395 Err(Error::invalid_argument(
396 op,
397 "derivative_eps",
398 format!("must be positive and finite, got {derivative_eps}"),
399 ))
400 }
401}
402
403#[derive(Clone, Copy, Debug, PartialEq)]
404#[doc(hidden)]
405#[allow(dead_code)]
406pub(crate) enum LinalgOp {
407 Cholesky,
408 Lu,
409 LuFactor,
410 LuSolvePrepared {
411 transpose_a: bool,
412 conjugate_a: bool,
413 },
414 SignDetFromLuFactor,
415 LogAbsDetFromLuFactor,
416 FullPivLu,
417 FullPivLuSolve {
418 transpose_a: bool,
419 },
420 Solve,
424 LuFactorSolve,
434 Svd {
435 derivative_eps: f64,
436 gauge: SvdGauge,
437 driver: SvdDriver,
438 },
439 SvdFull,
443 SvdVals {
445 derivative_eps: f64,
446 driver: SvdDriver,
447 },
448 Qr {
449 gauge: QrGauge,
450 },
451 RankRevealingQr {
452 gauge: QrGauge,
453 rtol: f64,
454 atol: f64,
455 },
456 HouseholderQrFactor,
457 HouseholderQrFromFactors,
458 HouseholderQrAppend,
459 HouseholderQrR {
460 gauge: QrGauge,
461 },
462 HouseholderQrQColumns {
463 start: usize,
464 end: usize,
465 gauge: QrGauge,
466 },
467 HouseholderQrThinQ {
469 gauge: QrGauge,
470 },
471 HouseholderQrAppendTangent,
473 HouseholderQrSplitTangent {
475 right: bool,
476 },
477 Eigh {
478 derivative_eps: f64,
479 gauge: EighGauge,
480 driver: EighDriver,
481 },
482 EighVals {
484 derivative_eps: f64,
485 driver: EighDriver,
486 },
487 Eig {
488 input_dtype: DType,
489 },
490 EigVals {
491 input_dtype: DType,
492 },
493 TriangularSolve {
494 left_side: bool,
495 lower: bool,
496 transpose_a: bool,
497 unit_diagonal: bool,
498 },
499}
500
501impl LinalgOp {
502 fn output_count(self) -> usize {
503 match self {
504 Self::Cholesky
505 | Self::EighVals { .. }
506 | Self::EigVals { .. }
507 | Self::FullPivLuSolve { .. }
508 | Self::LogAbsDetFromLuFactor
509 | Self::LuSolvePrepared { .. }
510 | Self::SignDetFromLuFactor
511 | Self::Solve
512 | Self::SvdVals { .. }
513 | Self::TriangularSolve { .. } => 1,
514 Self::Svd { .. } | Self::SvdFull | Self::LuFactorSolve => 3,
515 Self::RankRevealingQr { .. } | Self::Lu => 4,
516 Self::Qr { .. }
517 | Self::HouseholderQrFactor
518 | Self::HouseholderQrFromFactors
519 | Self::HouseholderQrAppend
520 | Self::Eigh { .. }
521 | Self::Eig { .. } => 2,
522 Self::HouseholderQrR { .. }
523 | Self::HouseholderQrQColumns { .. }
524 | Self::HouseholderQrThinQ { .. }
525 | Self::HouseholderQrAppendTangent
526 | Self::HouseholderQrSplitTangent { .. } => 1,
527 Self::LuFactor => 3,
528 Self::FullPivLu => 5,
529 }
530 }
531
532 fn input_count(self) -> usize {
533 match self {
534 Self::FullPivLuSolve { .. }
535 | Self::Solve
536 | Self::LuFactorSolve
537 | Self::TriangularSolve { .. }
538 | Self::HouseholderQrFromFactors
539 | Self::HouseholderQrR { .. }
540 | Self::HouseholderQrQColumns { .. }
541 | Self::HouseholderQrThinQ { .. } => 2,
542 Self::LogAbsDetFromLuFactor => 2,
543 Self::SignDetFromLuFactor | Self::HouseholderQrAppend => 3,
544 Self::HouseholderQrAppendTangent => 4,
545 Self::HouseholderQrSplitTangent { .. } => 3,
546 Self::LuSolvePrepared { .. } => 4,
547 _ => 1,
548 }
549 }
550
551 fn tag(self) -> u8 {
552 match self {
553 Self::Cholesky => 0,
554 Self::Lu => 1,
555 Self::FullPivLu => 2,
556 Self::FullPivLuSolve { .. } => 3,
557 Self::Svd { .. } => 4,
558 Self::Qr { .. } => 5,
559 Self::Eigh { .. } => 6,
560 Self::Eig { .. } => 7,
561 Self::TriangularSolve { .. } => 9,
562 Self::LuFactor => 10,
563 Self::LuSolvePrepared { .. } => 11,
564 Self::SvdVals { .. } => 12,
565 Self::EighVals { .. } => 13,
566 Self::EigVals { .. } => 14,
567 Self::SvdFull => 15,
568 Self::LogAbsDetFromLuFactor => 16,
569 Self::SignDetFromLuFactor => 17,
570 Self::Solve => 18,
571 Self::HouseholderQrFactor => 19,
572 Self::HouseholderQrFromFactors => 20,
573 Self::HouseholderQrAppend => 21,
574 Self::HouseholderQrR { .. } => 22,
575 Self::HouseholderQrQColumns { .. } => 23,
576 Self::HouseholderQrThinQ { .. } => 24,
577 Self::HouseholderQrAppendTangent => 25,
578 Self::HouseholderQrSplitTangent { .. } => 26,
579 Self::RankRevealingQr { .. } => 27,
580 Self::LuFactorSolve => 28,
581 }
582 }
583}
584
585#[derive(Clone, Debug, PartialEq)]
586#[doc(hidden)]
587pub(crate) struct LinalgExtensionOp {
588 op: LinalgOp,
589}
590
591impl LinalgExtensionOp {
592 pub(crate) fn new(op: LinalgOp) -> Self {
593 Self { op }
594 }
595
596 pub(crate) fn op(&self) -> LinalgOp {
597 self.op
598 }
599}
600
601impl ExtensionOp for LinalgExtensionOp {
602 fn family_id(&self) -> &'static str {
603 LINALG_EXTENSION_FAMILY_ID
604 }
605
606 fn payload_hash(&self, hasher: &mut dyn Hasher) {
607 hasher.write_u8(self.op.tag());
608 match self.op {
609 LinalgOp::Svd {
610 derivative_eps,
611 gauge,
612 driver,
613 } => {
614 hasher.write_u64(derivative_eps.to_bits());
615 hash_svd_gauge(hasher, gauge);
616 hash_svd_driver(hasher, driver);
617 }
618 LinalgOp::SvdVals {
619 derivative_eps,
620 driver,
621 } => {
622 hasher.write_u64(derivative_eps.to_bits());
623 hash_svd_driver(hasher, driver);
624 }
625 LinalgOp::EighVals {
626 derivative_eps,
627 driver,
628 } => {
629 hasher.write_u64(derivative_eps.to_bits());
630 hash_eigh_driver(hasher, driver);
631 }
632 LinalgOp::Qr { gauge }
633 | LinalgOp::HouseholderQrR { gauge }
634 | LinalgOp::HouseholderQrThinQ { gauge } => {
635 hash_qr_gauge(hasher, gauge);
636 }
637 LinalgOp::RankRevealingQr { gauge, rtol, atol } => {
638 hash_qr_gauge(hasher, gauge);
639 hasher.write_u64(rtol.to_bits());
640 hasher.write_u64(atol.to_bits());
641 }
642 LinalgOp::HouseholderQrQColumns { start, end, gauge } => {
643 hasher.write_usize(start);
644 hasher.write_usize(end);
645 hash_qr_gauge(hasher, gauge);
646 }
647 LinalgOp::Eigh {
648 derivative_eps,
649 gauge,
650 driver,
651 } => {
652 hasher.write_u64(derivative_eps.to_bits());
653 hash_eigh_gauge(hasher, gauge);
654 hash_eigh_driver(hasher, driver);
655 }
656 LinalgOp::Eig { input_dtype } | LinalgOp::EigVals { input_dtype } => {
657 hash_dtype(hasher, input_dtype);
658 }
659 LinalgOp::FullPivLuSolve { transpose_a }
660 | LinalgOp::HouseholderQrSplitTangent { right: transpose_a } => {
661 hasher.write_u8(u8::from(transpose_a));
662 }
663 LinalgOp::LuSolvePrepared {
664 transpose_a,
665 conjugate_a,
666 } => {
667 hasher.write_u8(u8::from(transpose_a));
668 hasher.write_u8(u8::from(conjugate_a));
669 }
670 LinalgOp::TriangularSolve {
671 left_side,
672 lower,
673 transpose_a,
674 unit_diagonal,
675 } => {
676 hasher.write_u8(u8::from(left_side));
677 hasher.write_u8(u8::from(lower));
678 hasher.write_u8(u8::from(transpose_a));
679 hasher.write_u8(u8::from(unit_diagonal));
680 }
681 LinalgOp::Cholesky
682 | LinalgOp::Lu
683 | LinalgOp::LuFactor
684 | LinalgOp::LogAbsDetFromLuFactor
685 | LinalgOp::SignDetFromLuFactor
686 | LinalgOp::FullPivLu
687 | LinalgOp::SvdFull
688 | LinalgOp::Solve
689 | LinalgOp::LuFactorSolve
690 | LinalgOp::HouseholderQrFactor
691 | LinalgOp::HouseholderQrFromFactors
692 | LinalgOp::HouseholderQrAppend
693 | LinalgOp::HouseholderQrAppendTangent => {}
694 }
695 }
696
697 fn payload_eq(&self, other: &dyn ExtensionOp) -> bool {
698 other
699 .as_any()
700 .downcast_ref::<Self>()
701 .is_some_and(|that| self == that)
702 }
703
704 fn clone_arc(&self) -> Arc<dyn ExtensionOp> {
705 Arc::new(self.clone())
706 }
707
708 fn as_any(&self) -> &dyn Any {
709 self
710 }
711
712 fn input_count(&self) -> usize {
713 self.op.input_count()
714 }
715
716 fn output_count(&self) -> usize {
717 self.op.output_count()
718 }
719
720 fn semantic_effects(&self) -> tenferro_ops::ext_op::ExtensionEffectDeclaration<'_> {
721 tenferro_ops::ext_op::ExtensionEffectDeclaration::Declared(&[])
722 }
723
724 fn semantic_aliases(&self) -> tenferro_ops::ext_op::ExtensionAliasDeclaration<'_> {
725 tenferro_ops::ext_op::ExtensionAliasDeclaration::AllFresh
726 }
727
728 fn prune_outputs(&self, live_outputs: &[bool]) -> Option<Arc<dyn ExtensionOp>> {
729 match self.op {
730 LinalgOp::Svd {
731 derivative_eps,
732 driver,
733 ..
734 } if live_outputs == [false, true, false] => {
735 Some(Arc::new(Self::new(LinalgOp::SvdVals {
736 derivative_eps,
737 driver,
738 })))
739 }
740 LinalgOp::Eigh {
741 derivative_eps,
742 driver,
743 ..
744 } if live_outputs == [true, false] => Some(Arc::new(Self::new(LinalgOp::EighVals {
745 derivative_eps,
746 driver,
747 }))),
748 LinalgOp::Eig { input_dtype } if live_outputs == [true, false] => {
749 Some(Arc::new(Self::new(LinalgOp::EigVals { input_dtype })))
750 }
751 _ => None,
759 }
760 }
761
762 fn infer_output_meta(
763 &self,
764 ctx: &mut tenferro_ops::ExtensionShapeContext<'_>,
765 ) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
766 let input_dtypes = (0..self.input_count())
767 .map(|input| ctx.input_dtype(input))
768 .collect::<Result<Vec<_>, _>>()?;
769 let input_shapes = (0..self.input_count())
770 .map(|input| ctx.input_shape(input))
771 .collect::<Result<Vec<_>, _>>()?;
772 let metas = match self.op {
773 LinalgOp::Cholesky => {
774 require_matrix_meta("tenferro-linalg.cholesky", input_shapes[0])?;
775 vec![(promote_dtypes(&input_dtypes), input_shapes[0].to_vec())]
776 }
777 LinalgOp::FullPivLuSolve { .. } => {
778 require_matrix_meta("tenferro-linalg.full_piv_lu_solve", input_shapes[0])?;
779 require_matrix_meta("tenferro-linalg.full_piv_lu_solve", input_shapes[1])?;
780 vec![(promote_dtypes(&input_dtypes), input_shapes[1].to_vec())]
781 }
782 LinalgOp::Solve => {
783 require_matrix_meta("tenferro-linalg.solve", input_shapes[0])?;
784 require_matrix_meta("tenferro-linalg.solve", input_shapes[1])?;
785 vec![(promote_dtypes(&input_dtypes), input_shapes[1].to_vec())]
786 }
787 LinalgOp::LuFactorSolve => {
788 require_matrix_meta("tenferro-linalg.lu_factor_solve", input_shapes[1])?;
789 let mut factors = lu_factor_meta(input_dtypes[0], input_shapes[0])?.into_iter();
790 let (Some(packed_lu), Some(pivots)) = (factors.next(), factors.next()) else {
791 return Err(Error::Internal(
792 "lu_factor_solve: lu_factor metadata returned fewer than two outputs"
793 .into(),
794 ));
795 };
796 vec![
797 (promote_dtypes(&input_dtypes), input_shapes[1].to_vec()),
798 packed_lu,
799 pivots,
800 ]
801 }
802 LinalgOp::TriangularSolve { .. } => {
803 require_matrix_meta("tenferro-linalg.triangular_solve", input_shapes[0])?;
804 require_matrix_meta("tenferro-linalg.triangular_solve", input_shapes[1])?;
805 vec![(promote_dtypes(&input_dtypes), input_shapes[1].to_vec())]
806 }
807 LinalgOp::LuSolvePrepared { .. } => {
808 require_matrix_meta("tenferro-linalg.lu_solve_prepared_lu", input_shapes[0])?;
809 require_matrix_meta("tenferro-linalg.lu_solve_prepared_rhs", input_shapes[3])?;
810 vec![(
811 promote_dtypes(&[input_dtypes[0], input_dtypes[3]]),
812 input_shapes[3].to_vec(),
813 )]
814 }
815 LinalgOp::Lu => lu_meta(input_dtypes[0], input_shapes[0])?,
816 LinalgOp::LuFactor => lu_factor_meta(input_dtypes[0], input_shapes[0])?,
817 LinalgOp::SignDetFromLuFactor => {
818 vec![signdet_from_lu_factor_meta(
819 input_dtypes[0],
820 input_shapes[0],
821 input_shapes[1],
822 input_shapes[2],
823 )?]
824 }
825 LinalgOp::LogAbsDetFromLuFactor => {
826 vec![logabsdet_from_lu_factor_meta(
827 input_dtypes[0],
828 input_shapes[0],
829 input_shapes[1],
830 )?]
831 }
832 LinalgOp::FullPivLu => full_piv_lu_meta(input_dtypes[0], input_shapes[0])?,
833 LinalgOp::Svd { .. } => svd_meta(input_dtypes[0], input_shapes[0])?,
834 LinalgOp::SvdFull => svd_full_meta(input_dtypes[0], input_shapes[0])?,
835 LinalgOp::SvdVals { .. } => {
836 vec![svd_values_meta(input_dtypes[0], input_shapes[0])?]
837 }
838 LinalgOp::Qr { .. } => qr_meta(input_dtypes[0], input_shapes[0])?,
839 LinalgOp::RankRevealingQr { .. } => {
840 rank_revealing_qr_meta(input_dtypes[0], input_shapes[0])?
841 }
842 LinalgOp::HouseholderQrFactor => {
843 householder_qr_factor_meta(input_dtypes[0], input_shapes[0])?
844 }
845 LinalgOp::HouseholderQrFromFactors => {
846 householder_qr_from_factors_meta(&input_dtypes, &input_shapes)?
847 }
848 LinalgOp::HouseholderQrAppend => {
849 householder_qr_append_meta(&input_dtypes, &input_shapes)?
850 }
851 LinalgOp::HouseholderQrR { .. } => vec![householder_qr_r_meta(
852 &input_dtypes,
853 input_shapes[0],
854 input_shapes[1],
855 )?],
856 LinalgOp::HouseholderQrQColumns { start, end, .. } => {
857 vec![householder_qr_q_columns_meta(
858 &input_dtypes,
859 input_shapes[0],
860 input_shapes[1],
861 start,
862 end,
863 )?]
864 }
865 LinalgOp::HouseholderQrThinQ { .. } => {
866 vec![householder_qr_thin_q_meta(
867 &input_dtypes,
868 input_shapes[0],
869 input_shapes[1],
870 )?]
871 }
872 LinalgOp::HouseholderQrAppendTangent => {
873 vec![householder_qr_append_tangent_meta(
874 &input_dtypes,
875 &input_shapes,
876 )?]
877 }
878 LinalgOp::HouseholderQrSplitTangent { right } => {
879 vec![householder_qr_split_tangent_meta(
880 &input_dtypes,
881 &input_shapes,
882 right,
883 )?]
884 }
885 LinalgOp::Eigh { .. } => eigh_meta(input_dtypes[0], input_shapes[0])?,
886 LinalgOp::EighVals { .. } => vec![eigh_values_meta(input_dtypes[0], input_shapes[0])?],
887 LinalgOp::Eig { input_dtype } => eig_meta(input_dtype, input_shapes[0])?,
888 LinalgOp::EigVals { input_dtype } => {
889 vec![eig_values_meta(input_dtype, input_shapes[0])?]
890 }
891 };
892 Ok(metas)
893 }
894}
895
896fn execute_linalg_extension_reads_on_session<B: BackendSession + ?Sized>(
897 op: &LinalgExtensionOp,
898 inputs: &[TensorRead<'_>],
899 session: &mut B,
900) -> tenferro_tensor::Result<Vec<Tensor>> {
901 if let Some(result) = with_cpu_exec_session(session, |session| {
902 execute_linalg_extension_reads_in_session(op, inputs, session)
903 }) {
904 return result;
905 }
906 #[cfg(feature = "cuda")]
907 if let Some(result) = with_cuda_exec_session(session, |session| {
908 execute_linalg_extension_reads_in_session(op, inputs, session)
909 }) {
910 return result;
911 }
912 Err(Error::unsupported(
913 "linalg_extension",
914 "selected backend session does not expose a linalg execution capability",
915 ))
916}
917
918fn execute_linalg_extension_reads_in_session<S: LinalgBackend>(
919 op: &LinalgExtensionOp,
920 inputs: &[TensorRead<'_>],
921 session: &mut S,
922) -> tenferro_tensor::Result<Vec<Tensor>> {
923 if op.op() == LinalgOp::HouseholderQrAppendTangent {
924 let left = session.to_contiguous_read(inputs[0].clone())?;
925 let right = session.to_contiguous_read(inputs[1].clone())?;
926 return Ok(vec![session.concatenate(&[&left, &right], 1)?]);
927 }
928 if let LinalgOp::HouseholderQrSplitTangent { right } = op.op() {
929 let cotangent_shape = inputs[0].clone().tensor_view().shape().to_vec();
930 let left_shape = inputs[1].clone().tensor_view().shape().to_vec();
931 let right_shape = inputs[2].clone().tensor_view().shape().to_vec();
932 let config =
933 householder_qr_split_config(&cotangent_shape, &left_shape, &right_shape, right)?;
934 let cotangent = session.to_contiguous_read(inputs[0].clone())?;
935 return Ok(vec![session.slice(&cotangent, &config)?]);
936 }
937 match op.op() {
941 LinalgOp::Cholesky => return Ok(vec![session.cholesky_read(inputs[0].clone())?]),
942 LinalgOp::Lu => return session.lu_read(inputs[0].clone()),
943 LinalgOp::FullPivLu => return session.full_piv_lu_read(inputs[0].clone()),
944 LinalgOp::Svd {
945 derivative_eps,
946 gauge,
947 driver,
948 } => {
949 return session.svd_with_options_read(
950 inputs[0].clone(),
951 SvdOptions {
952 derivative_eps,
953 gauge,
954 driver,
955 },
956 );
957 }
958 LinalgOp::SvdFull => return session.svd_full_read(inputs[0].clone()),
959 LinalgOp::SvdVals { driver, .. } => {
960 return Ok(vec![
961 session.svd_values_with_driver_read(inputs[0].clone(), driver)?
962 ]);
963 }
964 LinalgOp::Qr { gauge } => {
965 return session.qr_with_options_read(inputs[0].clone(), QrOptions { gauge });
966 }
967 LinalgOp::RankRevealingQr { gauge, rtol, atol } => {
968 return session.rank_revealing_qr_read(
969 inputs[0].clone(),
970 RankRevealingQrOptions { gauge, rtol, atol },
971 );
972 }
973 LinalgOp::Eigh {
974 derivative_eps,
975 gauge,
976 driver,
977 } => {
978 return session.eigh_with_options_read(
979 inputs[0].clone(),
980 EighOptions {
981 derivative_eps,
982 gauge,
983 driver,
984 },
985 );
986 }
987 LinalgOp::EighVals { driver, .. } => {
988 return Ok(vec![
989 session.eigh_values_with_driver_read(inputs[0].clone(), driver)?
990 ]);
991 }
992 LinalgOp::Eig { .. } => return session.eig_read(inputs[0].clone()),
993 LinalgOp::EigVals { .. } => return Ok(vec![session.eig_values_read(inputs[0].clone())?]),
994 LinalgOp::Solve => match session.solve_read(inputs[0].clone(), inputs[1].clone()) {
995 Ok(output) => return Ok(vec![output]),
996 Err(error) if error.kind() == ErrorKind::Unsupported => {}
997 Err(error) => return Err(error),
998 },
999 _ => {}
1000 }
1001 if let LinalgOp::TriangularSolve {
1002 left_side,
1003 lower,
1004 transpose_a,
1005 unit_diagonal,
1006 } = op.op()
1007 {
1008 match session.triangular_solve_read(
1009 inputs[0].clone(),
1010 inputs[1].clone(),
1011 left_side,
1012 lower,
1013 transpose_a,
1014 unit_diagonal,
1015 ) {
1016 Ok(output) => return Ok(vec![output]),
1017 Err(error) if error.kind() == ErrorKind::Unsupported => {}
1018 Err(error) => return Err(error),
1019 }
1020 }
1021
1022 let materialized_inputs = inputs
1025 .iter()
1026 .filter(|input| input.as_tensor().is_none())
1027 .cloned()
1028 .map(|input| session.to_contiguous_read(input))
1029 .collect::<tenferro_tensor::Result<Vec<_>>>()?;
1030 let mut views = materialized_inputs.iter();
1031 let input_refs: Vec<&Tensor> = inputs
1034 .iter()
1035 .filter_map(|input| input.as_tensor().or_else(|| views.next()))
1036 .collect();
1037 execute_linalg(op.op(), &input_refs, session)
1038}
1039
1040fn linalg_session_supported<B: tenferro_tensor::TensorBackend + 'static>(
1041 #[cfg_attr(not(feature = "cuda"), allow(unused_variables))] op: &LinalgExtensionOp,
1042) -> bool {
1043 let type_id = std::any::TypeId::of::<B>();
1049 if type_id == std::any::TypeId::of::<tenferro_cpu::CpuBackend>() {
1050 return true;
1054 }
1055 #[cfg(feature = "cuda")]
1056 {
1057 if type_id == std::any::TypeId::of::<tenferro_gpu::cuda::CudaBackend>() {
1058 return match op.op() {
1059 LinalgOp::FullPivLu | LinalgOp::FullPivLuSolve { .. } => false,
1061 LinalgOp::Eig { .. } | LinalgOp::EigVals { .. } => false,
1062 LinalgOp::Solve => true,
1067 LinalgOp::LuFactorSolve => true,
1070 LinalgOp::LuSolvePrepared {
1072 transpose_a: false,
1073 conjugate_a: true,
1074 } => false,
1075 _ => true,
1076 };
1077 }
1078 }
1079 false
1080}
1081
1082fn execute_linalg_extension_in_session(
1083 op: &LinalgExtensionOp,
1084 session: &mut dyn BackendSession,
1085 _extension_caches: &mut tenferro_runtime::ExtensionCacheStore,
1086 inputs: &[TensorRead<'_>],
1087) -> tenferro_tensor::Result<Vec<Tensor>> {
1088 execute_linalg_extension_reads_on_session(op, inputs, session)
1092}
1093
1094define_extension_runtime! {
1095 runtime = LinalgRuntime,
1096 family_id = LINALG_EXTENSION_FAMILY_ID,
1097 op_type = LinalgExtensionOp,
1098 execute_in_session = execute_linalg_extension_in_session,
1099 session_supported = linalg_session_supported,
1100 backend_bound = TensorBackend,
1101}
1102
1103fn execute_linalg<B: LinalgBackend>(
1104 op: LinalgOp,
1105 inputs: &[&Tensor],
1106 backend: &mut B,
1107) -> tenferro_tensor::Result<Vec<Tensor>> {
1108 match op {
1109 LinalgOp::Cholesky => Ok(vec![backend.cholesky(inputs[0])?]),
1110 LinalgOp::Lu => backend.lu(inputs[0]),
1111 LinalgOp::LuFactor => backend.lu_factor(inputs[0]),
1112 LinalgOp::SignDetFromLuFactor => {
1113 Ok(vec![signdet_from_lu_factor(inputs[1], inputs[2], backend)?])
1114 }
1115 LinalgOp::LogAbsDetFromLuFactor => Ok(vec![logabsdet_from_lu_factor(inputs[1], backend)?]),
1116 LinalgOp::LuSolvePrepared {
1117 transpose_a,
1118 conjugate_a,
1119 } => Ok(vec![backend.lu_solve_prepared(
1120 inputs[0],
1121 inputs[1],
1122 inputs[2],
1123 inputs[3],
1124 transpose_a,
1125 conjugate_a,
1126 )?]),
1127 LinalgOp::FullPivLu => backend.full_piv_lu(inputs[0]),
1128 LinalgOp::FullPivLuSolve { transpose_a } => Ok(vec![backend.full_piv_lu_solve(
1129 inputs[0],
1130 inputs[1],
1131 transpose_a,
1132 )?]),
1133 LinalgOp::Solve => Ok(vec![backend.solve(inputs[0], inputs[1])?]),
1134 LinalgOp::LuFactorSolve => backend.lu_factor_solve(inputs[0], inputs[1]),
1135 LinalgOp::Svd {
1136 derivative_eps,
1137 gauge,
1138 driver,
1139 } => backend.svd_with_options(
1140 inputs[0],
1141 SvdOptions {
1142 derivative_eps,
1143 gauge,
1144 driver,
1145 },
1146 ),
1147 LinalgOp::SvdFull => backend.svd_full(inputs[0]),
1148 LinalgOp::SvdVals { driver, .. } => {
1149 Ok(vec![backend.svd_values_with_driver(inputs[0], driver)?])
1150 }
1151 LinalgOp::Qr { gauge } => backend.qr_with_options(inputs[0], QrOptions { gauge }),
1152 LinalgOp::RankRevealingQr { gauge, rtol, atol } => {
1153 backend.rank_revealing_qr(inputs[0], RankRevealingQrOptions { gauge, rtol, atol })
1154 }
1155 LinalgOp::HouseholderQrFactor => {
1156 let state = backend.householder_qr(inputs[0])?;
1157 Ok(vec![state.packed, state.coeff])
1158 }
1159 LinalgOp::HouseholderQrFromFactors => {
1160 let state = backend.householder_qr_from_factors(inputs[0], inputs[1])?;
1161 Ok(vec![state.packed, state.coeff])
1162 }
1163 LinalgOp::HouseholderQrAppend => {
1164 let state = backend.householder_qr_append(inputs[0], inputs[1], inputs[2])?;
1165 Ok(vec![state.packed, state.coeff])
1166 }
1167 LinalgOp::HouseholderQrR { gauge } => Ok(vec![backend.householder_qr_r(
1168 inputs[0],
1169 inputs[1],
1170 QrOptions { gauge },
1171 )?]),
1172 LinalgOp::HouseholderQrQColumns { start, end, gauge } => Ok(vec![backend
1173 .householder_qr_q_columns(inputs[0], inputs[1], start..end, QrOptions { gauge })?]),
1174 LinalgOp::HouseholderQrThinQ { gauge } => {
1175 let end = inputs[1].shape().first().copied().ok_or_else(|| {
1176 Error::rank_mismatch("tenferro-linalg.householder_qr_thin_q", 1, 0)
1177 })?;
1178 Ok(vec![backend.householder_qr_q_columns(
1179 inputs[0],
1180 inputs[1],
1181 0..end,
1182 QrOptions { gauge },
1183 )?])
1184 }
1185 LinalgOp::HouseholderQrAppendTangent => {
1186 Ok(vec![backend.concatenate(&[inputs[0], inputs[1]], 1)?])
1187 }
1188 LinalgOp::HouseholderQrSplitTangent { right } => {
1189 let config = householder_qr_split_config(
1190 inputs[0].shape(),
1191 inputs[1].shape(),
1192 inputs[2].shape(),
1193 right,
1194 )?;
1195 Ok(vec![backend.slice(inputs[0], &config)?])
1196 }
1197 LinalgOp::Eigh {
1198 derivative_eps,
1199 gauge,
1200 driver,
1201 } => backend.eigh_with_options(
1202 inputs[0],
1203 EighOptions {
1204 derivative_eps,
1205 gauge,
1206 driver,
1207 },
1208 ),
1209 LinalgOp::EighVals { driver, .. } => {
1210 Ok(vec![backend.eigh_values_with_driver(inputs[0], driver)?])
1211 }
1212 LinalgOp::Eig { .. } => backend.eig(inputs[0]),
1213 LinalgOp::EigVals { .. } => Ok(vec![backend.eig_values(inputs[0])?]),
1214 LinalgOp::TriangularSolve {
1215 left_side,
1216 lower,
1217 transpose_a,
1218 unit_diagonal,
1219 } => Ok(vec![backend.triangular_solve(
1220 inputs[0],
1221 inputs[1],
1222 left_side,
1223 lower,
1224 transpose_a,
1225 unit_diagonal,
1226 )?]),
1227 }
1228}
1229
1230fn signdet_from_lu_factor<B: LinalgBackend + ?Sized>(
1238 packed_lu: &Tensor,
1239 parity: &Tensor,
1240 backend: &mut B,
1241) -> tenferro_tensor::Result<Tensor> {
1242 let diag = backend.extract_diagonal(packed_lu, 0, 1)?;
1243 let sign_diag = backend.sign_read(TensorRead::from_tensor(&diag))?;
1244 let sign_u = backend.reduce_prod_read(TensorRead::from_tensor(&sign_diag), &[0])?;
1245 backend.mul_read(
1246 TensorRead::from_tensor(parity),
1247 TensorRead::from_tensor(&sign_u),
1248 )
1249}
1250
1251fn logabsdet_from_lu_factor<B: LinalgBackend + ?Sized>(
1252 packed_lu: &Tensor,
1253 backend: &mut B,
1254) -> tenferro_tensor::Result<Tensor> {
1255 let diag = backend.extract_diagonal(packed_lu, 0, 1)?;
1256 let abs = backend.abs_read(TensorRead::from_tensor(&diag))?;
1257 let log = backend.log_read(TensorRead::from_tensor(&abs))?;
1258 backend.reduce_sum_read(TensorRead::from_tensor(&log), &[0])
1259}
1260
1261pub(crate) fn apply_svd_gauge(
1262 gauge: SvdGauge,
1263 outputs: &mut [Tensor],
1264) -> tenferro_tensor::Result<()> {
1265 match gauge {
1266 SvdGauge::Raw => Ok(()),
1267 SvdGauge::CanonicalPivot => apply_canonical_pivot_svd_gauge(outputs),
1268 }
1269}
1270
1271fn apply_canonical_pivot_svd_gauge(outputs: &mut [Tensor]) -> tenferro_tensor::Result<()> {
1272 if outputs.len() != 3 {
1273 return Err(Error::invalid_argument(
1274 "tenferro-linalg.svd",
1275 "outputs",
1276 format!(
1277 "canonical SVD gauge expected three outputs, got {}",
1278 outputs.len()
1279 ),
1280 ));
1281 }
1282
1283 let (u_slice, rest) = outputs.split_at_mut(1);
1284 let (singular_slice, vt_slice) = rest.split_at_mut(1);
1285 let u = &mut u_slice[0];
1286 let singular_values = &singular_slice[0];
1287 let vt = &mut vt_slice[0];
1288 let u_shape = u.shape().to_vec();
1289 let s_shape = singular_values.shape().to_vec();
1290 let vt_shape = vt.shape().to_vec();
1291 if u_shape.len() < 2 || vt_shape.len() < 2 || s_shape.is_empty() {
1292 return Err(Error::invalid_argument(
1293 "tenferro-linalg.svd",
1294 "outputs",
1295 format!(
1296 "canonical SVD gauge expected U rank >= 2, S rank >= 1, VT rank >= 2; got U={u_shape:?}, S={s_shape:?}, VT={vt_shape:?}"
1297 ),
1298 ));
1299 }
1300
1301 let m = u_shape[0];
1302 let k = u_shape[1];
1303 let n = vt_shape[1];
1304 if s_shape[0] != k
1305 || vt_shape[0] != k
1306 || u_shape[2..] != vt_shape[2..]
1307 || s_shape[1..] != u_shape[2..]
1308 {
1309 return Err(Error::invalid_argument(
1310 "tenferro-linalg.svd",
1311 "outputs",
1312 format!(
1313 "canonical SVD gauge expected compatible compact SVD shapes, got U={u_shape:?}, S={s_shape:?}, VT={vt_shape:?}"
1314 ),
1315 ));
1316 }
1317 let layout = canonical_svd_gauge_layout(m, k, n, &u_shape[2..])?;
1318
1319 match (u.dtype(), vt.dtype()) {
1320 (tenferro_tensor::DType::F64, tenferro_tensor::DType::F64) => {
1321 let (u, vt) = svd_gauge_pair_mut::<f64>(u, vt)?;
1322 canonicalize_svd_gauge_f64(u.host_data_mut()?, vt.host_data_mut()?, layout)
1323 }
1324 (tenferro_tensor::DType::F32, tenferro_tensor::DType::F32) => {
1325 let (u, vt) = svd_gauge_pair_mut::<f32>(u, vt)?;
1326 canonicalize_svd_gauge_f32(u.host_data_mut()?, vt.host_data_mut()?, layout)
1327 }
1328 (tenferro_tensor::DType::C64, tenferro_tensor::DType::C64) => {
1329 let (u, vt) = svd_gauge_pair_mut::<Complex64>(u, vt)?;
1330 canonicalize_svd_gauge_c64(u.host_data_mut()?, vt.host_data_mut()?, layout)
1331 }
1332 (tenferro_tensor::DType::C32, tenferro_tensor::DType::C32) => {
1333 let (u, vt) = svd_gauge_pair_mut::<Complex32>(u, vt)?;
1334 canonicalize_svd_gauge_c32(u.host_data_mut()?, vt.host_data_mut()?, layout)
1335 }
1336 (u_dtype, vt_dtype) => Err(Error::dtype_mismatch(
1337 "tenferro-linalg.svd",
1338 u_dtype,
1339 vt_dtype,
1340 )),
1341 }
1342}
1343
1344fn svd_gauge_pair_mut<'a, T: tenferro_tensor::TensorScalar>(
1346 u: &'a mut Tensor,
1347 vt: &'a mut Tensor,
1348) -> tenferro_tensor::Result<(
1349 &'a mut tenferro_tensor::TypedTensor<T>,
1350 &'a mut tenferro_tensor::TypedTensor<T>,
1351)> {
1352 let (u_dtype, vt_dtype) = (u.dtype(), vt.dtype());
1353 let u_t = u
1354 .as_typed_mut::<T>()
1355 .ok_or_else(|| Error::dtype_mismatch("tenferro-linalg.svd", u_dtype, vt_dtype))?;
1356 let vt_t = vt
1357 .as_typed_mut::<T>()
1358 .ok_or_else(|| Error::dtype_mismatch("tenferro-linalg.svd", u_dtype, vt_dtype))?;
1359 Ok((u_t, vt_t))
1360}
1361
1362#[derive(Clone, Copy, Debug, PartialEq, Eq)]
1363struct CanonicalSvdGaugeLayout {
1364 m: usize,
1365 k: usize,
1366 batch_count: usize,
1367 u_batch_len: usize,
1368 vt_batch_len: usize,
1369 u_len: usize,
1370 vt_len: usize,
1371}
1372
1373impl CanonicalSvdGaugeLayout {
1374 fn validate_storage(self, u_len: usize, vt_len: usize) -> tenferro_tensor::Result<()> {
1375 if u_len != self.u_len {
1376 return Err(Error::invalid_argument(
1377 "tenferro-linalg.svd",
1378 "U storage",
1379 format!(
1380 "canonical SVD gauge expected U storage length {}, got {u_len}",
1381 self.u_len
1382 ),
1383 ));
1384 }
1385 if vt_len != self.vt_len {
1386 return Err(Error::invalid_argument(
1387 "tenferro-linalg.svd",
1388 "VT storage",
1389 format!(
1390 "canonical SVD gauge expected VT storage length {}, got {vt_len}",
1391 self.vt_len
1392 ),
1393 ));
1394 }
1395 Ok(())
1396 }
1397}
1398
1399fn canonical_svd_gauge_layout(
1400 m: usize,
1401 k: usize,
1402 n: usize,
1403 batch_shape: &[usize],
1404) -> tenferro_tensor::Result<CanonicalSvdGaugeLayout> {
1405 let batch_count = tenferro_tensor::validate::checked_shape_product(
1406 "tenferro-linalg.svd",
1407 "canonical SVD batch",
1408 batch_shape,
1409 )?;
1410 let u_batch_len = tenferro_tensor::validate::checked_shape_product(
1411 "tenferro-linalg.svd",
1412 "canonical SVD U batch",
1413 &[m, k],
1414 )?;
1415 let vt_batch_len = tenferro_tensor::validate::checked_shape_product(
1416 "tenferro-linalg.svd",
1417 "canonical SVD VT batch",
1418 &[k, n],
1419 )?;
1420 let u_len = tenferro_tensor::validate::checked_shape_product(
1421 "tenferro-linalg.svd",
1422 "canonical SVD U storage",
1423 &[u_batch_len, batch_count],
1424 )?;
1425 let vt_len = tenferro_tensor::validate::checked_shape_product(
1426 "tenferro-linalg.svd",
1427 "canonical SVD VT storage",
1428 &[vt_batch_len, batch_count],
1429 )?;
1430 Ok(CanonicalSvdGaugeLayout {
1431 m,
1432 k,
1433 batch_count,
1434 u_batch_len,
1435 vt_batch_len,
1436 u_len,
1437 vt_len,
1438 })
1439}
1440
1441fn canonicalize_svd_gauge_f64(
1442 u: &mut [f64],
1443 vt: &mut [f64],
1444 layout: CanonicalSvdGaugeLayout,
1445) -> tenferro_tensor::Result<()> {
1446 layout.validate_storage(u.len(), vt.len())?;
1447 if layout.batch_count == 0 || layout.u_batch_len == 0 || layout.vt_batch_len == 0 {
1448 return Ok(());
1449 }
1450 for (u_batch, vt_batch) in u
1451 .chunks_exact_mut(layout.u_batch_len)
1452 .zip(vt.chunks_exact_mut(layout.vt_batch_len))
1453 {
1454 for (col, u_column) in u_batch.chunks_exact_mut(layout.m).enumerate() {
1455 let pivot = max_abs_pivot_f64(u_column);
1456 let pivot_value = u_column[pivot];
1457 if pivot_value < 0.0 {
1458 for value in u_column {
1459 *value = -*value;
1460 }
1461 for vt_column in vt_batch.chunks_exact_mut(layout.k) {
1462 vt_column[col] = -vt_column[col];
1463 }
1464 }
1465 }
1466 }
1467 Ok(())
1468}
1469
1470fn canonicalize_svd_gauge_f32(
1471 u: &mut [f32],
1472 vt: &mut [f32],
1473 layout: CanonicalSvdGaugeLayout,
1474) -> tenferro_tensor::Result<()> {
1475 layout.validate_storage(u.len(), vt.len())?;
1476 if layout.batch_count == 0 || layout.u_batch_len == 0 || layout.vt_batch_len == 0 {
1477 return Ok(());
1478 }
1479 for (u_batch, vt_batch) in u
1480 .chunks_exact_mut(layout.u_batch_len)
1481 .zip(vt.chunks_exact_mut(layout.vt_batch_len))
1482 {
1483 for (col, u_column) in u_batch.chunks_exact_mut(layout.m).enumerate() {
1484 let pivot = max_abs_pivot_f32(u_column);
1485 let pivot_value = u_column[pivot];
1486 if pivot_value < 0.0 {
1487 for value in u_column {
1488 *value = -*value;
1489 }
1490 for vt_column in vt_batch.chunks_exact_mut(layout.k) {
1491 vt_column[col] = -vt_column[col];
1492 }
1493 }
1494 }
1495 }
1496 Ok(())
1497}
1498
1499fn canonicalize_svd_gauge_c64(
1500 u: &mut [Complex64],
1501 vt: &mut [Complex64],
1502 layout: CanonicalSvdGaugeLayout,
1503) -> tenferro_tensor::Result<()> {
1504 layout.validate_storage(u.len(), vt.len())?;
1505 if layout.batch_count == 0 || layout.u_batch_len == 0 || layout.vt_batch_len == 0 {
1506 return Ok(());
1507 }
1508 for (u_batch, vt_batch) in u
1509 .chunks_exact_mut(layout.u_batch_len)
1510 .zip(vt.chunks_exact_mut(layout.vt_batch_len))
1511 {
1512 for (col, u_column) in u_batch.chunks_exact_mut(layout.m).enumerate() {
1513 let pivot = max_abs_pivot_c64(u_column);
1514 let pivot_value = u_column[pivot];
1515 let pivot_norm = pivot_value.norm();
1516 if pivot_norm == 0.0 {
1517 continue;
1518 }
1519 let phase = pivot_value.conj() / pivot_norm;
1520 let vt_phase = phase.conj();
1521 for value in u_column {
1522 *value *= phase;
1523 }
1524 for vt_column in vt_batch.chunks_exact_mut(layout.k) {
1525 vt_column[col] *= vt_phase;
1526 }
1527 }
1528 }
1529 Ok(())
1530}
1531
1532fn canonicalize_svd_gauge_c32(
1533 u: &mut [Complex32],
1534 vt: &mut [Complex32],
1535 layout: CanonicalSvdGaugeLayout,
1536) -> tenferro_tensor::Result<()> {
1537 layout.validate_storage(u.len(), vt.len())?;
1538 if layout.batch_count == 0 || layout.u_batch_len == 0 || layout.vt_batch_len == 0 {
1539 return Ok(());
1540 }
1541 for (u_batch, vt_batch) in u
1542 .chunks_exact_mut(layout.u_batch_len)
1543 .zip(vt.chunks_exact_mut(layout.vt_batch_len))
1544 {
1545 for (col, u_column) in u_batch.chunks_exact_mut(layout.m).enumerate() {
1546 let pivot = max_abs_pivot_c32(u_column);
1547 let pivot_value = u_column[pivot];
1548 let pivot_norm = pivot_value.norm();
1549 if pivot_norm == 0.0 {
1550 continue;
1551 }
1552 let phase = pivot_value.conj() / pivot_norm;
1553 let vt_phase = phase.conj();
1554 for value in u_column {
1555 *value *= phase;
1556 }
1557 for vt_column in vt_batch.chunks_exact_mut(layout.k) {
1558 vt_column[col] *= vt_phase;
1559 }
1560 }
1561 }
1562 Ok(())
1563}
1564
1565fn max_abs_pivot_f64(u_column: &[f64]) -> usize {
1566 let mut pivot = 0;
1567 let mut pivot_abs = u_column[0].abs();
1568 for (row, value) in u_column.iter().enumerate().skip(1) {
1569 let candidate_abs = value.abs();
1570 if candidate_abs > pivot_abs {
1571 pivot = row;
1572 pivot_abs = candidate_abs;
1573 }
1574 }
1575 pivot
1576}
1577
1578fn max_abs_pivot_f32(u_column: &[f32]) -> usize {
1579 let mut pivot = 0;
1580 let mut pivot_abs = u_column[0].abs();
1581 for (row, value) in u_column.iter().enumerate().skip(1) {
1582 let candidate_abs = value.abs();
1583 if candidate_abs > pivot_abs {
1584 pivot = row;
1585 pivot_abs = candidate_abs;
1586 }
1587 }
1588 pivot
1589}
1590
1591fn max_abs_pivot_c64(u_column: &[Complex64]) -> usize {
1592 let mut pivot = 0;
1593 let mut pivot_abs = u_column[0].norm_sqr();
1594 for (row, value) in u_column.iter().enumerate().skip(1) {
1595 let candidate_abs = value.norm_sqr();
1596 if candidate_abs > pivot_abs {
1597 pivot = row;
1598 pivot_abs = candidate_abs;
1599 }
1600 }
1601 pivot
1602}
1603
1604fn max_abs_pivot_c32(u_column: &[Complex32]) -> usize {
1605 let mut pivot = 0;
1606 let mut pivot_abs = u_column[0].norm_sqr();
1607 for (row, value) in u_column.iter().enumerate().skip(1) {
1608 let candidate_abs = value.norm_sqr();
1609 if candidate_abs > pivot_abs {
1610 pivot = row;
1611 pivot_abs = candidate_abs;
1612 }
1613 }
1614 pivot
1615}
1616
1617fn require_matrix_meta(op: &'static str, shape: &[SymDim]) -> tenferro_tensor::Result<()> {
1618 if shape.len() < 2 {
1619 return Err(Error::rank_mismatch(op, 2, shape.len()));
1620 }
1621 Ok(())
1622}
1623
1624fn matrix_meta_parts<'a>(
1625 op: &'static str,
1626 shape: &'a [SymDim],
1627) -> tenferro_tensor::Result<(SymDim, SymDim, &'a [SymDim])> {
1628 require_matrix_meta(op, shape)?;
1629 Ok((shape[0].clone(), shape[1].clone(), &shape[2..]))
1630}
1631
1632fn lu_meta(dtype: DType, shape: &[SymDim]) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1633 let (m, n, batch) = matrix_meta_parts("tenferro-linalg.lu", shape)?;
1634 let k = m.clone().min(n.clone());
1635 Ok(vec![
1636 (dtype, matrix_shape(m.clone(), m, batch)),
1637 (dtype, matrix_shape(shape[0].clone(), k.clone(), batch)),
1638 (dtype, matrix_shape(k, n, batch)),
1639 (dtype, batch.to_vec()),
1640 ])
1641}
1642
1643fn lu_factor_meta(
1644 dtype: DType,
1645 shape: &[SymDim],
1646) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1647 let (m, n, batch) = matrix_meta_parts("tenferro-linalg.lu_factor", shape)?;
1648 let k = m.min(n);
1649 Ok(vec![
1650 (dtype, shape.to_vec()),
1651 (DType::I32, vector_shape(k, batch)),
1652 (dtype, batch.to_vec()),
1653 ])
1654}
1655
1656fn signdet_from_lu_factor_meta(
1657 input_dtype: DType,
1658 input_shape: &[SymDim],
1659 packed_shape: &[SymDim],
1660 parity_shape: &[SymDim],
1661) -> tenferro_tensor::Result<(DType, Vec<SymDim>)> {
1662 let (_, _, batch) = matrix_meta_parts("tenferro-linalg.signdet_from_lu_factor", input_shape)?;
1663 require_matrix_meta(
1664 "tenferro-linalg.signdet_from_lu_factor_packed",
1665 packed_shape,
1666 )?;
1667 if parity_shape.len() != batch.len() {
1668 return Err(Error::rank_mismatch(
1669 "tenferro-linalg.signdet_from_lu_factor_parity",
1670 batch.len(),
1671 parity_shape.len(),
1672 ));
1673 }
1674 Ok((input_dtype, batch.to_vec()))
1675}
1676
1677fn logabsdet_from_lu_factor_meta(
1678 input_dtype: DType,
1679 input_shape: &[SymDim],
1680 packed_shape: &[SymDim],
1681) -> tenferro_tensor::Result<(DType, Vec<SymDim>)> {
1682 let (_, _, batch) = matrix_meta_parts("tenferro-linalg.logabsdet_from_lu_factor", input_shape)?;
1683 require_matrix_meta(
1684 "tenferro-linalg.logabsdet_from_lu_factor_packed",
1685 packed_shape,
1686 )?;
1687 Ok((singular_values_dtype(input_dtype), batch.to_vec()))
1688}
1689
1690fn full_piv_lu_meta(
1691 dtype: DType,
1692 shape: &[SymDim],
1693) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1694 let (n, _, batch) = matrix_meta_parts("tenferro-linalg.full_piv_lu", shape)?;
1695 Ok(vec![
1696 (dtype, matrix_shape(n.clone(), n.clone(), batch)),
1697 (dtype, matrix_shape(n.clone(), n.clone(), batch)),
1698 (dtype, matrix_shape(n.clone(), n.clone(), batch)),
1699 (dtype, matrix_shape(n.clone(), n, batch)),
1700 (singular_values_dtype(dtype), batch.to_vec()),
1701 ])
1702}
1703
1704fn svd_meta(dtype: DType, shape: &[SymDim]) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1705 let (m, n, batch) = matrix_meta_parts("tenferro-linalg.svd", shape)?;
1706 let k = m.clone().min(n.clone());
1707 Ok(vec![
1708 (dtype, matrix_shape(m, k.clone(), batch)),
1709 (singular_values_dtype(dtype), vector_shape(k.clone(), batch)),
1710 (dtype, matrix_shape(k, n, batch)),
1711 ])
1712}
1713
1714fn svd_full_meta(
1715 dtype: DType,
1716 shape: &[SymDim],
1717) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1718 let (m, n, batch) = matrix_meta_parts("tenferro-linalg.svd_full", shape)?;
1719 let k = m.clone().min(n.clone());
1720 Ok(vec![
1721 (dtype, matrix_shape(m.clone(), m, batch)),
1722 (singular_values_dtype(dtype), vector_shape(k, batch)),
1723 (dtype, matrix_shape(n.clone(), n, batch)),
1724 ])
1725}
1726
1727fn svd_values_meta(
1728 dtype: DType,
1729 shape: &[SymDim],
1730) -> tenferro_tensor::Result<(DType, Vec<SymDim>)> {
1731 let (m, n, batch) = matrix_meta_parts("tenferro-linalg.svd_values", shape)?;
1732 let k = m.min(n);
1733 Ok((singular_values_dtype(dtype), vector_shape(k, batch)))
1734}
1735
1736fn qr_meta(dtype: DType, shape: &[SymDim]) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1737 let (m, n, batch) = matrix_meta_parts("tenferro-linalg.qr", shape)?;
1738 let k = m.clone().min(n.clone());
1739 Ok(vec![
1740 (dtype, matrix_shape(m, k.clone(), batch)),
1741 (dtype, matrix_shape(k, n, batch)),
1742 ])
1743}
1744
1745fn rank_revealing_qr_meta(
1746 dtype: DType,
1747 shape: &[SymDim],
1748) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1749 let (m, n, batch) = matrix_meta_parts("tenferro-linalg.rank_revealing_qr", shape)?;
1750 let k = m.clone().min(n.clone());
1751 Ok(vec![
1752 (dtype, matrix_shape(m, k.clone(), batch)),
1753 (dtype, matrix_shape(k, n.clone(), batch)),
1754 (DType::I64, vector_shape(n, batch)),
1755 (DType::I64, batch.to_vec()),
1756 ])
1757}
1758
1759fn require_householder_rank2(op: &'static str, shape: &[SymDim]) -> tenferro_tensor::Result<()> {
1760 if shape.len() != 2 {
1761 return Err(Error::rank_mismatch(op, 2, shape.len()));
1762 }
1763 Ok(())
1764}
1765
1766fn householder_qr_factor_meta(
1767 dtype: DType,
1768 shape: &[SymDim],
1769) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1770 require_householder_rank2("tenferro-linalg.householder_qr", shape)?;
1771 let k = shape[0].clone().min(shape[1].clone());
1772 Ok(vec![(dtype, shape.to_vec()), (dtype, vec![k])])
1773}
1774
1775fn householder_qr_from_factors_meta(
1776 dtypes: &[DType],
1777 shapes: &[&[SymDim]],
1778) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1779 const OP: &str = "tenferro-linalg.householder_qr_from_factors";
1780 require_householder_rank2(OP, shapes[0])?;
1781 require_householder_rank2(OP, shapes[1])?;
1782 if dtypes[0] != dtypes[1] {
1783 return Err(Error::dtype_mismatch(OP, dtypes[0], dtypes[1]));
1784 }
1785 require_static_extent_equal(OP, "q.cols/r.rows", &shapes[0][1], &shapes[1][0])?;
1786 if let (Some(q_cols), Some(q_rows), Some(r_cols)) = (
1787 shapes[0][1].constant_value(),
1788 shapes[0][0].constant_value(),
1789 shapes[1][1].constant_value(),
1790 ) && q_cols > q_rows.min(r_cols)
1791 {
1792 return Err(Error::invalid_argument(
1793 OP,
1794 "shape",
1795 "Q column count exceeds min(Q rows, R columns)",
1796 ));
1797 }
1798 let m = shapes[0][0].clone();
1799 let n = shapes[1][1].clone();
1800 let k = m.clone().min(n.clone());
1801 Ok(vec![(dtypes[0], vec![m, n]), (dtypes[0], vec![k])])
1802}
1803
1804fn householder_qr_append_meta(
1805 dtypes: &[DType],
1806 shapes: &[&[SymDim]],
1807) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1808 const OP: &str = "tenferro-linalg.householder_qr_append";
1809 require_householder_state_meta(OP, dtypes, shapes[0], shapes[1])?;
1810 require_householder_rank2(OP, shapes[2])?;
1811 if dtypes[0] != dtypes[2] {
1812 return Err(Error::dtype_mismatch(OP, dtypes[0], dtypes[2]));
1813 }
1814 require_static_extent_equal(OP, "rows", &shapes[0][0], &shapes[2][0])?;
1815 let m = shapes[0][0].clone();
1816 let width = shapes[0][1].clone() + shapes[2][1].clone();
1817 let k = m.clone().min(width.clone());
1818 Ok(vec![(dtypes[0], vec![m, width]), (dtypes[0], vec![k])])
1819}
1820
1821fn require_static_extent_equal(
1822 op: &'static str,
1823 field: &'static str,
1824 lhs: &SymDim,
1825 rhs: &SymDim,
1826) -> tenferro_tensor::Result<()> {
1827 if let (Some(lhs), Some(rhs)) = (lhs.constant_value(), rhs.constant_value())
1828 && lhs != rhs
1829 {
1830 return Err(Error::invalid_argument(
1831 op,
1832 field,
1833 format!("expected equal extents, got {lhs} and {rhs}"),
1834 ));
1835 }
1836 Ok(())
1837}
1838
1839fn require_householder_state_meta(
1840 op: &'static str,
1841 dtypes: &[DType],
1842 packed: &[SymDim],
1843 coeff: &[SymDim],
1844) -> tenferro_tensor::Result<()> {
1845 require_householder_rank2(op, packed)?;
1846 if coeff.len() != 1 {
1847 return Err(Error::rank_mismatch(op, 1, coeff.len()));
1848 }
1849 if dtypes[0] != dtypes[1] {
1850 return Err(Error::dtype_mismatch(op, dtypes[0], dtypes[1]));
1851 }
1852 let expected = packed[0].clone().min(packed[1].clone());
1853 require_static_extent_equal(op, "coeff", &coeff[0], &expected)
1854}
1855
1856fn householder_qr_r_meta(
1857 dtypes: &[DType],
1858 packed: &[SymDim],
1859 coeff: &[SymDim],
1860) -> tenferro_tensor::Result<(DType, Vec<SymDim>)> {
1861 require_householder_state_meta("tenferro-linalg.householder_qr_r", dtypes, packed, coeff)?;
1862 Ok((dtypes[0], vec![coeff[0].clone(), packed[1].clone()]))
1863}
1864
1865fn householder_qr_q_columns_meta(
1866 dtypes: &[DType],
1867 packed: &[SymDim],
1868 coeff: &[SymDim],
1869 start: usize,
1870 end: usize,
1871) -> tenferro_tensor::Result<(DType, Vec<SymDim>)> {
1872 require_householder_state_meta(
1873 "tenferro-linalg.householder_qr_q_columns",
1874 dtypes,
1875 packed,
1876 coeff,
1877 )?;
1878 if start > end {
1879 return Err(Error::invalid_argument(
1880 "tenferro-linalg.householder_qr_q_columns",
1881 "range",
1882 format!("invalid Q-column range {start}..{end}"),
1883 ));
1884 }
1885 if packed[0].constant_value().is_some_and(|rows| end > rows) {
1888 return Err(Error::invalid_argument(
1889 "tenferro-linalg.householder_qr_q_columns",
1890 "range",
1891 format!("Q-column range {start}..{end} exceeds full-Q width"),
1892 ));
1893 }
1894 Ok((
1895 dtypes[0],
1896 vec![packed[0].clone(), SymDim::from(end - start)],
1897 ))
1898}
1899
1900fn householder_qr_thin_q_meta(
1901 dtypes: &[DType],
1902 packed: &[SymDim],
1903 coeff: &[SymDim],
1904) -> tenferro_tensor::Result<(DType, Vec<SymDim>)> {
1905 require_householder_state_meta(
1906 "tenferro-linalg.householder_qr_thin_q",
1907 dtypes,
1908 packed,
1909 coeff,
1910 )?;
1911 Ok((dtypes[0], vec![packed[0].clone(), coeff[0].clone()]))
1912}
1913
1914fn householder_qr_split_config(
1915 cotangent: &[usize],
1916 left: &[usize],
1917 right_shape: &[usize],
1918 take_right: bool,
1919) -> tenferro_tensor::Result<tenferro_tensor::SliceConfig> {
1920 const OP: &str = "tenferro-linalg.householder_qr_split_tangent";
1921 for shape in [cotangent, left, right_shape] {
1922 if shape.len() != 2 {
1923 return Err(Error::rank_mismatch(OP, 2, shape.len()));
1924 }
1925 }
1926 let total_width = left[1]
1927 .checked_add(right_shape[1])
1928 .ok_or_else(|| Error::invalid_argument(OP, "shape", "column range overflow"))?;
1929 if cotangent[0] != left[0] || cotangent[0] != right_shape[0] || cotangent[1] != total_width {
1930 return Err(Error::invalid_argument(
1931 OP,
1932 "shape",
1933 "cotangent shape does not match appended factors",
1934 ));
1935 }
1936 let selected = if take_right { right_shape } else { left };
1937 let start = if take_right { left[1] } else { 0 };
1938 let end = start
1939 .checked_add(selected[1])
1940 .ok_or_else(|| Error::invalid_argument(OP, "shape", "column range overflow"))?;
1941 Ok(tenferro_tensor::SliceConfig {
1942 starts: vec![0, start],
1943 limits: vec![selected[0], end],
1944 strides: vec![1, 1],
1945 })
1946}
1947
1948fn householder_qr_append_tangent_meta(
1949 dtypes: &[DType],
1950 shapes: &[&[SymDim]],
1951) -> tenferro_tensor::Result<(DType, Vec<SymDim>)> {
1952 const OP: &str = "tenferro-linalg.householder_qr_append_tangent";
1953 for shape in shapes {
1954 require_householder_rank2(OP, shape)?;
1955 }
1956 if dtypes.iter().any(|dtype| *dtype != dtypes[0]) {
1957 return Err(Error::dtype_mismatch(OP, dtypes[0], dtypes[1]));
1958 }
1959 require_static_extent_equal(OP, "rows", &shapes[0][0], &shapes[1][0])?;
1960 require_static_extent_equal(OP, "left tangent", &shapes[0][0], &shapes[2][0])?;
1961 require_static_extent_equal(OP, "right tangent", &shapes[1][0], &shapes[3][0])?;
1962 require_static_extent_equal(OP, "anchor rows", &shapes[2][0], &shapes[3][0])?;
1963 Ok((
1964 dtypes[0],
1965 vec![
1966 shapes[2][0].clone(),
1967 shapes[2][1].clone() + shapes[3][1].clone(),
1968 ],
1969 ))
1970}
1971
1972fn householder_qr_split_tangent_meta(
1973 dtypes: &[DType],
1974 shapes: &[&[SymDim]],
1975 right: bool,
1976) -> tenferro_tensor::Result<(DType, Vec<SymDim>)> {
1977 const OP: &str = "tenferro-linalg.householder_qr_split_tangent";
1978 for shape in shapes {
1979 require_householder_rank2(OP, shape)?;
1980 }
1981 if dtypes.iter().any(|dtype| *dtype != dtypes[0]) {
1982 return Err(Error::dtype_mismatch(OP, dtypes[0], dtypes[1]));
1983 }
1984 require_static_extent_equal(OP, "left rows", &shapes[0][0], &shapes[1][0])?;
1985 require_static_extent_equal(OP, "right rows", &shapes[0][0], &shapes[2][0])?;
1986 let expected_width = shapes[1][1].clone() + shapes[2][1].clone();
1987 require_static_extent_equal(OP, "width", &shapes[0][1], &expected_width)?;
1988 let selected = if right { shapes[2] } else { shapes[1] };
1989 Ok((dtypes[0], selected.to_vec()))
1990}
1991
1992fn eigh_meta(dtype: DType, shape: &[SymDim]) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1993 let (n, _, batch) = matrix_meta_parts("tenferro-linalg.eigh", shape)?;
1994 Ok(vec![
1995 (singular_values_dtype(dtype), vector_shape(n.clone(), batch)),
1996 (dtype, matrix_shape(n.clone(), n, batch)),
1997 ])
1998}
1999
2000fn eigh_values_meta(
2001 dtype: DType,
2002 shape: &[SymDim],
2003) -> tenferro_tensor::Result<(DType, Vec<SymDim>)> {
2004 let (n, _, batch) = matrix_meta_parts("tenferro-linalg.eigh_values", shape)?;
2005 Ok((singular_values_dtype(dtype), vector_shape(n, batch)))
2006}
2007
2008fn eig_meta(
2009 input_dtype: DType,
2010 shape: &[SymDim],
2011) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
2012 let dtype = eig_output_dtype(input_dtype);
2013 let (n, _, batch) = matrix_meta_parts("tenferro-linalg.eig", shape)?;
2014 Ok(vec![
2015 (dtype, vector_shape(n.clone(), batch)),
2016 (dtype, matrix_shape(n.clone(), n, batch)),
2017 ])
2018}
2019
2020fn eig_values_meta(
2021 input_dtype: DType,
2022 shape: &[SymDim],
2023) -> tenferro_tensor::Result<(DType, Vec<SymDim>)> {
2024 let dtype = eig_output_dtype(input_dtype);
2025 let (n, _, batch) = matrix_meta_parts("tenferro-linalg.eig_values", shape)?;
2026 Ok((dtype, vector_shape(n, batch)))
2027}
2028
2029fn matrix_shape(rows: SymDim, cols: SymDim, batch: &[SymDim]) -> Vec<SymDim> {
2030 let mut shape = vec![rows, cols];
2031 shape.extend_from_slice(batch);
2032 shape
2033}
2034
2035fn vector_shape(len: SymDim, batch: &[SymDim]) -> Vec<SymDim> {
2036 let mut shape = vec![len];
2037 shape.extend_from_slice(batch);
2038 shape
2039}
2040
2041fn eig_output_dtype(dtype: DType) -> DType {
2042 match dtype {
2043 DType::F64 | DType::C64 => DType::C64,
2044 DType::F32 | DType::C32 => DType::C32,
2045 DType::I32 | DType::I64 | DType::Bool => DType::C64,
2046 DType::External(_) => unreachable!("linalg validates its input dtype first"),
2049 }
2050}
2051
2052fn singular_values_dtype(dtype: DType) -> DType {
2053 match dtype {
2054 DType::C64 => DType::F64,
2055 DType::C32 => DType::F32,
2056 DType::External(id) => {
2057 DType::External(id)
2060 }
2061 other => other,
2062 }
2063}
2064
2065fn promote_dtypes(dtypes: &[DType]) -> DType {
2066 dtypes
2067 .iter()
2068 .copied()
2069 .reduce(tenferro_tensor::validate::promote_dtype)
2070 .unwrap_or(DType::F64)
2071}
2072
2073fn hash_dtype(hasher: &mut dyn Hasher, dtype: DType) {
2074 let tag = match dtype {
2075 DType::F64 => 0,
2076 DType::F32 => 1,
2077 DType::I64 => 2,
2078 DType::C64 => 3,
2079 DType::C32 => 4,
2080 DType::I32 => 5,
2081 DType::Bool => 6,
2082 DType::External(id) => {
2083 let mut identity = std::collections::hash_map::DefaultHasher::new();
2089 Hash::hash(&id, &mut identity);
2090 hasher.write_u8(7);
2091 hasher.write_u64(identity.finish());
2092 return;
2093 }
2094 };
2095 hasher.write_u8(tag);
2096}
2097
2098fn hash_svd_gauge(hasher: &mut dyn Hasher, gauge: SvdGauge) {
2099 let tag = match gauge {
2100 SvdGauge::Raw => 0,
2101 SvdGauge::CanonicalPivot => 1,
2102 };
2103 hasher.write_u8(tag);
2104}
2105
2106fn hash_svd_driver(hasher: &mut dyn Hasher, driver: SvdDriver) {
2107 let tag = match driver {
2108 SvdDriver::Auto => 0,
2109 SvdDriver::Gesvdj => 1,
2110 SvdDriver::Gesvd => 2,
2111 SvdDriver::Xgesvdp => 3,
2112 };
2113 hasher.write_u8(tag);
2114}
2115
2116fn hash_eigh_driver(hasher: &mut dyn Hasher, driver: EighDriver) {
2117 let tag = match driver {
2118 EighDriver::Auto => 0,
2119 EighDriver::Syevd => 1,
2120 EighDriver::Syevj => 2,
2121 };
2122 hasher.write_u8(tag);
2123}
2124
2125fn hash_eigh_gauge(hasher: &mut dyn Hasher, gauge: EighGauge) {
2126 let tag = match gauge {
2127 EighGauge::Raw => 0,
2128 EighGauge::CanonicalPivot => 1,
2129 };
2130 hasher.write_u8(tag);
2131}
2132
2133fn hash_qr_gauge(hasher: &mut dyn Hasher, gauge: QrGauge) {
2134 let tag = match gauge {
2135 QrGauge::Raw => 0,
2136 QrGauge::PositiveDiagonal => 1,
2137 };
2138 hasher.write_u8(tag);
2139}