Files

171 lines
5.0 KiB
Go
Raw Permalink Normal View History

2026-09-03 10:00:00 +02:00
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package linalg
import (
"math"
"strings"
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// Regression tests: entry points that reached a
// complex array's float accessor (a panic, not an error), a
// preconditioner applied to a system of another dimension, and Krylov
// callbacks returning the wrong length.
// mustComplex builds a complex array.
func mustComplex(t *testing.T, vals []complex128, shape ...int) *core.Array {
t.Helper()
a, err := core.FromComplexes(vals, shape...)
if err != nil {
t.Fatalf("FromComplexes: %v", err)
}
return a
}
// TestComplexInputsAreRefused pins the dtype gate of the four entry
// points that read their arguments as real numbers.
func TestComplexInputsAreRefused(t *testing.T) {
c2 := mustComplex(t, []complex128{1 + 1i, 1 + 1i}, 2)
c3 := mustComplex(t, []complex128{1 + 1i, 2 + 2i, 3 + 3i}, 3)
t.Run("SolveTridiagonal", func(t *testing.T) {
if _, err := SolveTridiagonal(c2, c3, c2, c3); err == nil {
t.Fatal("expected an error for complex diagonals")
} else if !strings.Contains(err.Error(), "complex") {
t.Fatalf("error = %v, want a complex-dtype refusal", err)
}
})
t.Run("SolveCyclicTridiagonal", func(t *testing.T) {
if _, err := SolveCyclicTridiagonal(c3, c3, c3, c3); err == nil {
t.Fatal("expected an error for complex diagonals")
}
})
t.Run("FitPolynomial", func(t *testing.T) {
if _, err := FitPolynomial(c2, c2, 1); err == nil {
t.Fatal("expected an error for complex samples")
}
})
t.Run("NewCubicSpline", func(t *testing.T) {
if _, err := NewCubicSpline(c3, c3); err == nil {
t.Fatal("expected an error for complex knots")
}
})
}
// TestSparseILUDimensionMismatch pins the preconditioner check: an ILU
// built for another system must be refused rather than indexed.
func TestSparseILUDimensionMismatch(t *testing.T) {
mk := func(n int, vals []float64) *core.SparseCOO {
t.Helper()
idx := make([]int64, 0, n*2)
for i := range n {
idx = append(idx, int64(i), int64(i))
}
ia, err := core.FromInts(idx, n, 2)
if err != nil {
t.Fatal(err)
}
va, err := core.FromFloats(vals, n)
if err != nil {
t.Fatal(err)
}
sp, err := core.NewSparseCOO(ia, va, []int{n, n})
if err != nil {
t.Fatal(err)
}
return sp
}
big := mk(4, []float64{4, 4, 4, 4})
small := mk(3, []float64{3, 3, 3})
ilu, err := NewSparseILU(small)
if err != nil {
t.Fatalf("NewSparseILU: %v", err)
}
b := mustFloats(t, []float64{1, 1, 1, 1}, 4)
if _, err := SpSolve(big, b, 0, 50, ilu); err == nil {
t.Fatal("expected an error for a preconditioner of the wrong dimension")
} else if !strings.Contains(err.Error(), "dimension") {
t.Fatalf("error = %v, want a dimension refusal", err)
}
if _, err := SpSolveBiCGSTAB(big, b, 0, 50, ilu); err == nil {
t.Fatal("expected an error from BiCGSTAB for a preconditioner of the wrong dimension")
}
}
// TestGMRESCallbackLength pins the op contract: a callback that
// returns fewer elements than the problem has must be an error.
func TestGMRESCallbackLength(t *testing.T) {
b := mustFloats(t, []float64{1, 2}, 2)
short := mustFloats(t, []float64{1}, 1)
if _, err := GMRES(func(*core.Array) (*core.Array, error) { return short, nil }, b, 2, 5, 1e-12); err == nil {
t.Fatal("expected an error for an op returning the wrong length")
} else if !strings.Contains(err.Error(), "elements") {
t.Fatalf("error = %v, want a length refusal", err)
}
}
// TestQRLargeScale pins the Householder reflector at scales where the
// raw squares of the old norm overflowed: every intermediate of the
// reflector is now O(1) in scaled units, so the reconstruction holds at
// any magnitude.
func TestQRLargeScale(t *testing.T) {
for _, scale := range []float64{1, 1e10, 1e100, 1.2589e154, 1e200, 1e300} {
vals := []float64{2, -1, 0.5, 1, 3, -2, 4, 1, 1}
for i := range vals {
vals[i] *= scale
}
a, err := core.FromFloats(vals, 3, 3)
if err != nil {
t.Fatal(err)
}
q, r, err := QR(a)
if err != nil {
t.Fatalf("scale %g: %v", scale, err)
}
maxA := 0.0
for _, v := range vals {
maxA = math.Max(maxA, math.Abs(v))
}
worst := 0.0
for i := range 3 {
for j := range 3 {
acc := 0.0
for k := range 3 {
acc += q.FloatAt(i*3+k) * r.FloatAt(k*3+j)
}
worst = math.Max(worst, math.Abs(acc-vals[i*3+j])/maxA)
}
}
if !(worst < 1e-12) {
t.Errorf("scale %g: |QR−A|/|A| = %g, want below 1e-12", scale, worst)
}
}
}
// TestSparseILUNil pins the nil guard: an explicitly nil preconditioner
// is an error, not a dereference.
func TestSparseILUNil(t *testing.T) {
a, err := core.FromFloats([]float64{2, 3}, 2)
if err != nil {
t.Fatal(err)
}
idx, err := core.FromInts([]int64{0, 0, 1, 1}, 2, 2)
if err != nil {
t.Fatal(err)
}
sp, err := core.NewSparseCOO(idx, a, []int{2, 2})
if err != nil {
t.Fatal(err)
}
b, err := core.FromFloats([]float64{1, 1}, 2)
if err != nil {
t.Fatal(err)
}
if _, err := SpSolve(sp, b, 0, 10, nil); err == nil {
t.Fatal("expected an error for a nil preconditioner")
}
}