121 lines
4.3 KiB
Go
121 lines
4.3 KiB
Go
// 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.
|
|||
|
|
for i := range n {
|
|||
|
|
for j := i + 1; j < n; j++ {
|
|||
|
|
if out.RawFloats()[i*n+j] != 0 {
|
|||
|
|
return nil, base.Errf("%s: the factor must be lower triangular (nonzero entry at row %d, column %d)", name, i, j)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
// 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
|
|||
|
|
}
|