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