404 lines
13 KiB
Go
404 lines
13 KiB
Go
// 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
|
|||
|
|
}
|