feat: initial release
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s

Assisted-by: GLM 5.3 Flash
This commit is contained in:
2026-09-03 10:00:00 +02:00
commit af4ee19703
617 changed files with 191195 additions and 0 deletions
+436
View File
@@ -0,0 +1,436 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package stats
import (
"cmp"
"math"
"slices"
"strings"
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// TestKMeansRecoversBlobCentres draws two well-separated Gaussian
// blobs over the house generator and requires the fit to recover both
// centres to tolerance and to partition every sample onto its own
// blob's label.
func TestKMeansRecoversBlobCentres(t *testing.T) {
g := core.NewGenerator(11)
const perBlob = 80
sample := core.New(core.Float, 2*perBlob, 2)
truth := [2][2]float64{{0, 0}, {8, 8}}
for i := range 2 * perBlob {
blob := 0
if i >= perBlob {
blob = 1
}
for j := range 2 {
sample.RawFloats()[i*2+j] = truth[blob][j] + 0.6*g.NormalUnit()
}
}
res, err := KMeans(core.NewGenerator(3), sample, 2)
if err != nil {
t.Fatalf("KMeans: %v", err)
}
if !res.Converged {
t.Fatalf("the fit did not converge in %d sweeps", res.Iterations)
}
// Each true centre must be matched by a fitted one within the
// sampling noise of the blob mean: 0.6/sqrt(80) is about 0.07, so
// 0.4 leaves eight standard errors of room.
for _, want := range truth {
best := math.Inf(1)
for _, c := range res.Centres {
if d := sqDistance(c, want[:]); d < best {
best = d
}
}
if best > 0.4*0.4 {
t.Fatalf("no fitted centre within 0.4 of (%g, %g): closest %.4g",
want[0], want[1], math.Sqrt(best))
}
}
// The labels partition the blobs: every sample of a blob carries
// one label, and the two blobs carry different labels.
first, second := res.Labels[0], res.Labels[perBlob]
if first == second {
t.Fatalf("both blobs share label %d", first)
}
for i := range perBlob {
if res.Labels[i] != first {
t.Fatalf("blob 0 sample %d carries label %d, want %d", i, res.Labels[i], first)
}
if res.Labels[perBlob+i] != second {
t.Fatalf("blob 1 sample %d carries label %d, want %d", perBlob+i, res.Labels[perBlob+i], second)
}
}
}
// TestKMeansSeedingBeatsFixedBadSeeding pins the value of k-means++
// seeding as a property of the seeding, not a flaky inequality. The
// fixture is three one-dimensional blobs at 0, 30 and 60 with width
// 0.4 and k = 3, so the only good fit plants one centre per blob and
// costs about 19 units of inertia, while a seeding that plants all
// three centres inside blob 0 converges to a stable bad optimum: the
// leftmost centres stay to split blob 0, the rightmost is dragged to
// 45, midway between the two unclaimed blobs, and both pay 40 points
// times 15² each, about 18,000. The bad optimum is stable because the
// boundary between the shared centre at 45 and the blob-0 centres
// falls near 22, and no point of blob 1 crosses it. k-means++ cannot
// share that fate by anything the geometry allows: after the first
// (uniform) centre, every later centre is drawn with probability
// proportional to squared distance from the nearest one already
// chosen, and every point of an unclaimed blob outweighs every point
// of a claimed one by roughly (15/0.8)² ≳ 350, so keeping two centres
// in one blob is odds of millions to one against per run. Over the ten
// seeds below the outcome is deterministic, and every run must land
// the small objective.
func TestKMeansSeedingBeatsFixedBadSeeding(t *testing.T) {
g := core.NewGenerator(23)
const perBlob = 40
data := make([]float64, 3*perBlob)
for i := range perBlob {
data[i] = 0.4 * g.NormalUnit()
data[perBlob+i] = 30 + 0.4*g.NormalUnit()
data[2*perBlob+i] = 60 + 0.4*g.NormalUnit()
}
// The bad seeding: all three centres inside blob 0.
bad, err := kmeansLloyd(data, 3*perBlob, 1, 3, [][]float64{{-0.4}, {0}, {0.4}})
if err != nil {
t.Fatalf("bad-seeded fit: %v", err)
}
if bad.Inertia < 1000 {
t.Fatalf("the bad seeding escaped its own trap: inertia %.4g", bad.Inertia)
}
for seed := int64(101); seed < 111; seed++ {
res, err := KMeans(core.NewGenerator(seed), mustFloats(t, data, 3*perBlob, 1), 3)
if err != nil {
t.Fatalf("KMeans seed %d: %v", seed, err)
}
if res.Inertia > bad.Inertia/100 {
t.Fatalf("seed %d: k-means++ inertia %.4g is within a factor of 100 of the bad seeding's %.4g",
seed, res.Inertia, bad.Inertia)
}
}
}
// TestKMeansDeterministicUnderSeed runs the same sample twice under
// the same seed and requires bit-identical fits, then checks that the
// validation refusals all refuse.
func TestKMeansDeterministicUnderSeedAndValidation(t *testing.T) {
g := core.NewGenerator(9)
sample := core.New(core.Float, 30, 2)
for i := range 30 {
sample.RawFloats()[i*2] = g.Unit()
sample.RawFloats()[i*2+1] = g.Unit()
}
first, err := KMeans(core.NewGenerator(4), sample, 3)
if err != nil {
t.Fatalf("KMeans: %v", err)
}
second, err := KMeans(core.NewGenerator(4), sample, 3)
if err != nil {
t.Fatalf("KMeans repeat: %v", err)
}
for c, centre := range first.Centres {
if centre[0] != second.Centres[c][0] || centre[1] != second.Centres[c][1] {
t.Fatalf("seeded fit moved between runs: centre %d", c)
}
}
if first.Inertia != second.Inertia {
t.Fatalf("inertia moved between runs: %.17g against %.17g", first.Inertia, second.Inertia)
}
// An int sample takes the widening accessor path.
ints, err := core.FromInts([]int64{0, 0, 10, 12, 20, 18}, 3, 2)
if err != nil {
t.Fatalf("FromInts: %v", err)
}
if _, err := KMeans(core.NewGenerator(1), ints, 2); err != nil {
t.Fatalf("KMeans over an int sample: %v", err)
}
// Refusals: nil generator, wrong rank, complex, non-finite, and a
// k outside the sample.
if _, err := KMeans(nil, sample, 2); err == nil || !strings.Contains(err.Error(), "generator is nil") {
t.Fatalf("nil generator: got %v, want the nil-generator refusal", err)
}
if _, err := KMeans(core.NewGenerator(1), core.New(core.Float, 10), 2); err == nil || !strings.Contains(err.Error(), "must be rank 2") {
t.Fatalf("rank-1 sample: got %v, want the rank refusal", err)
}
if _, err := KMeans(core.NewGenerator(1), core.New(core.Complex, 4, 1), 2); err == nil || !strings.Contains(err.Error(), "complex") {
t.Fatalf("complex sample: got %v, want the complex refusal", err)
}
cloud := core.New(core.Float, 6, 2)
cloud.RawFloats()[3] = math.NaN()
if _, err := KMeans(core.NewGenerator(1), cloud, 2); err == nil || !strings.Contains(err.Error(), "non-finite") {
t.Fatalf("non-finite sample: got %v, want the non-finite refusal", err)
}
fine := core.New(core.Float, 6, 2)
if _, err := KMeans(core.NewGenerator(1), fine, 0); err == nil || !strings.Contains(err.Error(), "at least 1") {
t.Fatalf("k = 0: got %v, want the k floor refusal", err)
}
if _, err := KMeans(core.NewGenerator(1), fine, 7); err == nil || !strings.Contains(err.Error(), "exceeds") {
t.Fatalf("k above the sample size: got %v, want the exceeds refusal", err)
}
if _, err := KMeans(core.NewGenerator(1), core.New(core.Float, 6, 0), 2); err == nil || !strings.Contains(err.Error(), "one column") {
t.Fatalf("a zero-column sample: got %v, want the column refusal", err)
}
}
// TestKMeansEmptyClusterRule drives the documented empty-cluster
// handling directly: a seeding that leaves one centre nobody's nearest
// must see that centre move onto the sample farthest from it, and a
// sample with fewer distinct points than k must be refused rather than
// silently merged.
func TestKMeansEmptyClusterRule(t *testing.T) {
// Points 0, 1, 2 with centres 5 and 4.9: every point is nearer
// 4.9, so the centre at 5 strands and takes the farthest sample,
// 0, while the other centre keeps {1, 2} and settles on their
// mean. The settled fit is the split {0} and {1, 2}.
res, err := kmeansLloyd([]float64{0, 1, 2}, 3, 1, 2, [][]float64{{5}, {4.9}})
if err != nil {
t.Fatalf("kmeansLloyd: %v", err)
}
if !res.Converged {
t.Fatal("the empty-cluster fit did not converge")
}
if res.Labels[0] == res.Labels[1] || res.Labels[2] == res.Labels[0] {
t.Fatalf("labels %v do not split {0} from {1, 2}", res.Labels)
}
want := [][]float64{{0}, {1.5}}
for c := range 2 {
if res.Centres[c][0] != want[c][0] {
t.Fatalf("centre %d = %.4g, want %.4g", c, res.Centres[c][0], want[c][0])
}
}
// Three coincident points and k = 2: no repopulation is possible,
// and the refusal says so.
if _, err := KMeans(core.NewGenerator(1), mustFloats(t, []float64{3, 3, 3}, 3, 1), 2); err == nil {
t.Fatal("a sample with fewer distinct points than k was accepted")
}
}
// TestGaussianMixtureRecoversComponents draws a long sample from a
// two-component univariate mixture of known weights, means and
// variances and requires the EM fit to recover all three, matched in
// mean order.
func TestGaussianMixtureRecoversComponents(t *testing.T) {
g := core.NewGenerator(5)
const n = 6000
truth := []struct {
weight float64
mean float64
variance float64
}{{0.3, -2, 0.25}, {0.7, 3, 2.25}}
sample := core.New(core.Float, n, 1)
for i := range n {
c := 0
if g.Unit() >= truth[0].weight {
c = 1
}
sample.RawFloats()[i] = truth[c].mean + math.Sqrt(truth[c].variance)*g.NormalUnit()
}
res, err := GaussianMixture(core.NewGenerator(5), sample, 2)
if err != nil {
t.Fatalf("GaussianMixture: %v", err)
}
if !res.Converged {
t.Fatalf("EM did not converge in %d sweeps", res.Iterations)
}
if res.Components != 2 {
t.Fatalf("Components = %d, want 2", res.Components)
}
// The weights must form a distribution and the responsibilities
// must be posteriors.
total := 0.0
for _, w := range res.Weights {
total += w
}
if math.Abs(total-1) > 1e-9 {
t.Fatalf("weights sum to %.17g", total)
}
for i := range n {
rowSum := 0.0
for c := range 2 {
rowSum += res.Responsibilities[i*2+c]
}
if math.Abs(rowSum-1) > 1e-9 {
t.Fatalf("responsibilities of sample %d sum to %.17g", i, rowSum)
}
}
// Match components by mean order, the labelling that survives the
// mixture's label symmetry.
order := []int{0, 1}
slices.SortFunc(order, func(a, b int) int {
return cmp.Compare(res.Means[a][0], res.Means[b][0])
})
for pos, want := range truth {
c := order[pos]
if math.Abs(res.Weights[c]-want.weight) > 0.05 {
t.Fatalf("component %d: weight %.4g, want about %.2f", c, res.Weights[c], want.weight)
}
if math.Abs(res.Means[c][0]-want.mean) > 0.1 {
t.Fatalf("component %d: mean %.4g, want about %.2f", c, res.Means[c][0], want.mean)
}
got := res.Covariances[c][0]
if math.Abs(got-want.variance) > 0.15 {
t.Fatalf("component %d: variance %.4g, want about %.2f", c, got, want.variance)
}
}
}
// TestGaussianMixtureSeparatedResponsibilities fits a mixture whose
// components sit seven standard deviations apart and requires every
// posterior to be 0/1 to tolerance: a sample drawn from one component
// is orders of magnitude more likely under it than under the other.
func TestGaussianMixtureSeparatedResponsibilities(t *testing.T) {
g := core.NewGenerator(17)
const perComponent = 400
sample := core.New(core.Float, 2*perComponent, 1)
for i := range perComponent {
sample.RawFloats()[i] = -7 + g.NormalUnit()
sample.RawFloats()[perComponent+i] = 7 + g.NormalUnit()
}
res, err := GaussianMixture(core.NewGenerator(8), sample, 2)
if err != nil {
t.Fatalf("GaussianMixture: %v", err)
}
// Match the fitted components to the drawn sides by mean order.
low := 0
if res.Means[1][0] < res.Means[0][0] {
low = 1
}
for i := range 2 * perComponent {
want := 0.0
if i < perComponent {
want = 1
}
got := res.Responsibilities[i*2+low]
if math.Abs(got-want) > 1e-6 {
t.Fatalf("sample %d: responsibility %.3g, want %.0f", i, got, want)
}
}
}
// TestGaussianMixtureBICSelectsTrueCount draws from a two-component
// mixture and requires the BIC sweep over one to four components to
// put the minimum at the true count.
func TestGaussianMixtureBICSelectsTrueCount(t *testing.T) {
g := core.NewGenerator(29)
const n = 1200
sample := core.New(core.Float, n, 1)
for i := range n {
mean := -3.0
if i >= n/2 {
mean = 3
}
sample.RawFloats()[i] = mean + g.NormalUnit()
}
res, err := GaussianMixtureBIC(core.NewGenerator(6), sample, 4)
if err != nil {
t.Fatalf("GaussianMixtureBIC: %v", err)
}
if len(res.BICGrid) != 4 {
t.Fatalf("BIC grid holds %d entries, want 4", len(res.BICGrid))
}
if res.Components != 2 {
t.Fatalf("BIC selected %d components, want 2 (grid %v)", res.Components, res.BICGrid)
}
best := slices.Min(res.BICGrid)
if math.Abs(best-res.BIC) > 1e-9 {
t.Fatalf("the reported BIC %.4g is not the grid minimum %.4g", res.BIC, best)
}
}
// TestGaussianMixtureValidation checks the refusals of both entry
// points, and that the same seed reproduces the same fit bit for bit.
func TestGaussianMixtureValidation(t *testing.T) {
g := core.NewGenerator(2)
sample := core.New(core.Float, 40, 2)
for i := range 40 {
sample.RawFloats()[i*2] = g.Unit()
sample.RawFloats()[i*2+1] = g.Unit()
}
if _, err := GaussianMixture(nil, sample, 2); err == nil || !strings.Contains(err.Error(), "generator is nil") {
t.Fatalf("nil generator: got %v, want the nil-generator refusal", err)
}
if _, err := GaussianMixture(core.NewGenerator(1), core.New(core.Float, 40), 2); err == nil || !strings.Contains(err.Error(), "must be rank 2") {
t.Fatalf("rank-1 sample: got %v, want the rank refusal", err)
}
if _, err := GaussianMixture(core.NewGenerator(1), core.New(core.Float, 1, 2), 1); err == nil || !strings.Contains(err.Error(), "at least two samples") {
t.Fatalf("a one-sample fit: got %v, want the sample floor refusal", err)
}
if _, err := GaussianMixture(core.NewGenerator(1), sample, 0); err == nil || !strings.Contains(err.Error(), "component count") {
t.Fatalf("zero components: got %v, want the component-count refusal", err)
}
if _, err := GaussianMixture(core.NewGenerator(1), sample, 41); err == nil || !strings.Contains(err.Error(), "exceed") {
t.Fatalf("components above the sample size: got %v, want the exceeds refusal", err)
}
sick := core.New(core.Float, 4, 2)
sick.RawFloats()[2] = math.Inf(1)
if _, err := GaussianMixture(core.NewGenerator(1), sick, 1); err == nil || !strings.Contains(err.Error(), "non-finite") {
t.Fatalf("non-finite sample: got %v, want the non-finite refusal", err)
}
if _, err := GaussianMixture(core.NewGenerator(1), core.New(core.Complex, 4, 1), 1); err == nil || !strings.Contains(err.Error(), "complex") {
t.Fatalf("complex sample: got %v, want the complex refusal", err)
}
// A degenerate sample with no spread at all has no Gaussian to
// fit: the covariance collapses and the factorisation refuses it
// by name, in both entry points.
if _, err := GaussianMixture(core.NewGenerator(1), mustFloats(t, []float64{2, 2}, 2, 1), 1); err == nil || !strings.Contains(err.Error(), "positive definite") {
t.Fatalf("a zero-variance sample: got %v, want the factorisation refusal", err)
}
if _, err := GaussianMixtureBIC(core.NewGenerator(1), mustFloats(t, []float64{2, 2}, 2, 1), 2); err == nil || !strings.Contains(err.Error(), "positive definite") {
t.Fatalf("BIC: a zero-variance sample: got %v, want the factorisation refusal", err)
}
if _, err := GaussianMixtureBIC(nil, sample, 3); err == nil || !strings.Contains(err.Error(), "generator is nil") {
t.Fatalf("BIC: nil generator: got %v, want the nil-generator refusal", err)
}
if _, err := GaussianMixtureBIC(core.NewGenerator(1), sample, 0); err == nil || !strings.Contains(err.Error(), "largest component count") {
t.Fatalf("BIC: empty grid: got %v, want the component-count refusal", err)
}
if _, err := GaussianMixtureBIC(core.NewGenerator(1), core.New(core.Float, 40), 3); err == nil || !strings.Contains(err.Error(), "must be rank 2") {
t.Fatalf("BIC: rank-1 sample: got %v, want the rank refusal", err)
}
if _, err := GaussianMixtureBIC(core.NewGenerator(1), core.New(core.Float, 1, 2), 3); err == nil || !strings.Contains(err.Error(), "at least two samples") {
t.Fatalf("BIC: a one-sample fit: got %v, want the sample floor refusal", err)
}
if _, err := GaussianMixtureBIC(core.NewGenerator(1), sample, 41); err == nil || !strings.Contains(err.Error(), "exceed") {
t.Fatalf("BIC: grid above the sample size: got %v, want the exceeds refusal", err)
}
// Determinism: the same seed must reproduce the same likelihood
// and the same weights.
a, err := GaussianMixture(core.NewGenerator(12), sample, 2)
if err != nil {
t.Fatalf("GaussianMixture: %v", err)
}
b, err := GaussianMixture(core.NewGenerator(12), sample, 2)
if err != nil {
t.Fatalf("GaussianMixture repeat: %v", err)
}
if a.LogLikelihood != b.LogLikelihood {
t.Fatalf("seeded fit moved: %.17g against %.17g", a.LogLikelihood, b.LogLikelihood)
}
for c := range 2 {
if a.Weights[c] != b.Weights[c] {
t.Fatalf("weights moved between seeded runs: %.17g against %.17g", a.Weights[c], b.Weights[c])
}
}
}
// TestGaussianMixtureConstantSampleRefused pins the refusal a sample
// with no spread earns when k exceeds its distinct points: the
// k-means seed cannot repopulate every centre, and the mixture
// refuses instead of panicking.
func TestGaussianMixtureConstantSampleRefused(t *testing.T) {
g := core.NewGenerator(1)
x, err := core.FromFloats([]float64{2, 2}, 2, 1)
if err != nil {
t.Fatal(err)
}
if _, err := GaussianMixture(g, x, 2); err == nil {
t.Fatal("a constant sample with k = 2 was accepted")
}
}