161 lines
6.2 KiB
Go
161 lines
6.2 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||
// SPDX-License-Identifier: MIT
|
||
|
||
package linalg
|
||
|
||
import (
|
||
"math"
|
||
"slices"
|
||
|
||
"sourcedock.dev/petrbalvin/tensor/internal/base"
|
||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||
)
|
||
|
||
// Rank-one modification of the elimination-tree sparse Cholesky
|
||
// factor, the sparse counterpart of `CholeskyUpdate` and
|
||
// `CholeskyDowndate`. Refactorising from scratch reruns the whole
|
||
// ordering and elimination; a rank-one change of the matrix needs one
|
||
// sweep of (hyperbolic for the downdate) rotations down the stored
|
||
// columns, the sweep the dense contract applies to a full n×n factor.
|
||
//
|
||
// The honest contract of the sparse form is narrower than the dense
|
||
// one. A rank-one change of A generally changes the elimination
|
||
// pattern, so the updated factor generally does not live on the old
|
||
// one's fill; no rotation sweep can invent the missing entries. What
|
||
// these methods do is transform the stored NUMERIC factor on the
|
||
// EXISTING pattern: the sweep refuses, with a clear error, the moment
|
||
// the update would need an entry the pattern does not hold, and it
|
||
// leaves the factor untouched when it does. Where the pattern
|
||
// suffices, and it does for the bread-and-butter cases of a supported
|
||
// vector on a banded or otherwise already-rich pattern, the result is
|
||
// the exact factor of the modified matrix; otherwise the fallback is a
|
||
// fresh `NewSparseCholesky` of the explicitly modified matrix.
|
||
|
||
// Update applies the rank-one update A ← A + x·xᵀ to the factor in
|
||
// place: afterwards the factor is the Cholesky factor of the modified
|
||
// matrix, on the same pattern and in the same elimination order. x is
|
||
// a vector in the original coordinates, the coordinates `Solve` takes,
|
||
// and is gathered through the permutation internally.
|
||
//
|
||
// The sweep runs the dense rotation column by column, touching only
|
||
// the stored entries of each column. If the rotation would write a
|
||
// value into a position the pattern does not store, the update is
|
||
// refused and the factor keeps its original numbers, mirroring the
|
||
// dense contract where the input factor is left intact on failure; a
|
||
// fresh factorisation of the modified matrix is the fallback. A
|
||
// non-finite entry, a complex vector or a wrong length is refused
|
||
// outright.
|
||
func (f *SparseCholesky) Update(x *core.Array) error {
|
||
return f.cholRankOneUpdate(x, "SparseCholesky.Update", +1)
|
||
}
|
||
|
||
// Downdate applies the rank-one downdate A ← A − x·xᵀ to the factor in
|
||
// place, with Update's pattern contract. Each hyperbolic rotation
|
||
// squares the remaining diagonal against the working vector: when a
|
||
// diagonal can no longer dominate, the modified matrix has left the
|
||
// positive definite cone and the downdate is an error, exactly the
|
||
// dense refusal; when the square leaves the float64 range the downdate
|
||
// refuses rather than poisoning the factor.
|
||
func (f *SparseCholesky) Downdate(x *core.Array) error {
|
||
return f.cholRankOneUpdate(x, "SparseCholesky.Downdate", -1)
|
||
}
|
||
|
||
// cholRankOneUpdate carries the shared sweep on a copy of the numeric
|
||
// factor, committing only when every column has settled, so a refusal
|
||
// anywhere leaves the receiver untouched.
|
||
func (f *SparseCholesky) cholRankOneUpdate(x *core.Array, name string, sign int) error {
|
||
if x.Dtype() == core.Complex {
|
||
return base.Errf("%s: complex inputs are not supported", name)
|
||
}
|
||
if x.NDim() != 1 || x.Len() != f.n {
|
||
return base.Errf("%s: the vector must have length %d, got shape %s",
|
||
name, f.n, base.ShapeText(x.Shape()))
|
||
}
|
||
// The factor lives in the elimination order: the working vector is
|
||
// x gathered through the permutation, the same gather `Solve`
|
||
// applies to a right-hand side.
|
||
v := make([]float64, f.n)
|
||
for k := range f.n {
|
||
z := x.FloatAt(f.perm[k])
|
||
if math.IsNaN(z) || math.IsInf(z, 0) {
|
||
return base.Errf("%s: entry %d is not finite", name, f.perm[k])
|
||
}
|
||
v[k] = z
|
||
}
|
||
vals := slices.Clone(f.values)
|
||
diag := slices.Clone(f.diag)
|
||
// touched holds every position the sweep has seen non-zero, so the
|
||
// fill check scans the support of the working vector instead of all
|
||
// n positions of every rotated column. The support only grows along
|
||
// stored entries, which keeps the check proportional to what the
|
||
// sweep actually touches.
|
||
touched := make([]bool, f.n)
|
||
nz := make([]int, 0, f.n)
|
||
for k := range f.n {
|
||
if v[k] != 0 {
|
||
touched[k] = true
|
||
nz = append(nz, k)
|
||
}
|
||
}
|
||
for k := range f.n {
|
||
if v[k] == 0 {
|
||
// The rotation on a zero pivot is the identity, c = 1 and
|
||
// s = 0, needs no fill and writes nothing: the skip is
|
||
// bit-exact against the dense sweep.
|
||
continue
|
||
}
|
||
// The fill check. The rotation writes ±s·v[i] into every
|
||
// position (i, k) with v[i] ≠ 0, so one such position outside
|
||
// the stored column would create fill the pattern does not
|
||
// hold. The column's rows are sorted, so the membership test is
|
||
// a binary search.
|
||
for _, i := range nz {
|
||
if i <= k || v[i] == 0 {
|
||
continue
|
||
}
|
||
if _, ok := slices.BinarySearch(f.rowIdx[f.colStart[k]:f.colStart[k+1]], i); ok {
|
||
continue
|
||
}
|
||
return base.Errf("%s: the modification needs a factor entry at permuted position [%d,%d] (original [%d,%d]), which the stored pattern does not hold; factorise the modified matrix afresh instead",
|
||
name, i, k, f.perm[i], f.perm[k])
|
||
}
|
||
d := diag[k]
|
||
var r, c, s float64
|
||
if sign > 0 {
|
||
r = math.Hypot(d, v[k])
|
||
if !finiteF64(r) {
|
||
return base.Errf("%s: the update at column %d (original %d) overflows the float64 range", name, k, f.perm[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 base.Errf("%s: the downdate at column %d (original %d) overflows the float64 range", name, k, f.perm[k])
|
||
}
|
||
if r2 <= 0 {
|
||
return base.Errf("%s: the modified matrix is not positive definite at column %d (original %d)", name, k, f.perm[k])
|
||
}
|
||
r = math.Sqrt(r2)
|
||
c, s = d/r, v[k]/r
|
||
}
|
||
diag[k] = r
|
||
for p := f.colStart[k]; p < f.colStart[k+1]; p++ {
|
||
i := f.rowIdx[p]
|
||
lik, xi := vals[p], v[i]
|
||
if sign > 0 {
|
||
vals[p] = c*lik + s*xi
|
||
} else {
|
||
vals[p] = c*lik - s*xi
|
||
}
|
||
v[i] = c*xi - s*lik
|
||
if !touched[i] {
|
||
touched[i] = true
|
||
nz = append(nz, i)
|
||
}
|
||
}
|
||
}
|
||
f.values = vals
|
||
f.diag = diag
|
||
return nil
|
||
}
|