Files
tensor/linalg/schur_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

198 lines
6.2 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"
)
// 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))
}
}
}