feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -0,0 +1,176 @@
|
||||
// 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)
|
||||
}
|
||||
Reference in New Issue
Block a user