Files
tensor/stats/hierarchy_test.go
T
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

243 lines
7.2 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// 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")
}
}