Files
tensor/linalg/sparselu_test.go
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

308 lines
8.7 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// 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")
}
}