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