437 lines
17 KiB
Go
437 lines
17 KiB
Go
// 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")
|
||
|
|
}
|
||
|
|
}
|