Files

168 lines
4.9 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"
"sourcedock.dev/petrbalvin/tensor/internal/core"
"testing"
)
// TestSolveTridiagonal matches the Thomas solve against the dense
// solver on a representative system.
func TestSolveTridiagonal(t *testing.T) {
// Tridiagonal with b = 4 on the main diagonal, ±1 off it.
n := 6
b := make([]float64, n)
a := make([]float64, n-1)
c := make([]float64, n-1)
d := make([]float64, n)
for i := range n {
b[i] = 4
d[i] = float64(i + 1)
if i > 0 {
a[i-1] = -1
}
if i < n-1 {
c[i] = 1
}
}
x, err := SolveTridiagonal(floatsToArray(a, []int{n - 1}), floatsToArray(b, []int{n}),
floatsToArray(c, []int{n - 1}), floatsToArray(d, []int{n}))
if err != nil {
t.Fatalf("SolveTridiagonal: %v", err)
}
full := make([]float64, n*n)
for i := range n {
full[i*n+i] = 4
if i > 0 {
full[i*n+i-1] = -1
}
if i < n-1 {
full[i*n+i+1] = 1
}
}
ref, err := Solve(mustFloats(t, full, n, n), mustFloats(t, d, n))
if err != nil {
t.Fatalf("Solve: %v", err)
}
for i := range n {
if math.Abs(x.FloatAt(i)-ref.FloatAt(i)) > 1e-12 {
t.Fatalf("x[%d] = %.16g, want %.16g", i, x.FloatAt(i), ref.FloatAt(i))
}
}
if _, err := SolveTridiagonal(floatsToArray(a, []int{n - 1}), floatsToArray(make([]float64, n), []int{n}),
floatsToArray(c, []int{n - 1}), floatsToArray(d, []int{n})); err == nil {
t.Fatal("zero pivot: want an error")
}
}
// TestSolveTridiagonalOneByOne pins the 1×1 system: the c and a
// diagonals are empty per the length contract, and the seed must not
// read them (it used to panic on c.FloatAt(0)).
func TestSolveTridiagonalOneByOne(t *testing.T) {
x, err := SolveTridiagonal(floatsToArray(nil, []int{0}), mustFloats(t, []float64{2}),
floatsToArray(nil, []int{0}), mustFloats(t, []float64{6}))
if err != nil {
t.Fatalf("SolveTridiagonal 1×1: %v", err)
}
if math.Abs(x.FloatAt(0)-3) > 1e-14 {
t.Fatalf("x[0] = %v, want 3", x.FloatAt(0))
}
// A 1×1 system with a zero main-diagonal entry is refused, not
// divided through.
if _, err := SolveTridiagonal(floatsToArray(nil, []int{0}), mustFloats(t, []float64{0}),
floatsToArray(nil, []int{0}), mustFloats(t, []float64{6})); err == nil {
t.Fatal("1×1 zero pivot: want an error")
}
}
// TestNewCubicSplineThreePoints regresses the same seed panic through
// its smallest caller: three knots build a 1×1 interior system.
func TestNewCubicSplineThreePoints(t *testing.T) {
xs := mustFloats(t, []float64{0, 1, 2})
ys := mustFloats(t, []float64{0, 1, 0})
s, err := NewCubicSpline(xs, ys)
if err != nil {
t.Fatalf("NewCubicSpline with 3 points: %v", err)
}
for i := range 3 {
if v := s.At(xs.FloatAt(i)); math.Abs(v-ys.FloatAt(i)) > 1e-14 {
t.Fatalf("S(%v) = %v, want %v", xs.FloatAt(i), v, ys.FloatAt(i))
}
}
}
// cyclicApply multiplies the cyclic tridiagonal system out: a[0] is
// the A[0][n−1] wrap corner, c[n−1] the A[n−1][0] corner.
func cyclicApply(a, b, c, x []float64) []float64 {
n := len(b)
out := make([]float64, n)
for i := range n {
out[i] = b[i] * x[i]
if i > 0 {
out[i] += a[i] * x[i-1]
}
if i < n-1 {
out[i] += c[i] * x[i+1]
}
}
out[0] += a[0] * x[n-1]
out[n-1] += c[n-1] * x[0]
return out
}
// TestSolveCyclicTridiagonalAsymmetricCorners pins the Sherman-Morrison
// corner assignment: the correction vector u ends with beta (= c[n−1]),
// not alpha (= a[0]). With alpha ≠ beta the swapped corner used to
// return a silently wrong answer.
func TestSolveCyclicTridiagonalAsymmetricCorners(t *testing.T) {
a := []float64{7, 2, 1, 3}
b := []float64{10, 11, 12, 13}
c := []float64{1, 4, 2, 5}
d := []float64{1, 2, 3, 4}
x, err := SolveCyclicTridiagonal(mustFloats(t, a), mustFloats(t, b), mustFloats(t, c), mustFloats(t, d))
if err != nil {
t.Fatalf("SolveCyclicTridiagonal: %v", err)
}
got := cyclicApply(a, b, c, []float64{x.FloatAt(0), x.FloatAt(1), x.FloatAt(2), x.FloatAt(3)})
for i := range len(d) {
if math.Abs(got[i]-d[i]) > 1e-9 {
t.Fatalf("(A·x)[%d] = %.16g, want %.16g", i, got[i], d[i])
}
}
}
// TestSolveCyclicTridiagonalSymmetric covers the classic periodic
// second-difference system, whose wrap corners are equal.
func TestSolveCyclicTridiagonalSymmetric(t *testing.T) {
n := 5
a := make([]float64, n)
b := make([]float64, n)
c := make([]float64, n)
d := make([]float64, n)
for i := range n {
a[i], b[i], c[i] = -1, 3, -1
d[i] = float64(i+1) * float64(i)
}
x, err := SolveCyclicTridiagonal(mustFloats(t, a), mustFloats(t, b), mustFloats(t, c), mustFloats(t, d))
if err != nil {
t.Fatalf("SolveCyclicTridiagonal: %v", err)
}
got := cyclicApply(a, b, c, mustSlice(x, n))
for i := range n {
if math.Abs(got[i]-d[i]) > 1e-9 {
t.Fatalf("(A·x)[%d] = %.16g, want %.16g", i, got[i], d[i])
}
}
}
func mustSlice(x *core.Array, n int) []float64 {
out := make([]float64, n)
for i := range n {
out[i] = x.FloatAt(i)
}
return out
}