Files
tensor/linalg/matrixfunc.go
T

404 lines
13 KiB
Go
Raw Permalink 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"
"math/cmplx"
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// Matrix functions through the eigen-decomposition. MatrixSqrt and
// MatrixLog evaluate √A and ln(A) for a symmetric positive
// semi-definite real matrix by diagonalising: A = VΛVᵀ, so
// f(A) = V·f(Λ)·Vᵀ. Every other input, a Hermitian complex matrix
// included, runs through the complex Schur decomposition described at
// the individual functions. The positive semi-definite requirement
// matters: the eigenvalues of such a matrix are non-negative, which
// makes the square root and the logarithm real. A matrix with a
// significantly negative eigenvalue is refused rather than silently
// producing complex or NaN entries.
// spectrumRelTol is the relative tolerance MatrixSqrt and the Schur
// routes of both matrix functions use to judge their spectrum: an
// eigenvalue counts as non-positive or singular when it sits within
// this fraction of the spectrum's magnitude. Purely relative on
// purpose: an absolute floor would refuse legitimate small-scale
// inputs such as diag(1e-12, 2e-12). The symmetric route of MatrixLog
// judges negativity against the rounding floor n·eps·λmax instead,
// which admits far smaller positive eigenvalues; see MatrixLog.
const spectrumRelTol = 1e-10
// MatrixSqrt returns the principal square root √A of a symmetric
// positive semi-definite real matrix, by diagonalisation: A = VΛVᵀ,
// so √A = V·√Λ·Vᵀ. Eigenvalues below −tol·‖A‖ are refused (the
// matrix is not positive semi-definite); small negative values from
// rounding are clamped to zero.
// For nonsymmetric matrices the answer runs through the complex Schur
// decomposition: A = Q·T·Qᴴ reduces the function to the triangular
// recurrence on T, the Schur-Parlett route. A real matrix answers in
// real entries whenever the principal function is real, which for the
// square root means no negative real eigenvalues and for the
// logarithm none on the non-positive real axis; otherwise the call is
// an error, honestly, rather than a complex answer in real clothing.
func MatrixSqrt(a *core.Array) (*core.Array, error) {
n := a.Shape()[0]
if a.NDim() != 2 || n != a.Shape()[1] {
return nil, base.Errf("MatrixSqrt: needs a square 2-D matrix, got shape %s", base.ShapeText(a.Shape()))
}
if !isSymmetricMatrix(a) {
return matrixSqrtSchur(a)
}
vals, vecs, err := Eigen(a)
if err != nil {
return nil, base.Errf("MatrixSqrt: %w", err)
}
// Eigen returns dense float arrays, so the reconstruction reads the
// payloads directly rather than paying floatAt's dispatch per entry
// of this O(n³) triple product.
scale := 0.0
for _, lam := range vals.RawFloats() {
if v := math.Abs(lam); v > scale {
scale = v
}
}
tol := spectrumRelTol * scale
sqrtVals := make([]float64, n)
for i, lam := range vals.RawFloats() {
if lam < -tol {
return nil, base.Errf("MatrixSqrt: eigenvalue %v is negative; the matrix is not positive semi-definite", lam)
}
if lam < 0 {
lam = 0
}
sqrtVals[i] = math.Sqrt(lam)
}
// √A = V·diag(√λ)·Vᵀ.
out := core.New(core.Float, []int{n, n}...)
vMat := vecs.RawFloats()
for i := range n {
for j := range n {
s := 0.0
for k := range n {
s += vMat[i*n+k] * sqrtVals[k] * vMat[j*n+k]
}
out.RawFloats()[i*n+j] = s
}
}
return out, nil
}
// MatrixLog returns the principal matrix logarithm ln(A) of a real
// symmetric positive-definite matrix, by the same diagonalisation
// route as MatrixSqrt with ln applied to the eigenvalues. A negative
// eigenvalue beyond the rounding floor n·eps·λmax is refused as
// non-positive, and so is an exact zero: the principal logarithm of a
// matrix with eigenvalues on the non-positive real axis leaves the
// reals. A positive eigenvalue is logged whatever its size against
// the spectrum, since ln 1e-300 is finite and exact to its own
// digits.
func MatrixLog(a *core.Array) (*core.Array, error) {
n := a.Shape()[0]
if a.NDim() != 2 || n != a.Shape()[1] {
return nil, base.Errf("MatrixLog: needs a square 2-D matrix, got shape %s", base.ShapeText(a.Shape()))
}
if !isSymmetricMatrix(a) {
return matrixLogSchur(a)
}
vals, vecs, err := Eigen(a)
if err != nil {
return nil, base.Errf("MatrixLog: %w", err)
}
// The negativity floor is the eigensolver's own rounding scale,
// n·eps·λmax, the floor the QR rank guard uses: a negative value
// beyond it is a genuine sign of a non-positive spectrum, one
// inside it is rounding of a zero, and a positive value is logged
// honestly however far below the spectrum's top it sits.
floor := float64(n) * base.EpsF * eigenMaxAbs(vals, n)
logVals := make([]float64, n)
for i, lam := range vals.RawFloats() {
if lam < -floor {
return nil, base.Errf("MatrixLog: eigenvalue %v is negative; the matrix is not positive definite", lam)
}
if lam <= 0 {
return nil, base.Errf("MatrixLog: eigenvalue %v is non-positive; the principal logarithm needs a positive-definite matrix", lam)
}
logVals[i] = math.Log(lam)
}
out := core.New(core.Float, []int{n, n}...)
vMat := vecs.RawFloats()
for i := range n {
for j := range n {
s := 0.0
for k := range n {
s += vMat[i*n+k] * logVals[k] * vMat[j*n+k]
}
out.RawFloats()[i*n+j] = s
}
}
return out, nil
}
// eigenMaxAbs returns the largest absolute eigenvalue.
func eigenMaxAbs(vals *core.Array, n int) float64 {
m := 0.0
for _, v := range vals.RawFloats() {
if v := math.Abs(v); v > m {
m = v
}
}
return m
}
// isSymmetricMatrix reports whether a square real matrix is symmetric
// to rounding, which routes the matrix functions to the cheaper
// symmetric eigendecomposition. The tolerance is purely relative to
// the matrix scale, deliberately without an absolute floor: a floor
// would call a small-scale matrix symmetric on asymmetry it carries in
// full relative measure, and the Eigen route then applies its own
// relative check and refuses the matrix the routing just approved.
// Relative-only keeps the two decisions in agreement, because any
// asymmetry this check passes is within the fraction of the scale the
// route's own guard tolerates.
func isSymmetricMatrix(a *core.Array) bool {
if a.NDim() != 2 || a.Shape()[0] != a.Shape()[1] || a.Dtype() == core.Complex {
return false
}
n := a.Shape()[0]
scale := 0.0
for i := range n * n {
scale = math.Max(scale, math.Abs(a.FloatAt(i)))
}
tol := 1e-12 * scale
for i := range n {
for j := i + 1; j < n; j++ {
if math.Abs(a.FloatAt(i*n+j)-a.FloatAt(j*n+i)) > tol {
return false
}
}
}
return true
}
// matrixSqrtSchur evaluates the principal square root through the
// complex Schur decomposition and the triangular recurrence. A real
// matrix with a negative real eigenvalue has a complex principal root
// and is refused; a real answer is otherwise recovered to rounding.
func matrixSqrtSchur(a *core.Array) (*core.Array, error) {
const name = "MatrixSqrt"
n := a.Shape()[0]
t, q, err := SchurComplex(a)
if err != nil {
return nil, base.Errf("%s: %w", name, err)
}
tm := make([]complex128, n*n)
copy(tm, t.RawComplexes())
if a.Dtype() != core.Complex {
tol := spectrumRelTol * cmplxAbsMax(tm, n)
for i := range n {
lam := tm[i*n+i]
if math.Abs(imag(lam)) <= tol && real(lam) < -tol {
return nil, base.Errf("%s: eigenvalue %v is negative real; the principal square root is complex", name, lam)
}
}
}
if err := triSqrt(tm, n); err != nil {
return nil, base.Errf("%s: %w", name, err)
}
return schurBackTransform(q, tm, n, a.Dtype() != core.Complex, name)
}
// matrixLogSchur evaluates the principal logarithm through the complex
// Schur decomposition by inverse scaling and squaring: repeated
// triangular square roots walk the matrix near the identity, where the
// Mercator series finishes the job, and the squarings unwind. A real
// matrix with an eigenvalue on the non-positive real axis has no real
// principal logarithm and is refused.
func matrixLogSchur(a *core.Array) (*core.Array, error) {
const name = "MatrixLog"
n := a.Shape()[0]
t, q, err := SchurComplex(a)
if err != nil {
return nil, base.Errf("%s: %w", name, err)
}
tm := make([]complex128, n*n)
copy(tm, t.RawComplexes())
logScale := 0.0
// Spectrum screen. Real input refuses eigenvalues on the
// non-positive real axis, where the principal logarithm leaves the
// reals; complex input refuses eigenvalues at or below the
// rounding floor of the spectrum, where the square-root walk below
// would diverge into garbage instead of converging.
tol := spectrumRelTol * cmplxAbsMax(tm, n)
for i := range n {
lam := tm[i*n+i]
if a.Dtype() != core.Complex {
if math.Abs(imag(lam)) <= tol && real(lam) <= tol {
return nil, base.Errf("%s: eigenvalue %v is on the non-positive real axis; the principal logarithm is complex", name, lam)
}
} else if cmplx.Abs(lam) <= tol {
return nil, base.Errf("%s: eigenvalue %v is at the rounding floor of the spectrum; the matrix is singular and has no principal logarithm", name, lam)
}
}
// Scale the spectrum towards one so the square-root walk converges
// in a handful of steps; the scalar shift unwinds afterwards.
c := 0.0
for i := range n {
c = math.Max(c, cmplx.Abs(tm[i*n+i]))
}
if c > 0 && (c < 0.5 || c > 2) {
logScale = math.Log(c)
for i := range n {
for j := range n {
tm[i*n+j] /= complex(c, 0)
}
}
}
// Square-root walk towards the identity.
squarings := 0
for range 60 {
dev := triDeviation(tm, n)
if dev <= 0.25 {
break
}
if err := triSqrt(tm, n); err != nil {
return nil, base.Errf("%s: %w", name, err)
}
squarings++
}
// A spectrum the walk cannot bring near the identity (typically a
// near-singular one the screen let through) would send the Mercator
// series below into open-ended divergence; refuse instead.
if dev := triDeviation(tm, n); dev > 0.25 {
return nil, base.Errf("%s: the square-root walk did not converge in 60 passes (deviation %.3g); the matrix is outside the domain of the principal logarithm", name, dev)
}
// Mercator series for log(I + E) on the near-identity triangle.
e := make([]complex128, n*n)
for i := range n {
for j := range n {
e[i*n+j] = tm[i*n+j]
}
e[i*n+i] -= 1
}
lg := make([]complex128, n*n)
// The series multiplies term by e each step, so the two buffers
// alternate instead of the product landing in a fresh matrix: every
// value read is the previous step's product exactly as it was when
// the product was copied back, because only the upper triangle is
// ever written and both buffers start with a zero lower triangle.
term, next := make([]complex128, n*n), make([]complex128, n*n)
copy(term, e)
for m := 1; m <= 40; m++ {
factor := complex(1/float64(m), 0)
if m%2 == 0 {
factor = complex(-1/float64(m), 0)
}
for i := range n * n {
lg[i] += factor * term[i]
}
if m == 40 {
break
}
for i := range n {
for j := i; j < n; j++ {
s := complex(0, 0)
for k := i; k <= j; k++ {
s += term[i*n+k] * e[k*n+j]
}
next[i*n+j] = s
}
}
tnorm := 0.0
for i := range n * n {
tnorm = math.Max(tnorm, cmplx.Abs(next[i]))
}
term, next = next, term
if tnorm <= 1e-17 {
break
}
}
// Unwind the walk: log(T) = 2^k·L + log(c)·I.
for i := range n * n {
lg[i] *= complex(math.Pow(2, float64(squarings)), 0)
}
if logScale != 0 {
for i := range n {
lg[i*n+i] += complex(logScale, 0)
}
}
return schurBackTransform(q, lg, n, a.Dtype() != core.Complex, name)
}
// schurBackTransform maps the triangular result back through the Schur
// unitary, q·S·qᴴ, and returns real entries for real input when the
// imaginary rounding cancels.
func schurBackTransform(q *core.Array, s []complex128, n int, wantReal bool, name string) (*core.Array, error) {
qm := q.RawComplexes()
w := make([]complex128, n*n)
for i := range n {
for j := range n {
sum := complex(0, 0)
for k := i; k < n; k++ {
sum += s[i*n+k] * cmplx.Conj(qm[j*n+k])
}
w[i*n+j] = sum
}
}
out := make([]complex128, n*n)
for i := range n {
for j := range n {
sum := complex(0, 0)
for k := range n {
sum += qm[i*n+k] * w[k*n+j]
}
out[i*n+j] = sum
}
}
if !wantReal {
res, _ := core.FromComplexes(out, n, n)
return res, nil
}
worstIm, worstRe := 0.0, 0.0
for i := range n * n {
worstIm = math.Max(worstIm, math.Abs(imag(out[i])))
worstRe = math.Max(worstRe, math.Abs(real(out[i])))
}
if worstIm > 1e-8*math.Max(1, worstRe) {
return nil, base.Errf("%s: the result is complex (imaginary magnitude %.3g); the real answer does not exist", name, worstIm)
}
res := core.New(core.Float, []int{n, n}...)
for i := range n * n {
res.RawFloats()[i] = real(out[i])
}
return res, nil
}
// cmplxAbsMax returns the largest diagonal magnitude of a complex
// triangular matrix.
func cmplxAbsMax(t []complex128, n int) float64 {
m := 0.0
for i := range n {
m = math.Max(m, cmplx.Abs(t[i*n+i]))
}
return m
}
// triDeviation measures how far an upper triangular matrix still sits
// from the identity: the largest magnitude among the diagonal offsets
// and the strict upper part, the convergence gauge of the square-root
// walk in matrixLogSchur.
func triDeviation(t []complex128, n int) float64 {
dev := 0.0
for i := range n {
dev = math.Max(dev, cmplx.Abs(t[i*n+i]-1))
}
for i := range n {
for j := i + 1; j < n; j++ {
dev = math.Max(dev, cmplx.Abs(t[i*n+j]))
}
}
return dev
}