Files
tensor/stats/gmm_posterior_reference_test.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

177 lines
6.3 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 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)
}