308 lines
8.7 KiB
Go
308 lines
8.7 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (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")
|
||
}
|
||
}
|