Files
tensor/linalg/polyroots_test.go
T
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

136 lines
4.1 KiB
Go
Raw 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 (
"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")
}
}