// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) // SPDX-License-Identifier: MIT package linalg import ( "math" "slices" "strings" "testing" "sourcedock.dev/petrbalvin/tensor/internal/core" ) // nonsymmetricCOO builds a sparse nonsymmetric matrix from the seeded // generator: a diagonally dominant backbone with extra couplings that // force fill into the factor, large enough that pivoting has real // choices. func nonsymmetricCOO(t *testing.T, seed int64, n int) *core.SparseCOO { t.Helper() g := core.NewGenerator(seed) idx := make([]int64, 0, 6*n) vals := make([]float64, 0, 6*n) add := func(r, c int, v float64) { idx = append(idx, int64(r), int64(c)) vals = append(vals, v) } for i := range n { add(i, i, 6+2*g.Unit()) if i+1 < n { add(i, i+1, -1-g.Unit()) add(i+1, i, -1-g.Unit()) } if i+2 < n { add(i, i+2, -0.5*g.Unit()) } if i+3 < n { add(i+3, i, -0.3*g.Unit()) } if i+5 < n { add(i, i+5, -0.2*g.Unit()) } } indices, err := core.FromInts(idx, len(vals), 2) if err != nil { t.Fatalf("FromInts: %v", err) } coo, err := core.NewSparseCOO(indices, floatsToArray(vals, []int{len(vals)}), []int{n, n}) if err != nil { t.Fatalf("NewSparseCOO: %v", err) } return coo } func TestSparseLUSolvesNonsymmetric(t *testing.T) { const n = 20 for _, seed := range []int64{1, 2, 3} { coo := nonsymmetricCOO(t, seed, n) csr, err := CSRFromCOO(coo) if err != nil { t.Fatalf("CSRFromCOO: %v", err) } g := core.NewGenerator(seed + 100) xTrue := core.New(core.Float, n) for i := range n { xTrue.RawFloats()[i] = g.NormalUnit() } b, err := csr.MatVec(xTrue) if err != nil { t.Fatalf("MatVec: %v", err) } f, err := NewSparseLU(coo) if err != nil { t.Fatalf("NewSparseLU: %v", err) } x, err := f.Solve(b) if err != nil { t.Fatalf("Solve: %v", err) } worst := 0.0 for i := range n { if d := math.Abs(x.FloatAt(i) - xTrue.FloatAt(i)); d > worst { worst = d } } if worst > 1e-9 { t.Fatalf("seed %d: worst solution error %.3g, want under 1e-9", seed, worst) } } } // TestSparseLUMatchesDense requires the sparse factor to answer what // the dense LU with partial pivoting answers for the same system. func TestSparseLUMatchesDense(t *testing.T) { const n = 16 coo := nonsymmetricCOO(t, 7, n) csr, err := CSRFromCOO(coo) if err != nil { t.Fatalf("CSRFromCOO: %v", err) } g := core.NewGenerator(8) xTrue := core.New(core.Float, n) for i := range n { xTrue.RawFloats()[i] = g.NormalUnit() } b, err := csr.MatVec(xTrue) if err != nil { t.Fatalf("MatVec: %v", err) } f, err := NewSparseLU(coo) if err != nil { t.Fatalf("NewSparseLU: %v", err) } sparse, err := f.Solve(b) if err != nil { t.Fatalf("Solve: %v", err) } denseVals := make([]float64, n*n) nnz := coo.Indices.Shape()[0] for i := range nnz { r := int(coo.Indices.RawInts()[i*2]) c := int(coo.Indices.RawInts()[i*2+1]) denseVals[r*n+c] = coo.Values.FloatAt(i) } dense, err := Solve(floatsToArray(denseVals, []int{n, n}), b) if err != nil { t.Fatalf("dense Solve: %v", err) } scale := 0.0 for i := range n { if d := math.Abs(dense.FloatAt(i)); d > scale { scale = d } } for i := range n { if math.Abs(sparse.FloatAt(i)-dense.FloatAt(i)) > 1e-9*scale { t.Fatalf("entry %d: sparse %.12g vs dense %.12g", i, sparse.FloatAt(i), dense.FloatAt(i)) } } } // TestSparseLUPivotsRowSwap pins the reason partial pivoting exists: // the anti-diagonal matrix has a zero at (0,0) and only a row swap // can start the elimination, and the swap must be visible in the // reported permutation. func TestSparseLUPivotsRowSwap(t *testing.T) { anti, err := core.NewSparseCOO( mustInts(t, []int64{0, 1, 1, 0, 1, 1}, 3, 2), floatsToArray([]float64{2, 1, 3}, []int{3}), []int{2, 2}) if err != nil { t.Fatalf("NewSparseCOO: %v", err) } // Entries: (0,1)=2, (1,0)=1, (1,1)=3: the matrix [[0,2],[1,3]]. f, err := NewSparseLU(anti) if err != nil { t.Fatalf("NewSparseLU: %v", err) } if got := f.Permutation(); !slices.Equal(got, []int{1, 0}) { t.Fatalf("permutation %v, want the rows swapped to [1 0]", got) } b := floatsToArray([]float64{4, 5}, []int{2}) x, err := f.Solve(b) if err != nil { t.Fatalf("Solve: %v", err) } // [[0,2],[1,3]]·x = 4x2=4... solve by hand: x2 = 4/2 = 2, // x1 + 3·2 = 5 → x1 = −1. if math.Abs(x.FloatAt(0)-(-1)) > 1e-12 || math.Abs(x.FloatAt(1)-2) > 1e-12 { t.Fatalf("solve = [%.6g %.6g], want [-1 2]", x.FloatAt(0), x.FloatAt(1)) } } func TestSparseLURefusals(t *testing.T) { good := nonsymmetricCOO(t, 9, 6) if _, err := NewSparseLU(good); err != nil { t.Fatalf("NewSparseLU: %v", err) } // Complex input. cplx, err := core.NewSparseCOO( mustInts(t, []int64{0, 0}, 1, 2), mustFromComplexes(t, []complex128{1 + 1i}, 1), []int{1, 1}) if err != nil { t.Fatalf("NewSparseCOO: %v", err) } if _, err := NewSparseLU(cplx); err == nil || !strings.Contains(err.Error(), "complex") { t.Fatalf("complex input: %v", err) } // Rectangular input. rect, err := core.NewSparseCOO( mustInts(t, []int64{0, 0}, 1, 2), floatsToArray([]float64{1}, []int{1}), []int{1, 2}) if err != nil { t.Fatalf("NewSparseCOO: %v", err) } if _, err := NewSparseLU(rect); err == nil { t.Fatal("a rectangular matrix was accepted") } // Non-finite stored value. bad, err := core.NewSparseCOO( mustInts(t, []int64{0, 0, 1, 1}, 2, 2), floatsToArray([]float64{4, math.NaN()}, []int{2}), []int{2, 2}) if err != nil { t.Fatalf("NewSparseCOO: %v", err) } if _, err := NewSparseLU(bad); err == nil || !strings.Contains(err.Error(), "finite") { t.Fatalf("NaN entry: %v", err) } // Structurally singular: the last column stores nothing, and no // permutation can pivot an empty column. sing, err := core.NewSparseCOO( mustInts(t, []int64{0, 0, 1, 0, 1, 1, 2, 0}, 4, 2), floatsToArray([]float64{4, 1, 5, 6}, []int{4}), []int{3, 3}) if err != nil { t.Fatalf("NewSparseCOO: %v", err) } if _, err := NewSparseLU(sing); err == nil || !strings.Contains(err.Error(), "singular") { t.Fatalf("an empty column: %v", err) } // Solve-side refusals. f, err := NewSparseLU(good) if err != nil { t.Fatalf("NewSparseLU: %v", err) } if _, err := f.Solve(core.New(core.Float, 3, 3)); err == nil { t.Fatal("a rank 2 right hand side was accepted") } if _, err := f.Solve(floatsToArray([]float64{1, 2, 3}, []int{3})); err == nil { t.Fatal("a short right hand side was accepted") } } // TestSparseLUIsDeterministic factors the same matrix twice and // requires the factor's stored values to come out bit for bit equal, // the contract every Tensor entry point carries. func TestSparseLUIsDeterministic(t *testing.T) { coo := nonsymmetricCOO(t, 11, 18) f1, err := NewSparseLU(coo) if err != nil { t.Fatalf("first factorisation: %v", err) } f2, err := NewSparseLU(coo) if err != nil { t.Fatalf("second factorisation: %v", err) } if f1.NNZ() != f2.NNZ() { t.Fatalf("factor sizes differ: %d vs %d", f1.NNZ(), f2.NNZ()) } for j := range f1.n { if slices.Compare(f1.colRows[j], f2.colRows[j]) != 0 { t.Fatalf("column %d patterns differ", j) } for p := range f1.colRows[j] { if f1.colVals[j][p] != f2.colVals[j][p] { t.Fatalf("L value %d of column %d differs: %.17g vs %.17g", p, j, f1.colVals[j][p], f2.colVals[j][p]) } } if slices.Compare(f1.rowCols[j], f2.rowCols[j]) != 0 { t.Fatalf("row %d patterns differ", j) } for p := range f1.rowCols[j] { if f1.rowVals[j][p] != f2.rowVals[j][p] { t.Fatalf("U value %d of row %d differs: %.17g vs %.17g", p, j, f1.rowVals[j][p], f2.rowVals[j][p]) } } } if slices.Compare(f1.piv, f2.piv) != 0 { t.Fatal("permutations differ") } } // TestSparseLUFillStaysSparse pins the point of the sparse structure: // on a banded matrix the factor must stay close to the bandwidth, not // drift toward a dense triangle. The bound carries slack; the // measured number is recorded in the log line. func TestSparseLUFillStaysSparse(t *testing.T) { const n = 60 coo := nonsymmetricCOO(t, 13, n) f, err := NewSparseLU(coo) if err != nil { t.Fatalf("NewSparseLU: %v", err) } t.Logf("n=%d: LU nnz %d (bandwidth-5 pattern)", n, f.NNZ()) // NNZ counts L strict + U strict + the U diagonal (L's unit diagonal // excluded): a bandwidth-5 pattern holds about 2·5n strict entries, // so the bound sits at 10n, still far from the dense n²/2 triangle. if f.NNZ() > 10*n { t.Fatalf("LU nnz %d drifted past 10n on a banded pattern", f.NNZ()) } perm := f.Permutation() sorted := slices.Clone(perm) slices.Sort(sorted) for i := range sorted { if sorted[i] != i { t.Fatalf("permutation entry %d holds %d; not a permutation", i, sorted[i]) } } perm[0] = -1 if f.Permutation()[0] == -1 { t.Fatal("Permutation exposed the factor's internal slice") } }