168 lines
4.9 KiB
Go
168 lines
4.9 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"
|
||
)
|
||
|
||
// 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
|
||
}
|