Files
tensor/linalg/schur_test.go
T

198 lines
6.2 KiB
Go
Raw 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 (
"math"
"sourcedock.dev/petrbalvin/tensor/internal/core"
"testing"
)
// nonsymSample builds a deterministic nonsymmetric square matrix.
func nonsymSample(n int) *core.Array {
vals := make([]float64, n*n)
for i := range n * n {
vals[i] = math.Sin(float64(3*i + 2))
}
a, _ := core.FromFloats(vals, n, n)
return a
}
// TestSchurComplexContracts pins the decomposition: a = q·t·qᴴ to
// rounding, q unitary, t upper triangular, and the diagonal of t the
// same set of eigenvalues EigenGeneral reports.
func TestSchurComplexContracts(t *testing.T) {
n := 5
a := nonsymSample(n)
tm, q, err := SchurComplex(a)
if err != nil {
t.Fatalf("SchurComplex: %v", err)
}
scale := 0.0
for i := range n * n {
scale = math.Max(scale, absComplex(tm.ComplexAt(i)))
}
// Reconstruction a ≈ q·t·qᴴ.
w := make([]complex128, n*n)
for i := range n {
for j := range n {
s := complex(0, 0)
for k := range n {
s += q.ComplexAt(i*n+k) * tm.ComplexAt(k*n+j)
}
w[i*n+j] = s
}
}
recon := 0.0
for i := range n {
for j := range n {
s := complex(0, 0)
for k := range n {
s += w[i*n+k] * cmplxConj(q.ComplexAt(j*n+k))
}
recon = math.Max(recon, absComplex(s-a.ComplexAt(i*n+j)))
}
}
if recon > 1e-10*math.Max(1, scale) {
t.Fatalf("reconstruction error %.3g", recon)
}
// t strictly upper triangular below the diagonal.
for i := 1; i < n; i++ {
for j := 0; j < i; j++ {
if absComplex(tm.ComplexAt(i*n+j)) > 1e-10*math.Max(1, scale) {
t.Fatalf("t[%d][%d] = %v not deflated", i, j, tm.ComplexAt(i*n+j))
}
}
}
// Eigenvalue set match against EigenGeneral.
want, _, err := EigenGeneral(a)
if err != nil {
t.Fatalf("EigenGeneral: %v", err)
}
got := make([]complex128, n)
for i := range n {
got[i] = tm.ComplexAt(i*n + i)
}
wants := make([]complex128, n)
for i := range n {
wants[i] = want.ComplexAt(i)
}
if !rootSetMatch(got, wants, 1e-6) {
t.Fatalf("Schur eigenvalues %v do not match EigenGeneral's %v", got, wants)
}
}
func cmplxConj(z complex128) complex128 { return complex(real(z), -imag(z)) }
// TestMatrixSqrtNonsymmetric pins exact principal roots: the square of
// an upper triangular matrix recovers it, and the rotation-like pair
// whose square has purely imaginary eigenvalues comes back too.
func TestMatrixSqrtNonsymmetric(t *testing.T) {
// [[2,1],[0,3]]² = [[4,5],[0,9]] and the principal root is the
// original: eigenvalues 2, 3 sit on the positive real axis.
a := mustFloats(t, []float64{4, 5, 0, 9}, 2, 2)
s, err := MatrixSqrt(a)
if err != nil {
t.Fatalf("MatrixSqrt: %v", err)
}
if math.Abs(s.FloatAt(0)-2) > 1e-8 || math.Abs(s.FloatAt(1)-1) > 1e-8 ||
math.Abs(s.FloatAt(2)) > 1e-8 || math.Abs(s.FloatAt(3)-3) > 1e-8 {
t.Fatalf("sqrt = [[%.10g, %.10g], [%.10g, %.10g]], want [[2, 1], [0, 3]]",
s.FloatAt(0), s.FloatAt(1), s.FloatAt(2), s.FloatAt(3))
}
// [[1,1],[−1,1]]² = [[0,2],[−2,0]]: the principal root of the
// rotation-doubler is the original, whose eigenvalues 1±i are the
// principal square roots of ±2i.
b := mustFloats(t, []float64{0, 2, -2, 0}, 2, 2)
s2, err := MatrixSqrt(b)
if err != nil {
t.Fatalf("MatrixSqrt rotation pair: %v", err)
}
if math.Abs(s2.FloatAt(0)-1) > 1e-8 || math.Abs(s2.FloatAt(1)-1) > 1e-8 ||
math.Abs(s2.FloatAt(2)+1) > 1e-8 || math.Abs(s2.FloatAt(3)-1) > 1e-8 {
t.Fatalf("sqrt = [[%.10g, %.10g], [%.10g, %.10g]], want [[1, 1], [−1, 1]]",
s2.FloatAt(0), s2.FloatAt(1), s2.FloatAt(2), s2.FloatAt(3))
}
}
// TestMatrixLogNonsymmetric walks the round trip through the matrix
// exponential: log(exp(B)) must return B for a nonsymmetric B, and the
// logarithm of a triangular matrix must match the analytic entries.
func TestMatrixLogNonsymmetric(t *testing.T) {
b := mustFloats(t, []float64{1, 0.5, 0, 2}, 2, 2)
a, err := MatrixExp(b)
if err != nil {
t.Fatalf("MatrixExp: %v", err)
}
lg, err := MatrixLog(a)
if err != nil {
t.Fatalf("MatrixLog: %v", err)
}
for i := range 4 {
if math.Abs(lg.FloatAt(i)-b.FloatAt(i)) > 1e-6 {
t.Fatalf("log(exp(B))[%d] = %.12g, want %.12g", i, lg.FloatAt(i), b.FloatAt(i))
}
}
// Rotation generator: exp(B) is a plane rotation with eigenvalues
// e^{±iθ}, squarely in the nonsymmetric complex-eigenvalue case.
const theta = 0.7
rot := mustFloats(t, []float64{
math.Cos(theta), -math.Sin(theta),
math.Sin(theta), math.Cos(theta),
}, 2, 2)
lg2, err := MatrixLog(rot)
if err != nil {
t.Fatalf("MatrixLog rotation: %v", err)
}
back, err := MatrixExp(lg2)
if err != nil {
t.Fatalf("MatrixExp round trip: %v", err)
}
for i := range 4 {
if math.Abs(back.FloatAt(i)-rot.FloatAt(i)) > 1e-6 {
t.Fatalf("exp(log(R))[%d] = %.12g, want %.12g", i, back.FloatAt(i), rot.FloatAt(i))
}
}
}
// TestMatrixFunctionRefusals pins the honest refusals: a negative real
// eigenvalue blocks the principal square root and anything on the
// non-positive real axis blocks the principal logarithm.
func TestMatrixFunctionRefusals(t *testing.T) {
negSqrt := mustFloats(t, []float64{-4, 1, 0, 1}, 2, 2)
if _, err := MatrixSqrt(negSqrt); err == nil {
t.Fatal("expected an error for a square root with a negative real eigenvalue")
}
negLog := mustFloats(t, []float64{-2, 1, 0, 4}, 2, 2)
if _, err := MatrixLog(negLog); err == nil {
t.Fatal("expected an error for a logarithm with a negative real eigenvalue")
}
singLog := mustFloats(t, []float64{0, 1, 0, 4}, 2, 2)
if _, err := MatrixLog(singLog); err == nil {
t.Fatal("expected an error for a logarithm at zero")
}
}
// TestMatrixSqrtComplexInput checks the complex path: the principal
// root of diag(1+i, 4) squares back to the input.
func TestMatrixSqrtComplexInput(t *testing.T) {
a, _ := core.FromComplexes([]complex128{1 + 1i, 0, 0, 4}, 2, 2)
s, err := MatrixSqrt(a)
if err != nil {
t.Fatalf("MatrixSqrt complex: %v", err)
}
if s.Dtype() != core.Complex {
t.Fatal("the complex answer must stay complex")
}
sq, err := core.MatMul2D(s, s)
if err != nil {
t.Fatalf("MatMul2D: %v", err)
}
for i := range 4 {
if absComplex(sq.ComplexAt(i)-a.ComplexAt(i)) > 1e-10 {
t.Fatalf("S·S[%d] = %v, want %v", i, sq.ComplexAt(i), a.ComplexAt(i))
}
}
}