209 lines
5.4 KiB
Go
209 lines
5.4 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||
// SPDX-License-Identifier: MIT
|
||
|
||
package linalg
|
||
|
||
import (
|
||
"math"
|
||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||
"testing"
|
||
)
|
||
|
||
// triDiagCOO builds a tridiagonal SparseCOO with the given diagonals.
|
||
func triDiagCOO(t *testing.T, n int, lo, diag, up float64) *core.SparseCOO {
|
||
t.Helper()
|
||
idx := make([]int64, 0, 3*n)
|
||
vals := make([]float64, 0, 3*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, diag)
|
||
if i+1 < n {
|
||
add(i, i+1, up)
|
||
add(i+1, i, lo)
|
||
}
|
||
}
|
||
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
|
||
}
|
||
|
||
// TestSparseILUTridiagonalExact uses the fact that a tridiagonal
|
||
// matrix creates no fill under LU: the ILU(0) factors must reproduce
|
||
// the complete LU exactly, so applying them to a unit vector answers
|
||
// what the dense Solve answers for the same right-hand side.
|
||
func TestSparseILUTridiagonalExact(t *testing.T) {
|
||
const n = 12
|
||
coo := triDiagCOO(t, n, -1, 4, -2)
|
||
ilu, err := NewSparseILU(coo)
|
||
if err != nil {
|
||
t.Fatalf("NewSparseILU: %v", err)
|
||
}
|
||
denseVals := make([]float64, n*n)
|
||
for i := range n {
|
||
denseVals[i*n+i] = 4
|
||
if i+1 < n {
|
||
denseVals[i*n+i+1] = -2
|
||
denseVals[(i+1)*n+i] = -1
|
||
}
|
||
}
|
||
dense := mustFloats(t, denseVals, n, n)
|
||
for j := range n {
|
||
e := make([]float64, n)
|
||
e[j] = 1
|
||
got := ilu.Apply(e)
|
||
ref, err := Solve(dense, mustFloats(t, e))
|
||
if err != nil {
|
||
t.Fatalf("Solve: %v", err)
|
||
}
|
||
for i := range n {
|
||
if math.Abs(got[i]-ref.FloatAt(i)) > 1e-11 {
|
||
t.Fatalf("column %d, row %d: %.14g, want %.14g", j, i, got[i], ref.FloatAt(i))
|
||
}
|
||
}
|
||
}
|
||
}
|
||
|
||
// TestSparseILUWiderStencil checks the wider-than-tridiagonal case:
|
||
// a five-point stencil with dropped fill cannot reproduce A exactly,
|
||
// but the preconditioned residual of the factorisation must still be
|
||
// far smaller than the identity preconditioner's.
|
||
func TestSparseILUWiderStencil(t *testing.T) {
|
||
const n = 20
|
||
idx := make([]int64, 0, 5*n)
|
||
vals := make([]float64, 0, 5*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, 5)
|
||
if i+1 < n {
|
||
add(i, i+1, -2)
|
||
add(i+1, i, -1)
|
||
}
|
||
if i+2 < n {
|
||
add(i, i+2, 0.3)
|
||
add(i+2, i, -0.2)
|
||
}
|
||
}
|
||
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)
|
||
}
|
||
ilu, err := NewSparseILU(coo)
|
||
if err != nil {
|
||
t.Fatalf("NewSparseILU: %v", err)
|
||
}
|
||
rhs := make([]float64, n)
|
||
for i := range n {
|
||
rhs[i] = math.Cos(float64(i))
|
||
}
|
||
x := ilu.Apply(rhs)
|
||
// The residual ‖r − A·x‖ must sit well below ‖r‖: the incomplete
|
||
// factorisation captures most of A.
|
||
res := 0.0
|
||
rhsNorm := 0.0
|
||
for i := range n {
|
||
s := rhs[i]
|
||
rowSum := 0.0
|
||
for k := range n {
|
||
v := 0.0
|
||
switch {
|
||
case k == i:
|
||
v = 5
|
||
case k == i+1:
|
||
v = -2
|
||
case k == i-1:
|
||
v = -1
|
||
case k == i+2:
|
||
v = 0.3
|
||
case k == i-2:
|
||
v = -0.2
|
||
}
|
||
rowSum += v * x[k]
|
||
}
|
||
d := s - rowSum
|
||
res += d * d
|
||
rhsNorm += s * s
|
||
}
|
||
if math.Sqrt(res) > 0.5*math.Sqrt(rhsNorm) {
|
||
t.Fatalf("ILU residual %g too large against ‖r‖ %g", math.Sqrt(res), math.Sqrt(rhsNorm))
|
||
}
|
||
}
|
||
|
||
// TestSparseILUPreconditionedSteps checks the preconditioner's purpose:
|
||
// on a convective 100×100 system the ILU-preconditioned BiCGSTAB must
|
||
// reach the same tolerance as the Jacobi run, with a residual that
|
||
// meets the bound.
|
||
func TestSparseILUPreconditionedSteps(t *testing.T) {
|
||
const n = 100
|
||
coo := triDiagCOO(t, n, -1+0.5, 4, -1)
|
||
bv := make([]float64, n)
|
||
for i := range n {
|
||
bv[i] = math.Sin(float64(i+1) / float64(n+1))
|
||
}
|
||
b := mustFloats(t, bv)
|
||
tol := 1e-10
|
||
xJ, err := SpSolveBiCGSTAB(coo, b, tol, 0)
|
||
if err != nil {
|
||
t.Fatalf("Jacobi run: %v", err)
|
||
}
|
||
ilu, err := NewSparseILU(coo)
|
||
if err != nil {
|
||
t.Fatalf("NewSparseILU: %v", err)
|
||
}
|
||
xI, err := SpSolveBiCGSTAB(coo, b, tol, 0, ilu)
|
||
if err != nil {
|
||
t.Fatalf("ILU run: %v", err)
|
||
}
|
||
for i := range n {
|
||
if math.Abs(xJ.FloatAt(i)-xI.FloatAt(i)) > 1e-8 {
|
||
t.Fatalf("solutions disagree at %d: %.14g vs %.14g", i,
|
||
xJ.FloatAt(i), xI.FloatAt(i))
|
||
}
|
||
}
|
||
// The tridiagonal ILU is the exact LU: the ILU-preconditioned
|
||
// system must converge in very few steps. Cap the budget where
|
||
// plain Jacobi needs far more.
|
||
if _, err := SpSolveBiCGSTAB(coo, b, tol, 3, ilu); err != nil {
|
||
t.Fatalf("exact-LU preconditioner should converge in 3 steps: %v", err)
|
||
}
|
||
}
|
||
|
||
func TestSparseILUErrors(t *testing.T) {
|
||
if _, err := NewSparseILU(triDiagCOO(t, 3, -1, 0, -1)); err == nil {
|
||
t.Fatal("zero pivot: want an error")
|
||
}
|
||
rect, err := core.NewSparseCOO(mustInts2(t, []int64{0, 0, 0, 1, 0, 2}, 3, 2),
|
||
floatsToArray([]float64{1, 2, 3}, []int{3}), []int{2, 3})
|
||
if err != nil {
|
||
t.Fatalf("NewSparseCOO: %v", err)
|
||
}
|
||
if _, err := NewSparseILU(rect); err == nil {
|
||
t.Fatal("non-square: want an error")
|
||
}
|
||
}
|
||
|
||
// mustInts2 builds a 2-D int array for sparse constructors.
|
||
func mustInts2(t *testing.T, vals []int64, rows, cols int) *core.Array {
|
||
t.Helper()
|
||
a, err := core.FromInts(vals, rows, cols)
|
||
if err != nil {
|
||
t.Fatalf("FromInts: %v", err)
|
||
}
|
||
return a
|
||
}
|