// Copyright (c) 2026 Petr Balvín (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 }