feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -0,0 +1,307 @@
|
||||
// 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")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user