Files
tensor/linalg/least_squares_test.go
T

191 lines
5.1 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 (
"math"
"math/big"
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// The solve applies the Householder reflectors to the right-hand sides
// instead of forming Q, so its forward error is compared here against a
// minimiser solved exactly in rational arithmetic: the reference is not
// another floating-point route, it is the answer.
// exactMinimiser returns the minimiser of ||Ax − b|| for an integer A and
// b, solved from the normal equations in exact rational arithmetic, along
// with the exact minimal residual norm.
func exactMinimiser(t *testing.T, a [][]float64, b []float64) ([]*big.Rat, float64) {
t.Helper()
m, n := len(a), len(a[0])
g := make([][]*big.Rat, n)
for i := range n {
g[i] = make([]*big.Rat, n+1)
for j := range n + 1 {
g[i][j] = new(big.Rat)
}
}
exact := func(v float64) *big.Rat {
r := new(big.Rat).SetFloat64(v)
if r == nil {
t.Fatalf("value %g is not representable as a rational", v)
}
return r
}
for i := range n {
for j := range n {
s := new(big.Rat)
for r := range m {
s.Add(s, new(big.Rat).Mul(exact(a[r][i]), exact(a[r][j])))
}
g[i][j] = s
}
s := new(big.Rat)
for r := range m {
s.Add(s, new(big.Rat).Mul(exact(a[r][i]), exact(b[r])))
}
g[i][n] = s
}
for c := range n {
p := c
for r := c + 1; r < n; r++ {
if new(big.Rat).Abs(g[r][c]).Cmp(new(big.Rat).Abs(g[p][c])) > 0 {
p = r
}
}
g[c], g[p] = g[p], g[c]
for r := c + 1; r < n; r++ {
f := new(big.Rat).Quo(g[r][c], g[c][c])
for k := c; k <= n; k++ {
g[r][k].Sub(g[r][k], new(big.Rat).Mul(f, g[c][k]))
}
}
}
x := make([]*big.Rat, n)
for i := n - 1; i >= 0; i-- {
s := new(big.Rat).Set(g[i][n])
for j := i + 1; j < n; j++ {
s.Sub(s, new(big.Rat).Mul(g[i][j], x[j]))
}
x[i] = s.Quo(s, g[i][i])
}
res := new(big.Rat)
for r := range m {
s := new(big.Rat)
for j := range n {
s.Add(s, new(big.Rat).Mul(exact(a[r][j]), x[j]))
}
s.Sub(s, exact(b[r]))
res.Add(res, new(big.Rat).Mul(s, s))
}
rf, _ := new(big.Float).Sqrt(new(big.Float).SetRat(res)).Float64()
return x, rf
}
// leastSquaresFixture builds an integer A with the given column scaling
// and a right-hand side that is A·x0 plus a small integer perturbation, so
// the exact minimiser is known and the residual is not zero.
func leastSquaresFixture(m, n int, scaleJ func(int) float64) ([][]float64, []float64) {
a := make([][]float64, m)
s := uint64(12345)
for i := range m {
a[i] = make([]float64, n)
for j := range n {
s = s*6364136223846793005 + 1442695040888963407
a[i][j] = float64((s>>40)%9+1) * scaleJ(j)
}
}
x0 := make([]float64, n)
for j := range n {
x0[j] = float64((j*7)%5 - 2)
}
b := make([]float64, m)
for i := range m {
s = s*6364136223846793005 + 1442695040888963407
v := 0.0
for j := range n {
v += a[i][j] * x0[j]
}
b[i] = v + float64(int((s>>50)%3)-1)
}
return a, b
}
func TestLeastSquaresAgainstRationalMinimiser(t *testing.T) {
cases := []struct {
name string
m, n int
scal func(int) float64
}{
{"512x32", 512, 32, func(int) float64 { return 1 }},
{"1024x48", 1024, 48, func(int) float64 { return 1 }},
{"512x32-scaled", 512, 32, func(j int) float64 { return math.Pow(0.5, float64(j)) }},
}
for _, c := range cases {
t.Run(c.name, func(t *testing.T) {
a, b := leastSquaresFixture(c.m, c.n, c.scal)
xr, resExact := exactMinimiser(t, a, b)
am, err := core.FromFloats(flattenRows(a), c.m, c.n)
if err != nil {
t.Fatal(err)
}
bm, err := core.FromFloats(b, c.m, 1)
if err != nil {
t.Fatal(err)
}
sol, err := LeastSquares(am, bm)
if err != nil {
t.Fatalf("LeastSquares: %v", err)
}
x := sol.RawFloats()
worst, norm := 0.0, 0.0
for i := range c.n {
ref, _ := xr[i].Float64()
norm = math.Max(norm, math.Abs(ref))
worst = math.Max(worst, math.Abs(x[i]-ref))
}
rel := worst / norm
// The housekeeping floor: a backward-stable solve of a problem
// this well conditioned lands within a few ulps per element,
// and the scaled case carries a condition number near 2^31, so
// the bound is the conditioning times the unit roundoff with a
// generous constant.
kond := 1.0
if c.scal(1) != 1 {
kond = math.Pow(2, float64(c.n))
}
bound := 64 * math.Max(1, kond) * 2.220446049250313e-16
if rel > bound {
t.Errorf("relative solution error %g exceeds the bound %g", rel, bound)
}
// The residual must not be worse than the exact minimal one by
// more than the same order.
res := 0.0
for i := range c.m {
v := b[i]
for j := range c.n {
v -= a[i][j] * x[j]
}
res += v * v
}
res = math.Sqrt(res)
if res > resExact*(1+1e-8) {
t.Errorf("residual %g exceeds the exact minimal residual %g", res, resExact)
}
t.Logf("relative solution error %.3e, residual %.6g (exact minimal %.6g)", rel, res, resExact)
})
}
}
func flattenRows(a [][]float64) []float64 {
out := make([]float64, 0, len(a)*len(a[0]))
for _, row := range a {
out = append(out, row...)
}
return out
}