191 lines
5.1 KiB
Go
191 lines
5.1 KiB
Go
// 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
|
|||
|
|
}
|