Files
tensor/linalg/schur.go
T
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

264 lines
7.9 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// 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
}