1use std::any::Any;
2use std::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::{ExtensionExecutionContext, ExtensionOp};
10use tenferro_tensor::{BackendSession, DType, Error, ErrorKind, Tensor, TensorBackend, TensorRead};
11
12#[cfg(feature = "cuda")]
13use tenferro_gpu::with_cuda_exec_session;
14
15use crate::backend::LinalgBackend;
16
17mod gauge;
18#[cfg(all(test, not(feature = "cuda")))]
19mod tests;
20
21pub(crate) use gauge::{apply_eigh_gauge, apply_qr_gauge};
22
23pub const LINALG_EXTENSION_FAMILY_ID: &str = "tenferro-linalg.linalg.v1";
24
25pub const DEFAULT_DECOMPOSITION_DERIVATIVE_EPS: f64 = 1e-12;
39
40#[derive(Clone, Copy, Debug, PartialEq, Eq)]
51pub enum SvdGauge {
52 Raw,
54 CanonicalPivot,
57}
58
59#[derive(Clone, Copy, Debug, PartialEq, Eq)]
70pub enum EighGauge {
71 Raw,
73 CanonicalPivot,
75}
76
77#[derive(Clone, Copy, Debug, PartialEq, Eq)]
88pub enum QrGauge {
89 Raw,
91 PositiveDiagonal,
93}
94
95#[derive(Clone, Copy, Debug, PartialEq)]
109pub struct SvdOptions {
110 pub gauge: SvdGauge,
112 pub derivative_eps: f64,
114}
115
116impl Default for SvdOptions {
117 fn default() -> Self {
118 Self {
119 gauge: SvdGauge::Raw,
120 derivative_eps: DEFAULT_DECOMPOSITION_DERIVATIVE_EPS,
121 }
122 }
123}
124
125impl SvdOptions {
126 pub fn gauge(mut self, gauge: SvdGauge) -> Self {
137 self.gauge = gauge;
138 self
139 }
140
141 pub fn derivative_eps(mut self, derivative_eps: f64) -> Self {
152 self.derivative_eps = derivative_eps;
153 self
154 }
155}
156
157#[derive(Clone, Copy, Debug, PartialEq)]
168pub struct EighOptions {
169 pub gauge: EighGauge,
171 pub derivative_eps: f64,
173}
174
175impl Default for EighOptions {
176 fn default() -> Self {
177 Self {
178 gauge: EighGauge::Raw,
179 derivative_eps: DEFAULT_DECOMPOSITION_DERIVATIVE_EPS,
180 }
181 }
182}
183
184impl EighOptions {
185 pub fn gauge(mut self, gauge: EighGauge) -> Self {
196 self.gauge = gauge;
197 self
198 }
199
200 pub fn derivative_eps(mut self, derivative_eps: f64) -> Self {
211 self.derivative_eps = derivative_eps;
212 self
213 }
214}
215
216#[derive(Clone, Copy, Debug, PartialEq, Eq)]
227pub struct QrOptions {
228 pub gauge: QrGauge,
230}
231
232impl Default for QrOptions {
233 fn default() -> Self {
234 Self {
235 gauge: QrGauge::Raw,
236 }
237 }
238}
239
240impl QrOptions {
241 pub fn gauge(mut self, gauge: QrGauge) -> Self {
252 self.gauge = gauge;
253 self
254 }
255}
256
257pub(crate) fn validate_derivative_eps(
258 op: &'static str,
259 derivative_eps: f64,
260) -> tenferro_tensor::Result<()> {
261 if derivative_eps.is_finite() && derivative_eps > 0.0 {
262 Ok(())
263 } else {
264 Err(Error::invalid_argument(
265 op,
266 "derivative_eps",
267 format!("must be positive and finite, got {derivative_eps}"),
268 ))
269 }
270}
271
272#[derive(Clone, Copy, Debug, PartialEq)]
273#[doc(hidden)]
274pub(crate) enum LinalgOp {
275 Cholesky,
276 Lu,
277 LuFactor,
278 LuSolvePrepared {
279 transpose_a: bool,
280 conjugate_a: bool,
281 },
282 SignDetFromLuFactor,
283 LogAbsDetFromLuFactor,
284 FullPivLu,
285 FullPivLuSolve {
286 transpose_a: bool,
287 },
288 Svd {
289 derivative_eps: f64,
290 gauge: SvdGauge,
291 },
292 SvdFull,
296 SvdVals {
297 derivative_eps: f64,
298 },
299 Qr {
300 gauge: QrGauge,
301 },
302 Eigh {
303 derivative_eps: f64,
304 gauge: EighGauge,
305 },
306 EighVals {
307 derivative_eps: f64,
308 },
309 Eig {
310 input_dtype: DType,
311 },
312 EigVals {
313 input_dtype: DType,
314 },
315 TriangularSolve {
316 left_side: bool,
317 lower: bool,
318 transpose_a: bool,
319 unit_diagonal: bool,
320 },
321}
322
323impl LinalgOp {
324 fn output_count(self) -> usize {
325 match self {
326 Self::Cholesky
327 | Self::EighVals { .. }
328 | Self::EigVals { .. }
329 | Self::FullPivLuSolve { .. }
330 | Self::LogAbsDetFromLuFactor
331 | Self::LuSolvePrepared { .. }
332 | Self::SignDetFromLuFactor
333 | Self::SvdVals { .. }
334 | Self::TriangularSolve { .. } => 1,
335 Self::Svd { .. } | Self::SvdFull => 3,
336 Self::Qr { .. } | Self::Eigh { .. } | Self::Eig { .. } => 2,
337 Self::LuFactor => 3,
338 Self::Lu => 4,
339 Self::FullPivLu => 5,
340 }
341 }
342
343 fn input_count(self) -> usize {
344 match self {
345 Self::FullPivLuSolve { .. } | Self::TriangularSolve { .. } => 2,
346 Self::LogAbsDetFromLuFactor => 2,
347 Self::SignDetFromLuFactor => 3,
348 Self::LuSolvePrepared { .. } => 4,
349 _ => 1,
350 }
351 }
352
353 fn tag(self) -> u8 {
354 match self {
355 Self::Cholesky => 0,
356 Self::Lu => 1,
357 Self::FullPivLu => 2,
358 Self::FullPivLuSolve { .. } => 3,
359 Self::Svd { .. } => 4,
360 Self::Qr { .. } => 5,
361 Self::Eigh { .. } => 6,
362 Self::Eig { .. } => 7,
363 Self::TriangularSolve { .. } => 9,
364 Self::LuFactor => 10,
365 Self::LuSolvePrepared { .. } => 11,
366 Self::SvdVals { .. } => 12,
367 Self::EighVals { .. } => 13,
368 Self::EigVals { .. } => 14,
369 Self::SvdFull => 15,
370 Self::LogAbsDetFromLuFactor => 16,
371 Self::SignDetFromLuFactor => 17,
372 }
373 }
374}
375
376#[derive(Clone, Debug, PartialEq)]
377#[doc(hidden)]
378pub(crate) struct LinalgExtensionOp {
379 op: LinalgOp,
380}
381
382impl LinalgExtensionOp {
383 pub(crate) fn new(op: LinalgOp) -> Self {
384 Self { op }
385 }
386
387 pub(crate) fn op(&self) -> LinalgOp {
388 self.op
389 }
390}
391
392impl ExtensionOp for LinalgExtensionOp {
393 fn family_id(&self) -> &'static str {
394 LINALG_EXTENSION_FAMILY_ID
395 }
396
397 fn payload_hash(&self, hasher: &mut dyn Hasher) {
398 hasher.write_u8(self.op.tag());
399 match self.op {
400 LinalgOp::Svd {
401 derivative_eps,
402 gauge,
403 } => {
404 hasher.write_u64(derivative_eps.to_bits());
405 hash_svd_gauge(hasher, gauge);
406 }
407 LinalgOp::SvdVals { derivative_eps } | LinalgOp::EighVals { derivative_eps } => {
408 hasher.write_u64(derivative_eps.to_bits());
409 }
410 LinalgOp::Qr { gauge } => {
411 hash_qr_gauge(hasher, gauge);
412 }
413 LinalgOp::Eigh {
414 derivative_eps,
415 gauge,
416 } => {
417 hasher.write_u64(derivative_eps.to_bits());
418 hash_eigh_gauge(hasher, gauge);
419 }
420 LinalgOp::Eig { input_dtype } | LinalgOp::EigVals { input_dtype } => {
421 hash_dtype(hasher, input_dtype);
422 }
423 LinalgOp::FullPivLuSolve { transpose_a } => {
424 hasher.write_u8(u8::from(transpose_a));
425 }
426 LinalgOp::LuSolvePrepared {
427 transpose_a,
428 conjugate_a,
429 } => {
430 hasher.write_u8(u8::from(transpose_a));
431 hasher.write_u8(u8::from(conjugate_a));
432 }
433 LinalgOp::TriangularSolve {
434 left_side,
435 lower,
436 transpose_a,
437 unit_diagonal,
438 } => {
439 hasher.write_u8(u8::from(left_side));
440 hasher.write_u8(u8::from(lower));
441 hasher.write_u8(u8::from(transpose_a));
442 hasher.write_u8(u8::from(unit_diagonal));
443 }
444 LinalgOp::Cholesky
445 | LinalgOp::Lu
446 | LinalgOp::LuFactor
447 | LinalgOp::LogAbsDetFromLuFactor
448 | LinalgOp::SignDetFromLuFactor
449 | LinalgOp::FullPivLu
450 | LinalgOp::SvdFull => {}
451 }
452 }
453
454 fn payload_eq(&self, other: &dyn ExtensionOp) -> bool {
455 other
456 .as_any()
457 .downcast_ref::<Self>()
458 .is_some_and(|that| self == that)
459 }
460
461 fn clone_arc(&self) -> Arc<dyn ExtensionOp> {
462 Arc::new(self.clone())
463 }
464
465 fn as_any(&self) -> &dyn Any {
466 self
467 }
468
469 fn input_count(&self) -> usize {
470 self.op.input_count()
471 }
472
473 fn output_count(&self) -> usize {
474 self.op.output_count()
475 }
476
477 fn semantic_effects(&self) -> tenferro_ops::ext_op::ExtensionEffectDeclaration<'_> {
478 tenferro_ops::ext_op::ExtensionEffectDeclaration::Declared(&[])
479 }
480
481 fn semantic_aliases(&self) -> tenferro_ops::ext_op::ExtensionAliasDeclaration<'_> {
482 tenferro_ops::ext_op::ExtensionAliasDeclaration::AllFresh
483 }
484
485 fn prune_outputs(&self, live_outputs: &[bool]) -> Option<Arc<dyn ExtensionOp>> {
486 match self.op {
487 LinalgOp::Svd { derivative_eps, .. } if live_outputs == [false, true, false] => {
488 Some(Arc::new(Self::new(LinalgOp::SvdVals { derivative_eps })))
489 }
490 LinalgOp::Eigh { derivative_eps, .. } if live_outputs == [true, false] => {
491 Some(Arc::new(Self::new(LinalgOp::EighVals { derivative_eps })))
492 }
493 LinalgOp::Eig { input_dtype } if live_outputs == [true, false] => {
494 Some(Arc::new(Self::new(LinalgOp::EigVals { input_dtype })))
495 }
496 _ => None,
497 }
498 }
499
500 fn infer_output_meta(
501 &self,
502 ctx: &mut tenferro_ops::ExtensionShapeContext<'_>,
503 ) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
504 let input_dtypes = (0..self.input_count())
505 .map(|input| ctx.input_dtype(input))
506 .collect::<Result<Vec<_>, _>>()?;
507 let input_shapes = (0..self.input_count())
508 .map(|input| ctx.input_shape(input))
509 .collect::<Result<Vec<_>, _>>()?;
510 let metas = match self.op {
511 LinalgOp::Cholesky => {
512 require_matrix_meta("tenferro-linalg.cholesky", input_shapes[0])?;
513 vec![(promote_dtypes(&input_dtypes), input_shapes[0].to_vec())]
514 }
515 LinalgOp::FullPivLuSolve { .. } => {
516 require_matrix_meta("tenferro-linalg.full_piv_lu_solve", input_shapes[0])?;
517 require_matrix_meta("tenferro-linalg.full_piv_lu_solve", input_shapes[1])?;
518 vec![(promote_dtypes(&input_dtypes), input_shapes[1].to_vec())]
519 }
520 LinalgOp::TriangularSolve { .. } => {
521 require_matrix_meta("tenferro-linalg.triangular_solve", input_shapes[0])?;
522 require_matrix_meta("tenferro-linalg.triangular_solve", input_shapes[1])?;
523 vec![(promote_dtypes(&input_dtypes), input_shapes[1].to_vec())]
524 }
525 LinalgOp::LuSolvePrepared { .. } => {
526 require_matrix_meta("tenferro-linalg.lu_solve_prepared_lu", input_shapes[0])?;
527 require_matrix_meta("tenferro-linalg.lu_solve_prepared_rhs", input_shapes[3])?;
528 vec![(
529 promote_dtypes(&[input_dtypes[0], input_dtypes[3]]),
530 input_shapes[3].to_vec(),
531 )]
532 }
533 LinalgOp::Lu => lu_meta(input_dtypes[0], input_shapes[0])?,
534 LinalgOp::LuFactor => lu_factor_meta(input_dtypes[0], input_shapes[0])?,
535 LinalgOp::SignDetFromLuFactor => {
536 vec![signdet_from_lu_factor_meta(
537 input_dtypes[0],
538 input_shapes[0],
539 input_shapes[1],
540 input_shapes[2],
541 )?]
542 }
543 LinalgOp::LogAbsDetFromLuFactor => {
544 vec![logabsdet_from_lu_factor_meta(
545 input_dtypes[0],
546 input_shapes[0],
547 input_shapes[1],
548 )?]
549 }
550 LinalgOp::FullPivLu => full_piv_lu_meta(input_dtypes[0], input_shapes[0])?,
551 LinalgOp::Svd { .. } => svd_meta(input_dtypes[0], input_shapes[0])?,
552 LinalgOp::SvdFull => svd_full_meta(input_dtypes[0], input_shapes[0])?,
553 LinalgOp::SvdVals { .. } => {
554 vec![svd_values_meta(input_dtypes[0], input_shapes[0])?]
555 }
556 LinalgOp::Qr { .. } => qr_meta(input_dtypes[0], input_shapes[0])?,
557 LinalgOp::Eigh { .. } => eigh_meta(input_dtypes[0], input_shapes[0])?,
558 LinalgOp::EighVals { .. } => vec![eigh_values_meta(input_dtypes[0], input_shapes[0])?],
559 LinalgOp::Eig { input_dtype } => eig_meta(input_dtype, input_shapes[0])?,
560 LinalgOp::EigVals { input_dtype } => {
561 vec![eig_values_meta(input_dtype, input_shapes[0])?]
562 }
563 };
564 Ok(metas)
565 }
566}
567
568pub(crate) fn execute_linalg_extension_reads<B: BackendSession + ?Sized>(
569 op: &LinalgExtensionOp,
570 inputs: &[TensorRead<'_>],
571 ctx: &mut ExtensionExecutionContext<'_, B>,
572) -> tenferro_tensor::Result<Vec<Tensor>> {
573 execute_linalg_extension_reads_on_session(op, inputs, ctx.backend_mut())
574}
575
576pub(crate) fn execute_linalg_extension_reads_owner<B: TensorBackend>(
577 op: &LinalgExtensionOp,
578 inputs: &[TensorRead<'_>],
579 ctx: &mut ExtensionExecutionContext<'_, B>,
580) -> tenferro_tensor::Result<Vec<Tensor>> {
581 let (backend, caches) = ctx.parts_mut();
582 backend.with_backend_session(|session| {
583 let mut session_ctx = ExtensionExecutionContext::new(session, caches);
584 execute_linalg_extension_reads(op, inputs, &mut session_ctx)
585 })
586}
587
588fn execute_linalg_extension_reads_on_session<B: BackendSession + ?Sized>(
589 op: &LinalgExtensionOp,
590 inputs: &[TensorRead<'_>],
591 session: &mut B,
592) -> tenferro_tensor::Result<Vec<Tensor>> {
593 if let Some(result) = with_cpu_exec_session(session, |session| {
594 execute_linalg_extension_reads_in_session(op, inputs, session)
595 }) {
596 return result;
597 }
598 #[cfg(feature = "cuda")]
599 if let Some(result) = with_cuda_exec_session(session, |session| {
600 execute_linalg_extension_reads_in_session(op, inputs, session)
601 }) {
602 return result;
603 }
604 Err(Error::unsupported(
605 "linalg_extension",
606 "selected backend session does not expose a linalg execution capability",
607 ))
608}
609
610fn execute_linalg_extension_reads_in_session<S: LinalgBackend>(
611 op: &LinalgExtensionOp,
612 inputs: &[TensorRead<'_>],
613 session: &mut S,
614) -> tenferro_tensor::Result<Vec<Tensor>> {
615 if op.op() == LinalgOp::Cholesky {
616 return Ok(vec![session.cholesky_read(inputs[0].clone())?]);
617 }
618 if let LinalgOp::TriangularSolve {
619 left_side,
620 lower,
621 transpose_a,
622 unit_diagonal,
623 } = op.op()
624 {
625 match session.triangular_solve_read(
626 inputs[0].clone(),
627 inputs[1].clone(),
628 left_side,
629 lower,
630 transpose_a,
631 unit_diagonal,
632 ) {
633 Ok(output) => return Ok(vec![output]),
634 Err(error) if error.kind() == ErrorKind::Unsupported => {}
635 Err(error) => return Err(error),
636 }
637 }
638
639 let materialized_inputs = inputs
642 .iter()
643 .cloned()
644 .map(|input| session.to_contiguous_read(input))
645 .collect::<tenferro_tensor::Result<Vec<_>>>()?;
646 let input_refs: Vec<&Tensor> = materialized_inputs.iter().collect();
647 execute_linalg(op.op(), &input_refs, session)
648}
649
650fn linalg_session_supported<B: BackendSession + 'static>(op: &LinalgExtensionOp) -> bool {
651 matches!(op.op(), LinalgOp::LuSolvePrepared { .. })
652 && std::any::TypeId::of::<B>() == std::any::TypeId::of::<tenferro_cpu::CpuBackend>()
653}
654
655fn execute_linalg_extension_in_session(
656 op: &LinalgExtensionOp,
657 session: &mut dyn BackendSession,
658 _extension_caches: &mut tenferro_runtime::ExtensionCacheStore,
659 inputs: &[TensorRead<'_>],
660) -> tenferro_tensor::Result<Vec<Tensor>> {
661 if !matches!(op.op(), LinalgOp::LuSolvePrepared { .. }) {
662 return Err(Error::unsupported(
663 "linalg_extension",
664 "linalg operation has no scheduler-session implementation",
665 ));
666 };
667 if inputs.len() != 4 {
668 return Err(Error::invalid_argument(
669 "linalg_extension",
670 "inputs",
671 format!(
672 "LuSolvePrepared session execution expected 4 inputs, got {}",
673 inputs.len()
674 ),
675 ));
676 }
677 execute_linalg_extension_reads_on_session(op, inputs, session)
678}
679
680define_extension_runtime! {
681 runtime = LinalgRuntime,
682 family_id = LINALG_EXTENSION_FAMILY_ID,
683 op_type = LinalgExtensionOp,
684 execute = execute_linalg_extension_reads_owner,
685 execute_reads = execute_linalg_extension_reads_owner,
686 execute_in_session = execute_linalg_extension_in_session,
687 session_supported = linalg_session_supported,
688 backend_bound = TensorBackend,
689}
690
691fn execute_linalg<B: LinalgBackend>(
692 op: LinalgOp,
693 inputs: &[&Tensor],
694 backend: &mut B,
695) -> tenferro_tensor::Result<Vec<Tensor>> {
696 match op {
697 LinalgOp::Cholesky => Ok(vec![backend.cholesky(inputs[0])?]),
698 LinalgOp::Lu => backend.lu(inputs[0]),
699 LinalgOp::LuFactor => backend.lu_factor(inputs[0]),
700 LinalgOp::SignDetFromLuFactor => Ok(vec![signdet_from_lu_factor(
701 inputs[0].dtype(),
702 inputs[1],
703 inputs[2],
704 backend,
705 )?]),
706 LinalgOp::LogAbsDetFromLuFactor => Ok(vec![logabsdet_from_lu_factor(inputs[1], backend)?]),
707 LinalgOp::LuSolvePrepared {
708 transpose_a,
709 conjugate_a,
710 } => Ok(vec![backend.lu_solve_prepared(
711 inputs[0],
712 inputs[1],
713 inputs[2],
714 inputs[3],
715 transpose_a,
716 conjugate_a,
717 )?]),
718 LinalgOp::FullPivLu => backend.full_piv_lu(inputs[0]),
719 LinalgOp::FullPivLuSolve { transpose_a } => Ok(vec![backend.full_piv_lu_solve(
720 inputs[0],
721 inputs[1],
722 transpose_a,
723 )?]),
724 LinalgOp::Svd {
725 derivative_eps,
726 gauge,
727 } => backend.svd_with_options(
728 inputs[0],
729 SvdOptions {
730 derivative_eps,
731 gauge,
732 },
733 ),
734 LinalgOp::SvdFull => backend.svd_full(inputs[0]),
735 LinalgOp::SvdVals { .. } => Ok(vec![backend.svd_values(inputs[0])?]),
736 LinalgOp::Qr { gauge } => backend.qr_with_options(inputs[0], QrOptions { gauge }),
737 LinalgOp::Eigh {
738 derivative_eps,
739 gauge,
740 } => backend.eigh_with_options(
741 inputs[0],
742 EighOptions {
743 derivative_eps,
744 gauge,
745 },
746 ),
747 LinalgOp::EighVals { .. } => Ok(vec![backend.eigh_values(inputs[0])?]),
748 LinalgOp::Eig { .. } => backend.eig(inputs[0]),
749 LinalgOp::EigVals { .. } => Ok(vec![backend.eig_values(inputs[0])?]),
750 LinalgOp::TriangularSolve {
751 left_side,
752 lower,
753 transpose_a,
754 unit_diagonal,
755 } => Ok(vec![backend.triangular_solve(
756 inputs[0],
757 inputs[1],
758 left_side,
759 lower,
760 transpose_a,
761 unit_diagonal,
762 )?]),
763 }
764}
765
766fn signdet_from_lu_factor<B: LinalgBackend + ?Sized>(
767 input_dtype: DType,
768 packed_lu: &Tensor,
769 parity: &Tensor,
770 backend: &mut B,
771) -> tenferro_tensor::Result<Tensor> {
772 let diag = backend.extract_diagonal(packed_lu, 0, 1)?;
773 let det_u = backend.reduce_prod_read(TensorRead::from_tensor(&diag), &[0])?;
774 let det = backend.mul_read(
775 TensorRead::from_tensor(parity),
776 TensorRead::from_tensor(&det_u),
777 )?;
778 if matches!(input_dtype, DType::C32 | DType::C64) {
779 let abs = backend.abs_read(TensorRead::from_tensor(&det))?;
780 let abs = backend.convert(&abs, input_dtype)?;
781 backend.div_read(TensorRead::from_tensor(&det), TensorRead::from_tensor(&abs))
782 } else {
783 backend.sign_read(TensorRead::from_tensor(&det))
784 }
785}
786
787fn logabsdet_from_lu_factor<B: LinalgBackend + ?Sized>(
788 packed_lu: &Tensor,
789 backend: &mut B,
790) -> tenferro_tensor::Result<Tensor> {
791 let diag = backend.extract_diagonal(packed_lu, 0, 1)?;
792 let abs = backend.abs_read(TensorRead::from_tensor(&diag))?;
793 let log = backend.log_read(TensorRead::from_tensor(&abs))?;
794 backend.reduce_sum_read(TensorRead::from_tensor(&log), &[0])
795}
796
797pub(crate) fn apply_svd_gauge(
798 gauge: SvdGauge,
799 outputs: &mut [Tensor],
800) -> tenferro_tensor::Result<()> {
801 match gauge {
802 SvdGauge::Raw => Ok(()),
803 SvdGauge::CanonicalPivot => apply_canonical_pivot_svd_gauge(outputs),
804 }
805}
806
807fn apply_canonical_pivot_svd_gauge(outputs: &mut [Tensor]) -> tenferro_tensor::Result<()> {
808 if outputs.len() != 3 {
809 return Err(Error::invalid_argument(
810 "tenferro-linalg.svd",
811 "outputs",
812 format!(
813 "canonical SVD gauge expected three outputs, got {}",
814 outputs.len()
815 ),
816 ));
817 }
818
819 let (u_slice, rest) = outputs.split_at_mut(1);
820 let (singular_slice, vt_slice) = rest.split_at_mut(1);
821 let u = &mut u_slice[0];
822 let singular_values = &singular_slice[0];
823 let vt = &mut vt_slice[0];
824 let u_shape = u.shape().to_vec();
825 let s_shape = singular_values.shape().to_vec();
826 let vt_shape = vt.shape().to_vec();
827 if u_shape.len() < 2 || vt_shape.len() < 2 || s_shape.is_empty() {
828 return Err(Error::invalid_argument(
829 "tenferro-linalg.svd",
830 "outputs",
831 format!(
832 "canonical SVD gauge expected U rank >= 2, S rank >= 1, VT rank >= 2; got U={u_shape:?}, S={s_shape:?}, VT={vt_shape:?}"
833 ),
834 ));
835 }
836
837 let m = u_shape[0];
838 let k = u_shape[1];
839 let n = vt_shape[1];
840 if s_shape[0] != k
841 || vt_shape[0] != k
842 || u_shape[2..] != vt_shape[2..]
843 || s_shape[1..] != u_shape[2..]
844 {
845 return Err(Error::invalid_argument(
846 "tenferro-linalg.svd",
847 "outputs",
848 format!(
849 "canonical SVD gauge expected compatible compact SVD shapes, got U={u_shape:?}, S={s_shape:?}, VT={vt_shape:?}"
850 ),
851 ));
852 }
853 let layout = canonical_svd_gauge_layout(m, k, n, &u_shape[2..])?;
854
855 match (u, vt) {
856 (Tensor::F64(u), Tensor::F64(vt)) => {
857 canonicalize_svd_gauge_f64(u.host_data_mut()?, vt.host_data_mut()?, layout)
858 }
859 (Tensor::F32(u), Tensor::F32(vt)) => {
860 canonicalize_svd_gauge_f32(u.host_data_mut()?, vt.host_data_mut()?, layout)
861 }
862 (Tensor::C64(u), Tensor::C64(vt)) => {
863 canonicalize_svd_gauge_c64(u.host_data_mut()?, vt.host_data_mut()?, layout)
864 }
865 (Tensor::C32(u), Tensor::C32(vt)) => {
866 canonicalize_svd_gauge_c32(u.host_data_mut()?, vt.host_data_mut()?, layout)
867 }
868 (u, vt) => Err(Error::dtype_mismatch(
869 "tenferro-linalg.svd",
870 u.dtype(),
871 vt.dtype(),
872 )),
873 }
874}
875
876#[derive(Clone, Copy, Debug, PartialEq, Eq)]
877struct CanonicalSvdGaugeLayout {
878 m: usize,
879 k: usize,
880 batch_count: usize,
881 u_batch_len: usize,
882 vt_batch_len: usize,
883 u_len: usize,
884 vt_len: usize,
885}
886
887impl CanonicalSvdGaugeLayout {
888 fn validate_storage(self, u_len: usize, vt_len: usize) -> tenferro_tensor::Result<()> {
889 if u_len != self.u_len {
890 return Err(Error::invalid_argument(
891 "tenferro-linalg.svd",
892 "U storage",
893 format!(
894 "canonical SVD gauge expected U storage length {}, got {u_len}",
895 self.u_len
896 ),
897 ));
898 }
899 if vt_len != self.vt_len {
900 return Err(Error::invalid_argument(
901 "tenferro-linalg.svd",
902 "VT storage",
903 format!(
904 "canonical SVD gauge expected VT storage length {}, got {vt_len}",
905 self.vt_len
906 ),
907 ));
908 }
909 Ok(())
910 }
911}
912
913fn canonical_svd_gauge_layout(
914 m: usize,
915 k: usize,
916 n: usize,
917 batch_shape: &[usize],
918) -> tenferro_tensor::Result<CanonicalSvdGaugeLayout> {
919 let batch_count = tenferro_tensor::validate::checked_shape_product(
920 "tenferro-linalg.svd",
921 "canonical SVD batch",
922 batch_shape,
923 )?;
924 let u_batch_len = tenferro_tensor::validate::checked_shape_product(
925 "tenferro-linalg.svd",
926 "canonical SVD U batch",
927 &[m, k],
928 )?;
929 let vt_batch_len = tenferro_tensor::validate::checked_shape_product(
930 "tenferro-linalg.svd",
931 "canonical SVD VT batch",
932 &[k, n],
933 )?;
934 let u_len = tenferro_tensor::validate::checked_shape_product(
935 "tenferro-linalg.svd",
936 "canonical SVD U storage",
937 &[u_batch_len, batch_count],
938 )?;
939 let vt_len = tenferro_tensor::validate::checked_shape_product(
940 "tenferro-linalg.svd",
941 "canonical SVD VT storage",
942 &[vt_batch_len, batch_count],
943 )?;
944 Ok(CanonicalSvdGaugeLayout {
945 m,
946 k,
947 batch_count,
948 u_batch_len,
949 vt_batch_len,
950 u_len,
951 vt_len,
952 })
953}
954
955fn canonicalize_svd_gauge_f64(
956 u: &mut [f64],
957 vt: &mut [f64],
958 layout: CanonicalSvdGaugeLayout,
959) -> tenferro_tensor::Result<()> {
960 layout.validate_storage(u.len(), vt.len())?;
961 if layout.batch_count == 0 || layout.u_batch_len == 0 || layout.vt_batch_len == 0 {
962 return Ok(());
963 }
964 for (u_batch, vt_batch) in u
965 .chunks_exact_mut(layout.u_batch_len)
966 .zip(vt.chunks_exact_mut(layout.vt_batch_len))
967 {
968 for (col, u_column) in u_batch.chunks_exact_mut(layout.m).enumerate() {
969 let pivot = max_abs_pivot_f64(u_column);
970 let pivot_value = u_column[pivot];
971 if pivot_value < 0.0 {
972 for value in u_column {
973 *value = -*value;
974 }
975 for vt_column in vt_batch.chunks_exact_mut(layout.k) {
976 vt_column[col] = -vt_column[col];
977 }
978 }
979 }
980 }
981 Ok(())
982}
983
984fn canonicalize_svd_gauge_f32(
985 u: &mut [f32],
986 vt: &mut [f32],
987 layout: CanonicalSvdGaugeLayout,
988) -> tenferro_tensor::Result<()> {
989 layout.validate_storage(u.len(), vt.len())?;
990 if layout.batch_count == 0 || layout.u_batch_len == 0 || layout.vt_batch_len == 0 {
991 return Ok(());
992 }
993 for (u_batch, vt_batch) in u
994 .chunks_exact_mut(layout.u_batch_len)
995 .zip(vt.chunks_exact_mut(layout.vt_batch_len))
996 {
997 for (col, u_column) in u_batch.chunks_exact_mut(layout.m).enumerate() {
998 let pivot = max_abs_pivot_f32(u_column);
999 let pivot_value = u_column[pivot];
1000 if pivot_value < 0.0 {
1001 for value in u_column {
1002 *value = -*value;
1003 }
1004 for vt_column in vt_batch.chunks_exact_mut(layout.k) {
1005 vt_column[col] = -vt_column[col];
1006 }
1007 }
1008 }
1009 }
1010 Ok(())
1011}
1012
1013fn canonicalize_svd_gauge_c64(
1014 u: &mut [Complex64],
1015 vt: &mut [Complex64],
1016 layout: CanonicalSvdGaugeLayout,
1017) -> tenferro_tensor::Result<()> {
1018 layout.validate_storage(u.len(), vt.len())?;
1019 if layout.batch_count == 0 || layout.u_batch_len == 0 || layout.vt_batch_len == 0 {
1020 return Ok(());
1021 }
1022 for (u_batch, vt_batch) in u
1023 .chunks_exact_mut(layout.u_batch_len)
1024 .zip(vt.chunks_exact_mut(layout.vt_batch_len))
1025 {
1026 for (col, u_column) in u_batch.chunks_exact_mut(layout.m).enumerate() {
1027 let pivot = max_abs_pivot_c64(u_column);
1028 let pivot_value = u_column[pivot];
1029 let pivot_norm = pivot_value.norm();
1030 if pivot_norm == 0.0 {
1031 continue;
1032 }
1033 let phase = pivot_value.conj() / pivot_norm;
1034 let vt_phase = phase.conj();
1035 for value in u_column {
1036 *value *= phase;
1037 }
1038 for vt_column in vt_batch.chunks_exact_mut(layout.k) {
1039 vt_column[col] *= vt_phase;
1040 }
1041 }
1042 }
1043 Ok(())
1044}
1045
1046fn canonicalize_svd_gauge_c32(
1047 u: &mut [Complex32],
1048 vt: &mut [Complex32],
1049 layout: CanonicalSvdGaugeLayout,
1050) -> tenferro_tensor::Result<()> {
1051 layout.validate_storage(u.len(), vt.len())?;
1052 if layout.batch_count == 0 || layout.u_batch_len == 0 || layout.vt_batch_len == 0 {
1053 return Ok(());
1054 }
1055 for (u_batch, vt_batch) in u
1056 .chunks_exact_mut(layout.u_batch_len)
1057 .zip(vt.chunks_exact_mut(layout.vt_batch_len))
1058 {
1059 for (col, u_column) in u_batch.chunks_exact_mut(layout.m).enumerate() {
1060 let pivot = max_abs_pivot_c32(u_column);
1061 let pivot_value = u_column[pivot];
1062 let pivot_norm = pivot_value.norm();
1063 if pivot_norm == 0.0 {
1064 continue;
1065 }
1066 let phase = pivot_value.conj() / pivot_norm;
1067 let vt_phase = phase.conj();
1068 for value in u_column {
1069 *value *= phase;
1070 }
1071 for vt_column in vt_batch.chunks_exact_mut(layout.k) {
1072 vt_column[col] *= vt_phase;
1073 }
1074 }
1075 }
1076 Ok(())
1077}
1078
1079fn max_abs_pivot_f64(u_column: &[f64]) -> usize {
1080 let mut pivot = 0;
1081 let mut pivot_abs = u_column[0].abs();
1082 for (row, value) in u_column.iter().enumerate().skip(1) {
1083 let candidate_abs = value.abs();
1084 if candidate_abs > pivot_abs {
1085 pivot = row;
1086 pivot_abs = candidate_abs;
1087 }
1088 }
1089 pivot
1090}
1091
1092fn max_abs_pivot_f32(u_column: &[f32]) -> usize {
1093 let mut pivot = 0;
1094 let mut pivot_abs = u_column[0].abs();
1095 for (row, value) in u_column.iter().enumerate().skip(1) {
1096 let candidate_abs = value.abs();
1097 if candidate_abs > pivot_abs {
1098 pivot = row;
1099 pivot_abs = candidate_abs;
1100 }
1101 }
1102 pivot
1103}
1104
1105fn max_abs_pivot_c64(u_column: &[Complex64]) -> usize {
1106 let mut pivot = 0;
1107 let mut pivot_abs = u_column[0].norm_sqr();
1108 for (row, value) in u_column.iter().enumerate().skip(1) {
1109 let candidate_abs = value.norm_sqr();
1110 if candidate_abs > pivot_abs {
1111 pivot = row;
1112 pivot_abs = candidate_abs;
1113 }
1114 }
1115 pivot
1116}
1117
1118fn max_abs_pivot_c32(u_column: &[Complex32]) -> usize {
1119 let mut pivot = 0;
1120 let mut pivot_abs = u_column[0].norm_sqr();
1121 for (row, value) in u_column.iter().enumerate().skip(1) {
1122 let candidate_abs = value.norm_sqr();
1123 if candidate_abs > pivot_abs {
1124 pivot = row;
1125 pivot_abs = candidate_abs;
1126 }
1127 }
1128 pivot
1129}
1130
1131fn require_matrix_meta(op: &'static str, shape: &[SymDim]) -> tenferro_tensor::Result<()> {
1132 if shape.len() < 2 {
1133 return Err(Error::rank_mismatch(op, 2, shape.len()));
1134 }
1135 Ok(())
1136}
1137
1138fn matrix_meta_parts<'a>(
1139 op: &'static str,
1140 shape: &'a [SymDim],
1141) -> tenferro_tensor::Result<(SymDim, SymDim, &'a [SymDim])> {
1142 require_matrix_meta(op, shape)?;
1143 Ok((shape[0].clone(), shape[1].clone(), &shape[2..]))
1144}
1145
1146fn lu_meta(dtype: DType, shape: &[SymDim]) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1147 let (m, n, batch) = matrix_meta_parts("tenferro-linalg.lu", shape)?;
1148 let k = m.clone().min(n.clone());
1149 Ok(vec![
1150 (dtype, matrix_shape(m.clone(), m, batch)),
1151 (dtype, matrix_shape(shape[0].clone(), k.clone(), batch)),
1152 (dtype, matrix_shape(k, n, batch)),
1153 (dtype, batch.to_vec()),
1154 ])
1155}
1156
1157fn lu_factor_meta(
1158 dtype: DType,
1159 shape: &[SymDim],
1160) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1161 let (m, n, batch) = matrix_meta_parts("tenferro-linalg.lu_factor", shape)?;
1162 let k = m.min(n);
1163 Ok(vec![
1164 (dtype, shape.to_vec()),
1165 (DType::I32, vector_shape(k, batch)),
1166 (dtype, batch.to_vec()),
1167 ])
1168}
1169
1170fn signdet_from_lu_factor_meta(
1171 input_dtype: DType,
1172 input_shape: &[SymDim],
1173 packed_shape: &[SymDim],
1174 parity_shape: &[SymDim],
1175) -> tenferro_tensor::Result<(DType, Vec<SymDim>)> {
1176 let (_, _, batch) = matrix_meta_parts("tenferro-linalg.signdet_from_lu_factor", input_shape)?;
1177 require_matrix_meta(
1178 "tenferro-linalg.signdet_from_lu_factor_packed",
1179 packed_shape,
1180 )?;
1181 if parity_shape.len() != batch.len() {
1182 return Err(Error::rank_mismatch(
1183 "tenferro-linalg.signdet_from_lu_factor_parity",
1184 batch.len(),
1185 parity_shape.len(),
1186 ));
1187 }
1188 Ok((input_dtype, batch.to_vec()))
1189}
1190
1191fn logabsdet_from_lu_factor_meta(
1192 input_dtype: DType,
1193 input_shape: &[SymDim],
1194 packed_shape: &[SymDim],
1195) -> tenferro_tensor::Result<(DType, Vec<SymDim>)> {
1196 let (_, _, batch) = matrix_meta_parts("tenferro-linalg.logabsdet_from_lu_factor", input_shape)?;
1197 require_matrix_meta(
1198 "tenferro-linalg.logabsdet_from_lu_factor_packed",
1199 packed_shape,
1200 )?;
1201 Ok((singular_values_dtype(input_dtype), batch.to_vec()))
1202}
1203
1204fn full_piv_lu_meta(
1205 dtype: DType,
1206 shape: &[SymDim],
1207) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1208 let (n, _, batch) = matrix_meta_parts("tenferro-linalg.full_piv_lu", shape)?;
1209 Ok(vec![
1210 (dtype, matrix_shape(n.clone(), n.clone(), batch)),
1211 (dtype, matrix_shape(n.clone(), n.clone(), batch)),
1212 (dtype, matrix_shape(n.clone(), n.clone(), batch)),
1213 (dtype, matrix_shape(n.clone(), n, batch)),
1214 (singular_values_dtype(dtype), batch.to_vec()),
1215 ])
1216}
1217
1218fn svd_meta(dtype: DType, shape: &[SymDim]) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1219 let (m, n, batch) = matrix_meta_parts("tenferro-linalg.svd", shape)?;
1220 let k = m.clone().min(n.clone());
1221 Ok(vec![
1222 (dtype, matrix_shape(m, k.clone(), batch)),
1223 (singular_values_dtype(dtype), vector_shape(k.clone(), batch)),
1224 (dtype, matrix_shape(k, n, batch)),
1225 ])
1226}
1227
1228fn svd_full_meta(
1229 dtype: DType,
1230 shape: &[SymDim],
1231) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1232 let (m, n, batch) = matrix_meta_parts("tenferro-linalg.svd_full", shape)?;
1233 let k = m.clone().min(n.clone());
1234 Ok(vec![
1235 (dtype, matrix_shape(m.clone(), m, batch)),
1236 (singular_values_dtype(dtype), vector_shape(k, batch)),
1237 (dtype, matrix_shape(n.clone(), n, batch)),
1238 ])
1239}
1240
1241fn svd_values_meta(
1242 dtype: DType,
1243 shape: &[SymDim],
1244) -> tenferro_tensor::Result<(DType, Vec<SymDim>)> {
1245 let (m, n, batch) = matrix_meta_parts("tenferro-linalg.svd_values", shape)?;
1246 let k = m.min(n);
1247 Ok((singular_values_dtype(dtype), vector_shape(k, batch)))
1248}
1249
1250fn qr_meta(dtype: DType, shape: &[SymDim]) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1251 let (m, n, batch) = matrix_meta_parts("tenferro-linalg.qr", shape)?;
1252 let k = m.clone().min(n.clone());
1253 Ok(vec![
1254 (dtype, matrix_shape(m, k.clone(), batch)),
1255 (dtype, matrix_shape(k, n, batch)),
1256 ])
1257}
1258
1259fn eigh_meta(dtype: DType, shape: &[SymDim]) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1260 let (n, _, batch) = matrix_meta_parts("tenferro-linalg.eigh", shape)?;
1261 Ok(vec![
1262 (singular_values_dtype(dtype), vector_shape(n.clone(), batch)),
1263 (dtype, matrix_shape(n.clone(), n, batch)),
1264 ])
1265}
1266
1267fn eigh_values_meta(
1268 dtype: DType,
1269 shape: &[SymDim],
1270) -> tenferro_tensor::Result<(DType, Vec<SymDim>)> {
1271 let (n, _, batch) = matrix_meta_parts("tenferro-linalg.eigh_values", shape)?;
1272 Ok((singular_values_dtype(dtype), vector_shape(n, batch)))
1273}
1274
1275fn eig_meta(
1276 input_dtype: DType,
1277 shape: &[SymDim],
1278) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1279 let dtype = eig_output_dtype(input_dtype);
1280 let (n, _, batch) = matrix_meta_parts("tenferro-linalg.eig", shape)?;
1281 Ok(vec![
1282 (dtype, vector_shape(n.clone(), batch)),
1283 (dtype, matrix_shape(n.clone(), n, batch)),
1284 ])
1285}
1286
1287fn eig_values_meta(
1288 input_dtype: DType,
1289 shape: &[SymDim],
1290) -> tenferro_tensor::Result<(DType, Vec<SymDim>)> {
1291 let dtype = eig_output_dtype(input_dtype);
1292 let (n, _, batch) = matrix_meta_parts("tenferro-linalg.eig_values", shape)?;
1293 Ok((dtype, vector_shape(n, batch)))
1294}
1295
1296fn matrix_shape(rows: SymDim, cols: SymDim, batch: &[SymDim]) -> Vec<SymDim> {
1297 let mut shape = vec![rows, cols];
1298 shape.extend_from_slice(batch);
1299 shape
1300}
1301
1302fn vector_shape(len: SymDim, batch: &[SymDim]) -> Vec<SymDim> {
1303 let mut shape = vec![len];
1304 shape.extend_from_slice(batch);
1305 shape
1306}
1307
1308fn eig_output_dtype(dtype: DType) -> DType {
1309 match dtype {
1310 DType::F64 | DType::C64 => DType::C64,
1311 DType::F32 | DType::C32 => DType::C32,
1312 DType::I32 | DType::I64 | DType::Bool => DType::C64,
1313 }
1314}
1315
1316fn singular_values_dtype(dtype: DType) -> DType {
1317 match dtype {
1318 DType::C64 => DType::F64,
1319 DType::C32 => DType::F32,
1320 other => other,
1321 }
1322}
1323
1324fn promote_dtypes(dtypes: &[DType]) -> DType {
1325 dtypes
1326 .iter()
1327 .copied()
1328 .reduce(tenferro_tensor::validate::promote_dtype)
1329 .unwrap_or(DType::F64)
1330}
1331
1332fn hash_dtype(hasher: &mut dyn Hasher, dtype: DType) {
1333 let tag = match dtype {
1334 DType::F64 => 0,
1335 DType::F32 => 1,
1336 DType::I64 => 2,
1337 DType::C64 => 3,
1338 DType::C32 => 4,
1339 DType::I32 => 5,
1340 DType::Bool => 6,
1341 };
1342 hasher.write_u8(tag);
1343}
1344
1345fn hash_svd_gauge(hasher: &mut dyn Hasher, gauge: SvdGauge) {
1346 let tag = match gauge {
1347 SvdGauge::Raw => 0,
1348 SvdGauge::CanonicalPivot => 1,
1349 };
1350 hasher.write_u8(tag);
1351}
1352
1353fn hash_eigh_gauge(hasher: &mut dyn Hasher, gauge: EighGauge) {
1354 let tag = match gauge {
1355 EighGauge::Raw => 0,
1356 EighGauge::CanonicalPivot => 1,
1357 };
1358 hasher.write_u8(tag);
1359}
1360
1361fn hash_qr_gauge(hasher: &mut dyn Hasher, gauge: QrGauge) {
1362 let tag = match gauge {
1363 QrGauge::Raw => 0,
1364 QrGauge::PositiveDiagonal => 1,
1365 };
1366 hasher.write_u8(tag);
1367}