Files
tensor/linalg/tridiag_test.go
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

168 lines
4.9 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// 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
}