Files
tensor/linalg/svdsolve.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

139 lines
4.6 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 (
"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
}