243 lines
7.2 KiB
Go
243 lines
7.2 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (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")
|
|||
|
|
}
|
|||
|
|
}
|