Files

136 lines
4.1 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 (
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
import "testing"
// rootSetMatch reports whether the computed roots equal the wanted set
// within tolerance: every wanted root has a computed partner and every
// computed root is claimed. Magnitude-sorted input order carries no
// meaning for a conjugate pair beyond rounding.
func rootSetMatch(got, want []complex128, tol float64) bool {
if len(got) != len(want) {
return false
}
used := make([]bool, len(got))
for _, w := range want {
found := false
for i, g := range got {
if !used[i] && abs2c(g-w) <= tol*tol {
used[i] = true
found = true
break
}
}
if !found {
return false
}
}
return true
}
func abs2c(z complex128) float64 {
return real(z)*real(z) + imag(z)*imag(z)
}
// TestPolynomialRootsFactored pins roots read off factored forms: a
// real pair, a cubic with a conjugate pair, and a degree-one case.
func TestPolynomialRootsFactored(t *testing.T) {
cases := []struct {
coeffs []float64
want []complex128
}{
{[]float64{6, -5, 1}, []complex128{2, 3}},
{[]float64{-13, 17, -5, 1}, []complex128{1, 2 + 3i, 2 - 3i}},
{[]float64{-4, 2}, []complex128{2}},
}
for k, tc := range cases {
roots, err := PolynomialRoots(mustFloats(t, tc.coeffs, len(tc.coeffs)))
if err != nil {
t.Fatalf("case %d: PolynomialRoots: %v", k, err)
}
got := make([]complex128, roots.Len())
for i := range got {
got[i] = roots.ComplexAt(i)
}
if !rootSetMatch(got, tc.want, 1e-8) {
t.Fatalf("case %d: roots %v, want %v", k, got, tc.want)
}
}
}
// TestPolynomialRootsRepeated checks a double root: the companion is
// defective, so the squared-off accuracy is all a similarity-based
// solver can give and the tolerance says so.
func TestPolynomialRootsRepeated(t *testing.T) {
roots, err := PolynomialRoots(mustFloats(t, []float64{4, -4, 1}, 3))
if err != nil {
t.Fatalf("PolynomialRoots: %v", err)
}
for i := range roots.Len() {
if d := roots.ComplexAt(i) - 2; abs2c(d) > 1e-8 {
t.Fatalf("root %d = %v, want 2 ± 1e-8", i, roots.ComplexAt(i))
}
}
}
// TestPolynomialRootsTrailingZeros strips trailing zero coefficients:
// the degree is the true one and the roots are unchanged.
func TestPolynomialRootsTrailingZeros(t *testing.T) {
roots, err := PolynomialRoots(mustFloats(t, []float64{6, -5, 1, 0, 0}, 5))
if err != nil {
t.Fatalf("PolynomialRoots: %v", err)
}
if roots.Len() != 2 {
t.Fatalf("got %d roots, want 2", roots.Len())
}
if !rootSetMatch([]complex128{roots.ComplexAt(0), roots.ComplexAt(1)},
[]complex128{2, 3}, 1e-8) {
t.Fatalf("roots %v, want {2, 3}", roots.Shape())
}
}
// TestPolynomialRootsComplexCoefficients exercises the complex path:
// x − (1+2i) has the obvious root.
func TestPolynomialRootsComplexCoefficients(t *testing.T) {
coeffs, err := core.FromComplexes([]complex128{-(1 + 2i), 1}, 2)
if err != nil {
t.Fatalf("FromComplexes: %v", err)
}
roots, err := PolynomialRoots(coeffs)
if err != nil {
t.Fatalf("PolynomialRoots: %v", err)
}
if roots.Len() != 1 || abs2c(roots.ComplexAt(0)-(1+2i)) > 1e-12 {
t.Fatalf("root = %v, want 1+2i", roots.ComplexAt(0))
}
}
// TestPolynomialRootsErrors pins the degenerate contracts: a constant
// answers an empty vector, the zero polynomial is an error, and so are
// wrong shapes.
func TestPolynomialRootsErrors(t *testing.T) {
roots, err := PolynomialRoots(mustFloats(t, []float64{3}, 1))
if err != nil {
t.Fatalf("constant polynomial: %v", err)
}
if roots.Len() != 0 {
t.Fatalf("constant polynomial returned %d roots, want 0", roots.Len())
}
if _, err := PolynomialRoots(mustFloats(t, []float64{0, 0}, 2)); err == nil {
t.Fatal("expected an error for the zero polynomial")
}
if _, err := PolynomialRoots(mustFloats(t, nil)); err == nil {
t.Fatal("expected an error for empty coefficients")
}
matrix, _ := core.FromFloats([]float64{1, 0, 0, 1}, 2, 2)
if _, err := PolynomialRoots(matrix); err == nil {
t.Fatal("expected an error for a rank-2 coefficient array")
}
}