139 lines
4.6 KiB
Go
139 lines
4.6 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"
|
||
)
|
||
|
||
// Rank-deficient and ill-posed systems through the singular value
|
||
// decomposition. Where the QR-based LeastSquares fails outright on a
|
||
// rank-deficient matrix, the SVD route scales each singular
|
||
// direction's contribution individually: truncation zeroes the
|
||
// directions below a chosen rank, Tikhonov damping shrinks every
|
||
// direction by σ/(σ²+λ). Both answer x = V·W·Uᵀb over the thin
|
||
// factorisation, the minimum-norm least-squares solution of the
|
||
// system they solve.
|
||
|
||
// SolveTruncated solves A·x = b keeping only the rank largest
|
||
// singular values, the truncated pseudoinverse: directions beyond the
|
||
// rank, whatever noise they carry, contribute nothing. a is m×n and b
|
||
// is m or m×k; a singular value that vanishes to round-off below the
|
||
// requested rank is an error, since the system is more deficient than
|
||
// the truncation asked for.
|
||
func SolveTruncated(a, b *core.Array, rank int) (*core.Array, error) {
|
||
return svdSolve(a, b, "SolveTruncated", func(sigma []float64) ([]float64, error) {
|
||
if rank < 1 || rank > len(sigma) {
|
||
return nil, base.Errf("SolveTruncated: rank must be between 1 and %d, got %d", len(sigma), rank)
|
||
}
|
||
floor := 64 * base.EpsF * sigma[0]
|
||
w := make([]float64, len(sigma))
|
||
for i, s := range sigma {
|
||
if i >= rank {
|
||
break
|
||
}
|
||
if s <= floor {
|
||
return nil, base.Errf("SolveTruncated: the system is rank-deficient below rank %d (σ%d = %g)",
|
||
rank, i, s)
|
||
}
|
||
w[i] = 1 / s
|
||
}
|
||
return w, nil
|
||
})
|
||
}
|
||
|
||
// SolveTikhonov solves the regularised least-squares problem
|
||
// min ‖A·x − b‖² + λ‖x‖² by Tikhonov damping with the identity
|
||
// prior: every singular direction is scaled by σ/(σ²+λ), which keeps
|
||
// ill-conditioned directions from amplifying whatever noise sits on
|
||
// b while biasing the answer only where the data carries no
|
||
// information. a is m×n and b is m or m×k; a non-positive λ is an
|
||
// error.
|
||
func SolveTikhonov(a, b *core.Array, lambda float64) (*core.Array, error) {
|
||
return svdSolve(a, b, "SolveTikhonov", func(sigma []float64) ([]float64, error) {
|
||
if lambda <= 0 {
|
||
return nil, base.Errf("SolveTikhonov: lambda must be positive, got %g", lambda)
|
||
}
|
||
w := make([]float64, len(sigma))
|
||
for i, s := range sigma {
|
||
w[i] = s / (s*s + lambda)
|
||
}
|
||
return w, nil
|
||
})
|
||
}
|
||
|
||
// svdSolve applies x = V·W·Uᵀb over the thin SVD of a, with W the
|
||
// diagonal of per-direction weights the caller chooses.
|
||
func svdSolve(a, b *core.Array, name string, weight func(sigma []float64) ([]float64, error)) (*core.Array, error) {
|
||
if a.Dtype() == core.Complex {
|
||
return nil, base.Errf("%s: complex matrices are not supported", name)
|
||
}
|
||
// The float fill of b below is the same gap the a gate closes: a
|
||
// complex right-hand side has no float payload to read.
|
||
if b.Dtype() == core.Complex {
|
||
return nil, base.Errf("%s: complex right-hand sides are not supported", name)
|
||
}
|
||
if a.NDim() != 2 || a.Shape()[0] == 0 || a.Shape()[1] == 0 {
|
||
return nil, base.Errf("%s: a must be a non-empty 2-D matrix, got shape %s", name, base.ShapeText(a.Shape()))
|
||
}
|
||
if b.NDim() != 1 && b.NDim() != 2 {
|
||
return nil, base.Errf("%s: b must be 1-D or 2-D, got shape %s", name, base.ShapeText(b.Shape()))
|
||
}
|
||
m, n := a.Shape()[0], a.Shape()[1]
|
||
if b.Shape()[0] != m {
|
||
return nil, base.Errf("%s: b rows (%d) must match a rows (%d)", name, b.Shape()[0], m)
|
||
}
|
||
cols := 1
|
||
if b.NDim() == 2 {
|
||
cols = b.Shape()[1]
|
||
}
|
||
r := min(n, m)
|
||
u, sigma, vt, err := SVD(a)
|
||
if err != nil {
|
||
return nil, base.Errf("%s: %w", name, err)
|
||
}
|
||
sVals := make([]float64, r)
|
||
for i := range r {
|
||
sVals[i] = sigma.RawFloats()[i] // SVD returns a dense float array
|
||
}
|
||
w, err := weight(sVals)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
// c = Uᵀb (r×cols), damped row-wise, then x = Vᵀ·d.
|
||
bFlat := make([]float64, m*cols)
|
||
if b.Dtype() == core.Float && !b.Strided() && len(b.RawFloats()) == m*cols {
|
||
copy(bFlat, b.RawFloats())
|
||
} else {
|
||
for i := range m {
|
||
for j := range cols {
|
||
bFlat[i*cols+j] = b.FloatAt(i*cols + j)
|
||
}
|
||
}
|
||
}
|
||
uFlat := denseFloats(u, m, r)
|
||
vtFlat := denseFloats(vt, r, n)
|
||
x := make([]float64, n*cols)
|
||
for i := range r {
|
||
if w[i] == 0 {
|
||
continue
|
||
}
|
||
for j := range cols {
|
||
c := 0.0
|
||
for l := range m {
|
||
c += uFlat[l*r+i] * bFlat[l*cols+j]
|
||
}
|
||
c *= w[i]
|
||
for jj := range n {
|
||
x[jj*cols+j] += vtFlat[i*n+jj] * c
|
||
}
|
||
}
|
||
}
|
||
if b.NDim() == 1 {
|
||
return floatsToArray(x, []int{n}), nil
|
||
}
|
||
return floatsToArray(x, []int{n, cols}), nil
|
||
}
|