Files
tensor/linalg/cholupdate.go
T

132 lines
4.8 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 (
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
import "math"
// Rank-one modification of a Cholesky factorisation. Refactorising
// from scratch costs a full O(n³) pass; a rank-one change of the
// matrix only needs one sweep of rotations over the factor, O(n²),
// which is what repeated posterior re-evaluation, sliding-window
// covariances and active-set methods live on.
// CholeskyUpdate returns the lower Cholesky factor of A + x·xᵀ, given
// l, the lower factor of A. The factor is rebuilt by orthogonal
// rotations applied to l with x appended as an extra row, so the
// result is exact to rounding without a fresh factorisation. A rank-1
// factor, a mismatched vector or a complex input is an error.
func CholeskyUpdate(l, x *core.Array) (*core.Array, error) {
return cholRankOne(l, x, "CholeskyUpdate", +1)
}
// CholeskyDowndate returns the lower Cholesky factor of A − x·xᵀ,
// given l, the lower factor of A. The sweep runs hyperbolic rotations
// down the factor, each one squaring the remaining diagonal against
// the vector: when a diagonal can no longer dominate, the modified
// matrix has left the positive definite cone and the downdate is an
// error rather than a silently broken factor.
func CholeskyDowndate(l, x *core.Array) (*core.Array, error) {
return cholRankOne(l, x, "CholeskyDowndate", -1)
}
// cholRankOne carries the shared sweep. The augmented matrix
// [l; xᵀ] holds the updated Gram matrix for +1; for −1 the vector row
// enters with a minus sign and only the hyperbolic rotations that keep
// l·lᵀ − x·xᵀ invariant may remove it, which is where the positive
// definiteness test lives.
func cholRankOne(l, x *core.Array, name string, sign int) (*core.Array, error) {
if l.Dtype() == core.Complex || x.Dtype() == core.Complex {
return nil, base.Errf("%s: complex inputs are not supported", name)
}
if l.NDim() != 2 || l.Shape()[0] != l.Shape()[1] {
return nil, base.Errf("%s: the factor must be a square 2-D matrix, got shape %s",
name, base.ShapeText(l.Shape()))
}
n := l.Shape()[0]
if x.NDim() != 1 || x.Len() != n {
return nil, base.Errf("%s: the vector must have length %d, got shape %s",
name, n, base.ShapeText(x.Shape()))
}
out := core.New(core.Float, []int{n, n}...)
if l.Dtype() == core.Float && !l.Strided() && len(l.RawFloats()) == n*n {
copy(out.RawFloats(), l.RawFloats())
} else {
for i := range n * n {
out.RawFloats()[i] = l.FloatAt(i)
}
}
v := make([]float64, n)
if x.Dtype() == core.Float && !x.Strided() && len(x.RawFloats()) == n {
copy(v, x.RawFloats())
} else {
for i := range v {
v[i] = x.FloatAt(i)
}
}
// The rotations below assume a lower triangular factor; a nonzero
// strict upper triangle would silently corrupt the sweep, so it is
// refused up front, and so is a non-finite entry anywhere: the
// diagonal test cannot see a NaN, and the hyperbolic square of an
// Inf overflows mid-sweep. The sparse twin refuses the same input.
2026-09-03 10:00:00 +02:00
for i := range n {
for j := range n {
w := out.RawFloats()[i*n+j]
if math.IsNaN(w) || math.IsInf(w, 0) {
return nil, base.Errf("%s: factor entry [%d,%d] is not finite", name, i, j)
}
if j > i && w != 0 {
2026-09-03 10:00:00 +02:00
return nil, base.Errf("%s: the factor must be lower triangular (nonzero entry at row %d, column %d)", name, i, j)
}
}
}
for i := range n {
if math.IsNaN(v[i]) || math.IsInf(v[i], 0) {
return nil, base.Errf("%s: entry %d is not finite", name, i)
}
}
2026-09-03 10:00:00 +02:00
// The modified matrix reads [L, x]·J·[L, x]ᵀ with J = diag(I, ±1):
// one sweep of (hyperbolic for −1, orthogonal for +1) rotations
// eliminates the vector column, and the surviving columns are the
// new lower factor. Each rotation squares the running diagonal
// against the vector, which is where a downdate leaves the
// positive definite cone.
for k := range n {
d := out.RawFloats()[k*n+k]
if d <= 0 {
return nil, base.Errf("%s: the factor must be lower triangular with a positive diagonal", name)
}
var r, c, s float64
if sign > 0 {
r = math.Hypot(d, v[k])
c, s = d/r, v[k]/r
} else {
r2 := d*d - v[k]*v[k]
if math.IsNaN(r2) || math.IsInf(r2, 0) {
return nil, base.Errf("%s: the downdate at row %d overflows the float64 range", name, k)
}
if r2 <= 0 {
return nil, base.Errf("%s: the modified matrix is not positive definite at row %d", name, k)
}
r = math.Sqrt(r2)
c, s = d/r, v[k]/r
}
out.RawFloats()[k*n+k] = r
for i := k + 1; i < n; i++ {
lik, xi := out.RawFloats()[i*n+k], v[i]
if sign > 0 {
out.RawFloats()[i*n+k] = c*lik + s*xi
} else {
out.RawFloats()[i*n+k] = c*lik - s*xi
}
v[i] = c*xi - s*lik
}
}
return out, nil
}