264 lines
7.9 KiB
Go
264 lines
7.9 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"
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
// The complex Schur decomposition and the principal matrix functions
|
|||
|
|
// built on it. A = Q·T·Qᴴ with Q unitary and T upper triangular turns
|
|||
|
|
// f(A) into Q·f(T)·Qᴴ, and f of a triangular matrix is a recurrence:
|
|||
|
|
// the diagonal maps entrywise and the off-diagonals follow from the
|
|||
|
|
// function's own defining equation. That is the Schur-Parlett route to
|
|||
|
|
// the principal square root and logarithm of nonsymmetric matrices,
|
|||
|
|
// where the symmetric eigendecomposition the easy route needs does not
|
|||
|
|
// exist.
|
|||
|
|
|
|||
|
|
// SchurComplex returns the complex Schur decomposition of the square
|
|||
|
|
// matrix a: t upper triangular and q unitary with a = q·t·qᴴ. For a
|
|||
|
|
// real matrix whose eigenvalues come in conjugate pairs t is complex;
|
|||
|
|
// the real Schur form is a different construction. The diagonal of t
|
|||
|
|
// carries the eigenvalues, matching EigenGeneral's set. An input
|
|||
|
|
// other than a square 2-D matrix or a non-converging iteration is an
|
|||
|
|
// error.
|
|||
|
|
func SchurComplex(a *core.Array) (t, q *core.Array, err error) {
|
|||
|
|
const name = "SchurComplex"
|
|||
|
|
if a.NDim() != 2 || a.Shape()[0] != a.Shape()[1] {
|
|||
|
|
return nil, nil, base.Errf("%s: needs a square 2-D matrix, got shape %s", name, base.ShapeText(a.Shape()))
|
|||
|
|
}
|
|||
|
|
n := a.Shape()[0]
|
|||
|
|
if n == 0 {
|
|||
|
|
return nil, nil, base.Errf("%s: zero-sized matrix", name)
|
|||
|
|
}
|
|||
|
|
h := make([]complex128, n*n)
|
|||
|
|
if a.Dtype() == core.Complex && !a.Strided() && len(a.RawComplexes()) == n*n {
|
|||
|
|
copy(h, a.RawComplexes())
|
|||
|
|
} else {
|
|||
|
|
for i := range n * n {
|
|||
|
|
h[i] = a.ComplexAt(i)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
// Outside the safe window the Hessenberg reduction's squared
|
|||
|
|
// magnitudes and the Givens denominators leave the normal range: a
|
|||
|
|
// tiny matrix has norm == 0 for every reflector, so the reduction is
|
|||
|
|
// skipped and the iteration then runs on a matrix that is not
|
|||
|
|
// Hessenberg, and a huge one has den == 0 destroy the subdiagonals.
|
|||
|
|
// The matrix is moved into the window for the reduction and the
|
|||
|
|
// iteration, and the triangular factor, which carries the scale, is
|
|||
|
|
// moved back before it is returned; the unitary factor is
|
|||
|
|
// scale-free.
|
|||
|
|
ws := windowScale(maxMagComplex(h))
|
|||
|
|
if ws != 1 {
|
|||
|
|
scaleComplexes(h, ws)
|
|||
|
|
}
|
|||
|
|
vecs := schurHessenberg(h, n)
|
|||
|
|
if err := schurQR(h, vecs, n); err != nil {
|
|||
|
|
return nil, nil, base.Errf("%s: %w", name, err)
|
|||
|
|
}
|
|||
|
|
if ws != 1 {
|
|||
|
|
unscaleComplexes(h, ws)
|
|||
|
|
}
|
|||
|
|
tArr := core.New(core.Complex, []int{n, n}...)
|
|||
|
|
copy(tArr.RawComplexes(), h)
|
|||
|
|
qArr := core.New(core.Complex, []int{n, n}...)
|
|||
|
|
copy(qArr.RawComplexes(), vecs)
|
|||
|
|
return tArr, qArr, nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// schurHessenberg reduces a in place to upper Hessenberg form by
|
|||
|
|
// complex Householder similarity transforms, accumulating the unitary
|
|||
|
|
// into q (seeded with the identity).
|
|||
|
|
func schurHessenberg(a []complex128, n int) []complex128 {
|
|||
|
|
q := make([]complex128, n*n)
|
|||
|
|
for i := range n {
|
|||
|
|
q[i*n+i] = 1
|
|||
|
|
}
|
|||
|
|
// Reflector scratch reused across the sweep; each step uses the
|
|||
|
|
// first n−k−1 entries.
|
|||
|
|
vecBuf := make([]complex128, n)
|
|||
|
|
for k := 0; k < n-2; k++ {
|
|||
|
|
norm := 0.0
|
|||
|
|
for i := k + 1; i < n; i++ {
|
|||
|
|
norm += base.Real2(a[i*n+k])
|
|||
|
|
}
|
|||
|
|
norm = math.Sqrt(norm)
|
|||
|
|
if norm == 0 {
|
|||
|
|
continue
|
|||
|
|
}
|
|||
|
|
head := a[(k+1)*n+k]
|
|||
|
|
beta := -norm
|
|||
|
|
if real(head) < 0 {
|
|||
|
|
beta = norm
|
|||
|
|
}
|
|||
|
|
vec := vecBuf[:n-k-1]
|
|||
|
|
for i := k + 1; i < n; i++ {
|
|||
|
|
vec[i-k-1] = a[i*n+k]
|
|||
|
|
if i == k+1 {
|
|||
|
|
vec[0] -= complex(beta, 0)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
if base.Real2(vec[0]) == 0 {
|
|||
|
|
continue
|
|||
|
|
}
|
|||
|
|
vhx := complex(0, 0)
|
|||
|
|
for i := k + 1; i < n; i++ {
|
|||
|
|
vhx += cmplx.Conj(vec[i-k-1]) * a[i*n+k]
|
|||
|
|
}
|
|||
|
|
tau := complex(1, 0) / vhx
|
|||
|
|
// Left: Ḣ·x = βe₁ with the plain tau, so the similarity is
|
|||
|
|
// H <- Ḣ·H·Ḣᴴ and the accumulated unitary picks up Ḣᴴ.
|
|||
|
|
for j := k; j < n; j++ {
|
|||
|
|
s := complex(0, 0)
|
|||
|
|
for i := k + 1; i < n; i++ {
|
|||
|
|
s += cmplx.Conj(vec[i-k-1]) * a[i*n+j]
|
|||
|
|
}
|
|||
|
|
s *= tau
|
|||
|
|
for i := k + 1; i < n; i++ {
|
|||
|
|
a[i*n+j] -= s * vec[i-k-1]
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
// Right: all rows over columns k+1..n-1, by Ḣᴴ.
|
|||
|
|
for i := range n {
|
|||
|
|
s := complex(0, 0)
|
|||
|
|
for j := k + 1; j < n; j++ {
|
|||
|
|
s += a[i*n+j] * vec[j-k-1]
|
|||
|
|
}
|
|||
|
|
s *= cmplx.Conj(tau)
|
|||
|
|
for j := k + 1; j < n; j++ {
|
|||
|
|
a[i*n+j] -= s * cmplx.Conj(vec[j-k-1])
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
for r := range n {
|
|||
|
|
s := complex(0, 0)
|
|||
|
|
for j := k + 1; j < n; j++ {
|
|||
|
|
s += q[r*n+j] * vec[j-k-1]
|
|||
|
|
}
|
|||
|
|
s *= cmplx.Conj(tau)
|
|||
|
|
for j := k + 1; j < n; j++ {
|
|||
|
|
q[r*n+j] -= s * cmplx.Conj(vec[j-k-1])
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
return q
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// schurQR drives the shifted QR iteration on the Hessenberg a to upper
|
|||
|
|
// triangular form, folding each rotation into the unitary q. Explicit
|
|||
|
|
// Givens QR of the active block with the Wilkinson shift keeps the
|
|||
|
|
// bookkeeping simple; convergence deflates subdiagonals to zero.
|
|||
|
|
func schurQR(a []complex128, q []complex128, n int) error {
|
|||
|
|
// Rotation scratch reused across sweeps; each sweep uses its first
|
|||
|
|
// 2·(r−l+1) entries.
|
|||
|
|
rots := make([]complex128, 2*n)
|
|||
|
|
for iter := 0; iter < 120*n+200; iter++ {
|
|||
|
|
// Deflate negligible subdiagonals.
|
|||
|
|
for i := 1; i < n; i++ {
|
|||
|
|
if cmplx.Abs(a[i*n+i-1]) <= base.EpsF*(cmplx.Abs(a[(i-1)*n+i-1])+cmplx.Abs(a[i*n+i])) {
|
|||
|
|
a[i*n+i-1] = 0
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
// Active trailing block [l..r].
|
|||
|
|
r := n - 1
|
|||
|
|
for r > 0 && a[r*n+r-1] == 0 {
|
|||
|
|
r--
|
|||
|
|
}
|
|||
|
|
if r == 0 {
|
|||
|
|
return nil
|
|||
|
|
}
|
|||
|
|
l := r
|
|||
|
|
for l > 0 && a[l*n+l-1] != 0 {
|
|||
|
|
l--
|
|||
|
|
}
|
|||
|
|
// Wilkinson shift from the trailing two-by-two.
|
|||
|
|
t11, t12 := a[(r-1)*n+r-1], a[(r-1)*n+r]
|
|||
|
|
t21, t22 := a[r*n+r-1], a[r*n+r]
|
|||
|
|
disc := cmplx.Sqrt((t11-t22)*(t11-t22) + 4*t21*t12)
|
|||
|
|
root1 := (t11 + t22 + disc) / 2
|
|||
|
|
root2 := (t11 + t22 - disc) / 2
|
|||
|
|
mu := root1
|
|||
|
|
if cmplx.Abs(root2-t22) < cmplx.Abs(root1-t22) {
|
|||
|
|
mu = root2
|
|||
|
|
}
|
|||
|
|
// Explicit shifted QR of the block by Givens rotations: the
|
|||
|
|
// shift leaves the diagonal first and returns at the end, so
|
|||
|
|
// the factorisation runs on H − μI exactly.
|
|||
|
|
for i := l; i <= r; i++ {
|
|||
|
|
a[i*n+i] -= mu
|
|||
|
|
}
|
|||
|
|
rr := rots[:2*(r-l+1)]
|
|||
|
|
for i := l; i < r; i++ {
|
|||
|
|
aa, ab := a[i*n+i], a[(i+1)*n+i]
|
|||
|
|
// The Givens denominator is the real hypot of the two
|
|||
|
|
// entries; a complex square root here would leave the
|
|||
|
|
// rotation non-unitary and the iteration divergent.
|
|||
|
|
den := math.Hypot(cmplx.Abs(aa), cmplx.Abs(ab))
|
|||
|
|
if den == 0 {
|
|||
|
|
rr[2*(i-l)] = 1
|
|||
|
|
rr[2*(i-l)+1] = 0
|
|||
|
|
continue
|
|||
|
|
}
|
|||
|
|
c := aa / complex(den, 0)
|
|||
|
|
s := ab / complex(den, 0)
|
|||
|
|
rr[2*(i-l)] = c
|
|||
|
|
rr[2*(i-l)+1] = s
|
|||
|
|
// Rows i, i+1 of the block columns l..n-1.
|
|||
|
|
for j := i; j < n; j++ {
|
|||
|
|
h1, h2 := a[i*n+j], a[(i+1)*n+j]
|
|||
|
|
a[i*n+j] = cmplx.Conj(c)*h1 + cmplx.Conj(s)*h2
|
|||
|
|
a[(i+1)*n+j] = -s*h1 + c*h2
|
|||
|
|
}
|
|||
|
|
a[(i+1)*n+i] = 0
|
|||
|
|
}
|
|||
|
|
// Multiply back by the rotations from the right: R·Gᴴ.
|
|||
|
|
for i := l; i < r; i++ {
|
|||
|
|
c, s := rr[2*(i-l)], rr[2*(i-l)+1]
|
|||
|
|
for j := 0; j <= min(i+1, n-1); j++ {
|
|||
|
|
h1, h2 := a[j*n+i], a[j*n+i+1]
|
|||
|
|
a[j*n+i] = h1*c + h2*s
|
|||
|
|
a[j*n+i+1] = -h1*cmplx.Conj(s) + h2*cmplx.Conj(c)
|
|||
|
|
}
|
|||
|
|
// Accumulate q <- q·Gᴴ.
|
|||
|
|
for row := range n {
|
|||
|
|
q1, q2 := q[row*n+i], q[row*n+i+1]
|
|||
|
|
q[row*n+i] = q1*c + q2*s
|
|||
|
|
q[row*n+i+1] = -q1*cmplx.Conj(s) + q2*cmplx.Conj(c)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
for i := l; i <= r; i++ {
|
|||
|
|
a[i*n+i] += mu
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
return base.Errf("the Schur iteration did not converge")
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// triSqrt overwrites the upper triangular t with its principal square
|
|||
|
|
// root by the recurrence s·s = t, diagonal first. A vanishing
|
|||
|
|
// s_ii + s_jj denominator means the principal root does not exist for
|
|||
|
|
// this spectrum and the sweep reports it.
|
|||
|
|
func triSqrt(t []complex128, n int) error {
|
|||
|
|
s := make([]complex128, n*n)
|
|||
|
|
for i := range n {
|
|||
|
|
s[i*n+i] = cmplx.Sqrt(t[i*n+i])
|
|||
|
|
}
|
|||
|
|
for j := 1; j < n; j++ {
|
|||
|
|
for i := j - 1; i >= 0; i-- {
|
|||
|
|
sum := t[i*n+j]
|
|||
|
|
for k := i + 1; k < j; k++ {
|
|||
|
|
sum -= s[i*n+k] * s[k*n+j]
|
|||
|
|
}
|
|||
|
|
den := s[i*n+i] + s[j*n+j]
|
|||
|
|
if den == 0 {
|
|||
|
|
return base.Errf("the principal square root does not exist: repeated zero eigenvalues on the diagonal")
|
|||
|
|
}
|
|||
|
|
s[i*n+j] = sum / den
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
copy(t, s)
|
|||
|
|
return nil
|
|||
|
|
}
|