feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -0,0 +1,242 @@
|
||||
// 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")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user