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