Files
tensor/linalg/schur.go
T

264 lines
7.9 KiB
Go
Raw 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"
)
// 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
}