// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) // SPDX-License-Identifier: MIT package stats import ( "math" "testing" "sourcedock.dev/petrbalvin/tensor/internal/core" ) // The hierarchical clustering against hand-worked referents. The // examples are small enough to read the merges off the distances, and // the 1-D samples make every height checkable by arithmetic. // hierSample builds an (n × 1) sample from one row per value. func hierSample(t *testing.T, vals []float64) *core.Array { t.Helper() a, err := core.FromFloats(vals, len(vals), 1) if err != nil { t.Fatal(err) } return a } func hierCluster(t *testing.T, vals []float64, method Linkage) *Dendrogram { t.Helper() d, err := HierarchicalClustering(hierSample(t, vals), method) if err != nil { t.Fatal(err) } return d } func TestHierarchicalClusteringOneDimensionalMerges(t *testing.T) { // Points 0, 1, 10, 12 on a line: the close pairs merge at 1 and // 2, and the final merge's height names the linkage. sample := []float64{0, 1, 10, 12} cases := []struct { method Linkage heights []float64 }{ {SingleLinkage, []float64{1, 2, 9}}, // nearest pair across the halves {CompleteLinkage, []float64{1, 2, 12}}, // farthest pair: 12 down to 0 {AverageLinkage, []float64{1, 2, 10.5}}, // the mean of the four cross pairs {CentroidLinkage, []float64{1, 2, 10.5}}, // centroids 0.5 and 11 are 10.5 apart // Ward's updates carry twice the within-cluster sum of squares' // increase: ESS of the four points is 112.75 against 2.5 in the // two pairs, so the height is √(2·110.25). {WardLinkage, []float64{1, 2, math.Sqrt(220.5)}}, } for _, c := range cases { d := hierCluster(t, sample, c.method) for step, h := range c.heights { if math.Abs(d.Heights[step]-h) > 1e-12 { t.Fatalf("%s merge %d height = %.16f, want %.16f", c.method, step, d.Heights[step], h) } } if d.Sizes[2] != 4 || d.Sizes[1] != 2 || d.Sizes[0] != 2 { t.Fatalf("%s sizes %v, want the close pairs first", c.method, d.Sizes) } } // The close pairs merge first under every linkage, and Left holds // the smaller id. d := hierCluster(t, sample, SingleLinkage) if d.Left[0] != 0 || d.Right[0] != 1 { t.Fatalf("the first merge joined %d and %d, want rows 0 and 1", d.Left[0], d.Right[0]) } } func TestHierarchicalClusteringCut(t *testing.T) { sample := []float64{0, 1, 10, 12, 40} d := hierCluster(t, sample, SingleLinkage) // k = 2 splits the 40 outlier from the rest. labels, err := d.Cut(2) if err != nil { t.Fatal(err) } want := []int{0, 0, 0, 0, 1} for i := range labels { if labels[i] != want[i] { t.Fatalf("Cut(2) labels %v, want %v", labels, want) } } // k = 3 splits the two close pairs apart. labels, err = d.Cut(3) if err != nil { t.Fatal(err) } want = []int{0, 0, 1, 1, 2} for i := range labels { if labels[i] != want[i] { t.Fatalf("Cut(3) labels %v, want %v", labels, want) } } // The extremes: one cluster, and the singletons in row order. labels, err = d.Cut(1) if err != nil { t.Fatal(err) } for i := range labels { if labels[i] != 0 { t.Fatalf("Cut(1) label %d is %d, want 0", i, labels[i]) } } labels, err = d.Cut(5) if err != nil { t.Fatal(err) } for i := range labels { if labels[i] != i { t.Fatalf("Cut(5) label %d is %d, want %d", i, labels[i], i) } } // The height cut at 1.5 keeps the pair {0, 1} and leaves the rest // singletons: four clusters, labels by smallest row. labels, err = d.CutHeight(1.5) if err != nil { t.Fatal(err) } want = []int{0, 0, 1, 2, 3} for i := range labels { if labels[i] != want[i] { t.Fatalf("CutHeight(1.5) labels %v, want %v", labels, want) } } // A height past the last merge is the one-cluster answer. labels, err = d.CutHeight(100) if err != nil { t.Fatal(err) } for i := range labels { if labels[i] != 0 { t.Fatalf("CutHeight(100) label %d is %d, want 0", i, labels[i]) } } } func TestHierarchicalClusteringWardRecoversBlobs(t *testing.T) { // Two tight clusters far apart: the Ward cut at two must match // the blobs, and the heights must be non-decreasing (Ward is a // monotone linkage). blob := func(centre, spread float64, count int, offset int) []float64 { out := make([]float64, count) for i := range count { out[i] = centre + spread*math.Sin(float64(i+offset)) } return out } vals := append(blob(0, 0.1, 8, 0), blob(50, 0.1, 8, 3)...) d := hierCluster(t, vals, WardLinkage) for step := 1; step < len(d.Heights); step++ { if d.Heights[step] < d.Heights[step-1] { t.Fatalf("Ward heights descend at %d: %g after %g", step, d.Heights[step], d.Heights[step-1]) } } labels, err := d.Cut(2) if err != nil { t.Fatal(err) } for i := range labels { want := 0 if i >= 8 { want = 1 } if labels[i] != want { t.Fatalf("row %d landed in cluster %d, want %d", i, labels[i], want) } } // Single and complete linkage recover the same partition, and // their heights are monotone too. for _, method := range []Linkage{SingleLinkage, CompleteLinkage, AverageLinkage} { d := hierCluster(t, vals, method) for step := 1; step < len(d.Heights); step++ { if d.Heights[step] < d.Heights[step-1] { t.Fatalf("%s heights descend at %d", method, step) } } labels, err := d.Cut(2) if err != nil { t.Fatal(err) } if labels[0] != labels[7] || labels[8] != labels[15] || labels[0] == labels[8] { t.Fatalf("%s split the blobs apart wrongly: %v", method, labels) } } } func TestHierarchicalClusteringDeterministic(t *testing.T) { vals := make([]float64, 16) for i := range vals { vals[i] = float64((i*37)%23) / 3 } run := func() *Dendrogram { return hierCluster(t, vals, AverageLinkage) } a, b := run(), run() for step := range a.Heights { if a.Heights[step] != b.Heights[step] || a.Left[step] != b.Left[step] || a.Right[step] != b.Right[step] || a.Sizes[step] != b.Sizes[step] { t.Fatalf("two identical runs disagreed at merge %d", step) } } } func TestHierarchicalClusteringRefusals(t *testing.T) { if _, err := HierarchicalClustering(hierSample(t, []float64{1}), WardLinkage); err == nil { t.Fatal("a one-row sample was accepted") } if _, err := HierarchicalClustering(mustStatArray(t, []float64{1, 2, 3, 4}, 4), WardLinkage); err == nil { t.Fatal("a rank-1 sample was accepted") } if _, err := HierarchicalClustering(hierSample(t, []float64{1, math.NaN()}), WardLinkage); err == nil { t.Fatal("a non-finite sample was accepted") } if _, err := HierarchicalClustering(hierSample(t, []float64{1, 2}), Linkage(9)); err == nil { t.Fatal("an unknown linkage was accepted") } // One row over the cap, refused before the quadratic matrix. big := make([]float64, HierarchicalMaxObservations+1) for i := range big { big[i] = float64(i) } if _, err := HierarchicalClustering(hierSample(t, big), WardLinkage); err == nil { t.Fatal("an over-cap sample was accepted") } d := hierCluster(t, []float64{0, 1, 10}, SingleLinkage) if _, err := d.Cut(0); err == nil { t.Fatal("Cut(0) was accepted") } if _, err := d.Cut(4); err == nil { t.Fatal("Cut(4) on a three-row sample was accepted") } if _, err := d.CutHeight(math.NaN()); err == nil { t.Fatal("a NaN height was accepted") } var nilD *Dendrogram if _, err := nilD.Cut(1); err == nil { t.Fatal("a nil dendrogram was accepted") } }