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