Files
tensor/stats/gmm_posterior_reference_test.go
T

177 lines
6.3 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 stats
import (
"math"
"math/big"
"testing"
)
// The expectation step's posterior normalisation against an exact
// referent. The fit computes each row's responsibilities from the log
// densities and their row maximum; the quotients it publishes must
// match the posteriories a 256-bit evaluation of the same numbers
// produces, and match them at least as closely as the form that
// re-exponentiates through the row normaliser. The log densities here
// come from a closed-form bivariate evaluator, never from the mixture
// code, so the referent shares no arithmetic with the implementation.
// bigLn2 is the natural logarithm of two to a hundred places, the
// reduction constant of the reference exponential. The head the series
// accuracy needs is far shorter than the tail quoted.
const bigLn2 = "0.6931471805599453094172321214581765680755001343602552541206800094933936219696947156058633269964186875"
// referenceExp evaluates eˣ at prec mantissa bits: the argument folds
// to x = k·ln 2 + r with |r| ≤ ln 2, the series Σ rⁿ/n! converges on
// that range in under a hundred terms, and the shift multiplies by 2^k.
func referenceExp(x float64, prec uint) *big.Float {
p := prec + 64
ln2, ok := new(big.Float).SetPrec(p).SetString(bigLn2)
if !ok {
panic("referenceExp: the reduction constant failed to parse")
}
xf := new(big.Float).SetPrec(p).SetFloat64(x)
k := int(math.Round(x / math.Ln2))
r := new(big.Float).SetPrec(p).Mul(ln2, new(big.Float).SetPrec(p).SetInt64(int64(k)))
r.Sub(xf, r)
sum := new(big.Float).SetPrec(p).SetInt64(1)
term := new(big.Float).SetPrec(p).SetInt64(1)
for n := 1; n < 1000; n++ {
term.Mul(term, r)
term.Quo(term, new(big.Float).SetPrec(p).SetUint64(uint64(n)))
sum.Add(sum, term)
if e := term.MantExp(nil); term.Sign() == 0 || e < -int(p)-4 {
pow2 := new(big.Float).SetPrec(p)
pow2.SetMantExp(new(big.Float).SetPrec(p).SetInt64(1), k)
return sum.Mul(sum, pow2)
}
}
panic("referenceExp: the series failed to converge")
}
// bivariateLogDensity evaluates the normal log density through the
// closed-form inverse and determinant of the 2×2 covariance, the
// evaluator that shares nothing with the fit's forward solve.
func bivariateLogDensity(x, y float64, mean [2]float64, cov [4]float64, weight float64) float64 {
det := cov[0]*cov[3] - cov[1]*cov[2]
dx, dy := x-mean[0], y-mean[1]
inv00, inv01, inv11 := cov[3]/det, -cov[1]/det, cov[0]/det
quad := dx*dx*inv00 + 2*dx*dy*inv01 + dy*dy*inv11
return math.Log(weight) - 0.5*(2*math.Log(2*math.Pi)+math.Log(det)) - 0.5*quad
}
func TestGaussianMixturePosteriorReferenceForms(t *testing.T) {
// The reference exponential agrees with the hardware one wherever
// the hardware answer is representable: this pins the reduction
// constant and the series against an independent evaluator before
// anything else is claimed on top of them.
for _, x := range []float64{-50, -7.3, -1, -0.01, 0, 0.5, 3.7, 40} {
want := math.Exp(x)
got, _ := referenceExp(x, 256).Float64()
if math.Abs(got-want) > 1e-13*math.Max(1, math.Abs(want)) {
t.Fatalf("referenceExp(%g) = %g, math.Exp gives %g", x, got, want)
}
}
weights := []float64{0.5, 0.3, 0.2}
means := [][2]float64{{-2, 0}, {1, 1}, {3, -1}}
covs := [][4]float64{
{1, 0.3, 0.3, 0.8},
{0.5, 0, 0, 1.2},
{2, -0.5, -0.5, 0.7},
}
points := [][2]float64{
{-2.1, 0.05}, // near the first component
{0.9, 1.2}, // near the second
{3.2, -0.8}, // near the third
{0.2, 0.4}, // between the first two
{-1, 0.8}, // a broad mixture row
{5.5, -2.5}, // moderately outside every component
{-6.5, 2.5}, // the same, opposite side
{18, 24}, // far tail: every density microscopic
}
const k = 3
worstNew, worstOld := 0.0, 0.0
for _, pt := range points {
// The log densities the step would see, from the independent
// closed form.
lps := make([]float64, k)
for c := range k {
lps[c] = bivariateLogDensity(pt[0], pt[1], means[c], covs[c], weights[c])
}
rowMax := math.Inf(-1)
for _, lp := range lps {
if lp > rowMax {
rowMax = lp
}
}
// The exact posteriories of these very numbers: exponentials
// and one quotient, all at 256 bits.
exacts := make([]*big.Float, k)
exactResp := make([]float64, k)
totalExact := new(big.Float).SetPrec(256)
for c := range lps {
exacts[c] = referenceExp(lps[c], 256)
totalExact.Add(totalExact, exacts[c])
}
for c := range lps {
exactResp[c], _ = new(big.Float).SetPrec(256).Quo(exacts[c], totalExact).Float64()
}
// The published forms on the same numbers: the quotient of the
// stored exponentials, and the retired form that re-exponentiates
// through the row normaliser.
expSum := 0.0
expVals := make([]float64, k)
for c := range lps {
expVals[c] = math.Exp(lps[c] - rowMax)
expSum += expVals[c]
}
logNorm := rowMax + math.Log(expSum)
// The far tail underflows every form to zero or near it, where
// a relative comparison carries no information; the row still
// participates through the sum-to-one pin.
live := false
for c := range lps {
if exactResp[c] > 1e-200 {
live = true
}
}
rowNew, rowOld := 0.0, 0.0
sum := 0.0
for c := range lps {
newResp := expVals[c] / expSum
oldResp := math.Exp(lps[c] - logNorm)
sum += newResp
if !live {
continue
}
exact := exactResp[c]
if exact <= 1e-200 {
continue
}
en := math.Abs(newResp-exact) / exact
eo := math.Abs(oldResp-exact) / exact
rowNew = math.Max(rowNew, en)
rowOld = math.Max(rowOld, eo)
}
worstNew = math.Max(worstNew, rowNew)
worstOld = math.Max(worstOld, rowOld)
if live && math.Abs(sum-1) > 1e-14 {
t.Fatalf("point %v: the quotients sum to %.17g, want 1", pt, sum)
}
if live && rowNew > 1e-12 {
t.Fatalf("point %v: the quotient form's worst relative error is %g, want under 1e-12", pt, rowNew)
}
if live && rowOld > 1e-12 {
t.Fatalf("point %v: the re-exponentiated form's worst relative error is %g, want under 1e-12", pt, rowOld)
}
}
if worstNew > worstOld {
t.Fatalf("the quotient form is the less accurate one at the fixture: worst relative error %g against the re-exponentiated form's %g", worstNew, worstOld)
}
t.Logf("worst relative error against the 256-bit referent: quotient form %g, re-exponentiated form %g", worstNew, worstOld)
}