// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) // SPDX-License-Identifier: MIT package linalg import ( "math" "math/cmplx" "sourcedock.dev/petrbalvin/tensor/internal/core" "strings" "testing" ) // complexMatrix builds a complex128 array from row-major values. func complexMatrix(t *testing.T, shape []int, vals ...complex128) *core.Array { t.Helper() total := 1 for _, d := range shape { total *= d } if len(vals) != total { t.Fatalf("value count %d does not fill %v", len(vals), shape) } a, err := core.ComplexFromArray(vals, shape...) if err != nil { t.Fatalf("ComplexFromArray(%v, %v): %v", vals, shape, err) } return a } func TestDetComplexRotation(t *testing.T) { // A rotation in the complex plane has unit determinant; scaling one // row by (2+i) multiplies the determinant by the same factor. angle := complex(0.6, -0.8) // unit magnitude m := complexMatrix(t, []int{2, 2}, angle, 0, 0, 1, ) det, err := DetComplex(m) if err != nil { t.Fatal(err) } // The complex diagonal keeps its phase, so the unit-modulus claim // holds on the magnitude only. if math.Abs(cmplx.Abs(det)-1) > 1e-12 { t.Fatalf("unitary |det| = %v", cmplx.Abs(det)) } scaled := complexMatrix(t, []int{2, 2}, 2+1i, 0, 3-4i, angle, ) det2, err := DetComplex(scaled) if err != nil { t.Fatal(err) } want := (2 + 1i) * angle if cmplx.Abs(det2-want) > 1e-12 { t.Fatalf("scaled det = %v, want %v", det2, want) } // A singular matrix yields zero without erroring. singular := complexMatrix(t, []int{2, 2}, 1, 2i, 2, 4i) detS, err := DetComplex(singular) if err != nil { t.Fatal(err) } if detS != 0 { t.Fatalf("singular det = %v, want 0", detS) } // The real Det keeps pointing complex callers at the complex twin. realShaped, _ := core.FromFloats([]float64{1, 0, 0, 1}, 2, 2) if _, err := Det(realShaped); err != nil { t.Fatalf("real det errored: %v", err) } if _, err := Det(complexMatrix(t, []int{2, 2}, 1, 0, 0, 1)); err == nil { t.Fatal("real Det accepted a complex matrix") } } func TestSolveInvComplexRoundTrip(t *testing.T) { a := complexMatrix(t, []int{3, 3}, 2+1i, 0, 1-1i, 0, 1-2i, 0.5, 1i, 1, 3+2i, ) inv, err := Inv(a) if err != nil { t.Fatal(err) } if inv.Dtype() != core.Complex { t.Fatalf("inverse dtype %v", inv.Dtype()) } product, err := core.MatMul2D(a, inv) if err != nil { t.Fatal(err) } for i := range 9 { row, col := i/3, i%3 want := complex(0, 0) if row == col { want = 1 } if got := product.RawComplexes()[i]; cmplx.Abs(got-want) > 1e-10 { t.Fatalf("A·A⁻¹[%d][%d] = %v, want %v", row, col, got, want) } } // Solve and verify by substitution for a vector right-hand side. b, _ := core.FromFloats([]float64{1, 2, 3}, 3) // real b promotes to complex x, err := Solve(a, b) if err != nil { t.Fatal(err) } check, err := core.MatMul2D(a, x) if err != nil { t.Fatal(err) } for i := range 3 { if got := check.RawComplexes()[i]; cmplx.Abs(got-complex(float64(i+1), 0)) > 1e-10 { t.Fatalf("(a·x)[%d] = %v, want %v", i, got, float64(i+1)) } } // Solving with an explicitly complex vector works symmetrically. bc := complexMatrix(t, []int{3}, 2+2i, 0, -1i) bcol, err := core.Reshape(bc, 3) if err != nil { t.Fatal(err) } xc, err := Solve(a, bcol) if err != nil { t.Fatal(err) } if xc.Dtype() != core.Complex || xc.Len() != 3 { t.Fatalf("complex solve shape/dtype: %v %v", xc.Shape(), xc.Dtype()) } // Singular systems error instead of returning garbage. singular := complexMatrix(t, []int{2, 2}, 1, 1i, 2, 2i) bad, _ := core.FromFloats([]float64{1, 1}, 2) if _, err := Solve(singular, bad); err == nil { t.Fatal("singular solve succeeded") } if _, err := Inv(singular); err == nil { t.Fatal("singular inverse succeeded") } } func TestKronComplexBlocks(t *testing.T) { a := complexMatrix(t, []int{2, 2}, 1+1i, 0, 0, 2-1i) b := complexMatrix(t, []int{2, 2}, 1, 2i, 3, 0) out, err := core.Kron(a, b) if err != nil { t.Fatal(err) } if out.Dtype() != core.Complex { t.Fatalf("kron dtype %v", out.Dtype()) } if got := out.Shape(); got[0] != 4 || got[1] != 4 { t.Fatalf("kron shape %v", got) } // Top-left block scales b by (1+1i); it sits on flat positions // k*4+l because blocks interleave in the outer product layout. type slot struct { idx int want complex128 } for _, s := range []slot{ {0, 1 + 1i}, {1, 2i * (1 + 1i)}, {4, 3 * (1 + 1i)}, {5, 0}, } { if got := out.RawComplexes()[s.idx]; got != s.want { t.Fatalf("block TL[%d] = %v, want %v", s.idx, got, s.want) } } // Bottom-right block scales b by (2−1i); identity via mixed pair. mixed, _ := core.FromFloats([]float64{1}, 1, 1) // real 1×1 eye := complexMatrix(t, []int{1, 1}, 2-1i) cross, err := core.Kron(mixed, eye) if err != nil { t.Fatal(err) } if cross.RawComplexes()[0] != 2-1i { t.Fatalf("mixed kron element = %v", cross.RawComplexes()[0]) } } func TestTraceComplexSum(t *testing.T) { m := complexMatrix(t, []int{2, 2}, 1+2i, 99, 99, 3-4i) s, err := core.TraceComplex(m) if err != nil { t.Fatal(err) } if s != 4-2i { t.Fatalf("trace = %v, want 4−2i", s) } nonSquare := complexMatrix(t, []int{1, 2}, 1, 2) if _, err := core.TraceComplex(nonSquare); err == nil { t.Fatal("non-square trace accepted") } // Real path keeps rejecting complexes by name. realM, _ := core.FromFloats([]float64{1, 2, 3, 4}, 2, 2) if _, err := core.TraceComplex(realM); err == nil { t.Fatal("TraceComplex accepted a real matrix") } } func TestPowIComplexExactness(t *testing.T) { z := mustComplexes(t, []complex128{1 + 1i, 2 - 1i}, 2) cubed, err := core.PowI(z, 3) if err != nil { t.Fatal(err) } wantFirst := (1 + 1i) * (1 + 1i) * (1 + 1i) // −2+2i if cubed.RawComplexes()[0] != wantFirst { t.Fatalf("(1+i)^3 = %v, want %v", cubed.RawComplexes()[0], wantFirst) } inverse, err := core.PowI(z, -1) if err != nil { t.Fatal(err) } one, err := core.Mul(inverse, z) if err != nil { t.Fatal(err) } for i := range 2 { if math.Abs(real(one.RawComplexes()[i])-1) > 1e-12 || math.Abs(imag(one.RawComplexes()[i])) > 1e-12 { t.Fatalf("z·z⁻¹[%d] = %v", i, one.RawComplexes()[i]) } } } // TestComplexSolveInvDetErrors moved with Solve, Inv and Det from the // root package: Solve and Inv run on complex systems through the same // LU kernel, while the real-only Det points at its complex twin. func TestComplexSolveInvDetErrors(t *testing.T) { sys := complexMatrix(t, []int{2, 2}, 2+1i, 0, 0, 3-1i, ) got, err := Inv(sys) if err != nil { t.Fatalf("Inv complex: %v", err) } if got.Dtype() != core.Complex { t.Fatalf("Inv complex dtype %v", got.Dtype()) } if _, err := Solve(sys, sys); err != nil { t.Fatalf("Solve complex: %v", err) } if _, err := Det(sys); err == nil || !strings.Contains(err.Error(), "DetComplex") { t.Fatalf("Det complex: %v", err) } }