Files
tensor/stats/hierarchy.go
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

350 lines
11 KiB
Go
Raw Permalink 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"
"slices"
"strconv"
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// Agglomerative hierarchical clustering: every observation starts in
// its own cluster and the closest pair merges repeatedly until one
// cluster holds the sample, with the rule that names "closest"
// carried by the linkage. The distances are Euclidean over the
// sample's rows, the merge distances update through the
// Lance-Williams recurrences, and ties resolve toward the lowest
// indices, so the dendrogram is a deterministic function of the
// input. The result is the merge record itself, the dendrogram, from
// which any number of flat partitions can be cut afterwards.
// Linkage names the rule a merge's distance follows.
type Linkage int
const (
// SingleLinkage joins on the smallest pair distance across the two
// clusters: the nearest-neighbour rule that follows chains.
SingleLinkage Linkage = iota
// CompleteLinkage joins on the largest pair distance, the
// farthest-neighbour rule that keeps clusters tight.
CompleteLinkage
// AverageLinkage joins on the mean pair distance over all pairs
// across the two clusters, the UPGMA compromise.
AverageLinkage
// CentroidLinkage joins on the distance between the clusters'
// centroids. The rule is not monotone: a later merge can sit
// below an earlier one in the dendrogram.
CentroidLinkage
// WardLinkage joins the pair whose merge raises the within-cluster
// sum of squares the least. The heights are the square roots of
// the recorded updates, the scale the two-singleton merges share
// with the plain distances.
WardLinkage
)
// String returns the linkage's name, the spelling the error messages
// carry.
func (l Linkage) String() string {
switch l {
case SingleLinkage:
return "single"
case CompleteLinkage:
return "complete"
case AverageLinkage:
return "average"
case CentroidLinkage:
return "centroid"
case WardLinkage:
return "ward"
}
return "linkage(" + strconv.Itoa(int(l)) + ")"
}
// HierarchicalMaxObservations is the sample cap of
// HierarchicalClustering: the working distance matrix is quadratic in
// the sample, and past the cap the refusal names the cost rather than
// handing gigabytes to the allocator. A larger sample wants an
// approximate construction the library does not pretend to carry.
const HierarchicalMaxObservations = 4096
// Dendrogram records an agglomerative clustering as its sequence of
// merges. Leaf i is the sample's row i; the merge at step t creates
// cluster n+t from the clusters Left[t] and Right[t], at
// Heights[t], holding Sizes[t] rows. Left always holds the smaller
// cluster id.
type Dendrogram struct {
// Left and Right name the two clusters each merge joins.
Left []int
Right []int
// Heights are the merge distances, in merge order. For the
// centroid and Ward linkages an inversion is possible and
// legitimate: the sequence need not be non-decreasing.
Heights []float64
// Sizes counts the rows each merge's cluster carries.
Sizes []int
}
// Cut partitions the sample into k flat clusters by undoing the last
// k−1 merges, and labels every row 0 to k−1. The labels are ordered
// by each cluster's smallest row, so the first row of the sample
// always lands in cluster 0. Refuses a nil dendrogram and a k outside
// [1, n].
func (d *Dendrogram) Cut(k int) ([]int, error) {
const name = "Dendrogram.Cut"
if d == nil {
return nil, base.Errf("%s: the dendrogram is nil", name)
}
n := len(d.Sizes) + 1
if k < 1 || k > n {
return nil, base.Errf("%s: k must lie in [1, %d], got %d", name, n, k)
}
return d.cutMerges(n - k), nil
}
// CutHeight partitions the sample by applying merges in order while
// they sit at or below the given height, and labels every row 0
// upward by each cluster's smallest row, as Cut does. A height below
// the first merge returns the singletons; one at or above the last
// returns the whole sample as one cluster. Refuses a nil dendrogram
// and a non-finite height.
func (d *Dendrogram) CutHeight(height float64) ([]int, error) {
const name = "Dendrogram.CutHeight"
if d == nil {
return nil, base.Errf("%s: the dendrogram is nil", name)
}
if math.IsNaN(height) || math.IsInf(height, 0) {
return nil, base.Errf("%s: the height must be finite, got %g", name, height)
}
m := 0
for m < len(d.Heights) && d.Heights[m] <= height {
m++
}
return d.cutMerges(m), nil
}
// cutMerges applies the first m merges and labels the rows by
// cluster, in order of each cluster's smallest row.
func (d *Dendrogram) cutMerges(m int) []int {
n := len(d.Sizes) + 1
// The merge tree as a parent table: the merge at step t lifts both
// children under the cluster n+t. Every id exceeds its parents, so
// finding a row's root is a plain upward walk.
parent := make([]int, 2*n-1)
for i := range parent {
parent[i] = i
}
// The smallest row each applied merge's cluster holds.
smallest := make([]int, 2*n-1)
for i := range n {
smallest[i] = i
}
for t := range m {
a, b := d.Left[t], d.Right[t]
parent[a] = n + t
parent[b] = n + t
c := n + t
smallest[c] = min(smallest[a], smallest[b])
}
root := make([]int, n)
for i := range n {
c := i
for parent[c] != c {
c = parent[c]
}
root[i] = c
}
// The distinct smallest-rows, sorted, are the label indices.
reps := make([]int, 0, n)
seen := make(map[int]bool)
for i := range n {
r := smallest[root[i]]
if !seen[r] {
seen[r] = true
reps = append(reps, r)
}
}
slices.Sort(reps)
labels := make([]int, n)
for i := range n {
labels[i] = slices.Index(reps, smallest[root[i]])
}
return labels
}
// HierarchicalClustering builds the dendrogram of the sample's rows
// under the given linkage, over Euclidean distances. The algorithm is
// the naive agglomerative sweep: each step scans the working distance
// matrix for the closest pair, merges it, and rewrites the merged
// cluster's distances through the linkage's Lance-Williams update, so
// the run costs O(n³) time and O(n²) memory and is deterministic for
// the input, ties included.
//
// Refuses a sample that is not rank 2, complex input, a non-finite
// entry, fewer than two rows, more than HierarchicalMaxObservations
// rows, and an unknown linkage.
func HierarchicalClustering(x *core.Array, method Linkage) (*Dendrogram, error) {
const name = "HierarchicalClustering"
switch method {
case SingleLinkage, CompleteLinkage, AverageLinkage, CentroidLinkage, WardLinkage:
default:
return nil, base.Errf("%s: unknown linkage %s", name, method)
}
if x.NDim() != 2 {
return nil, base.Errf("%s: the sample must be rank 2, got shape %s", name, base.ShapeText(x.Shape()))
}
if x.Dtype() == core.Complex {
return nil, base.Errf("%s: complex samples are not supported", name)
}
if err := checkFinite(name, "the sample", x); err != nil {
return nil, err
}
n := x.Shape()[0]
d := x.Shape()[1]
if n < 2 {
return nil, base.Errf("%s: at least two observations are needed, got %d", name, n)
}
if n > HierarchicalMaxObservations {
return nil, base.Errf("%s: %d observations exceed the %d cap; the working matrix alone would hold %.1f GiB",
name, n, HierarchicalMaxObservations, float64(n)*float64(n)*8/(1<<30))
}
// The sample is read once and the squared Euclidean distances are
// built from the flat rows: centroid and Ward update squared
// distances, and the plain linkages take their square root once at
// the start.
data := make([]float64, n*d)
if fs := rawFloats(x); fs != nil {
copy(data, fs[:n*d])
} else {
for i := range data {
data[i] = x.FloatAt(i)
}
}
squared := method == CentroidLinkage || method == WardLinkage
dist := make([]float64, n*n)
for i := range n {
for j := range i {
total := 0.0
for c := range d {
diff := data[i*d+c] - data[j*d+c]
total += diff * diff
}
if squared {
dist[i*n+j] = total
dist[j*n+i] = total
} else {
plain := math.Sqrt(total)
dist[i*n+j] = plain
dist[j*n+i] = plain
}
}
}
// The active clusters live in positions 0..k−1 of the working
// matrix; actID names each position's cluster id and actSize its
// row count. A merge rewrites the matrix in place: the merged
// cluster takes position i, the last position's rows move into j's,
// and the working width drops by one.
actID := make([]int, n)
actSize := make([]int, n)
for i := range n {
actID[i] = i
actSize[i] = 1
}
dendrogram := &Dendrogram{
Left: make([]int, n-1),
Right: make([]int, n-1),
Heights: make([]float64, n-1),
Sizes: make([]int, n-1),
}
k := n
for step := range n - 1 {
// The closest pair, ties toward the lowest positions. The
// matrix keeps its full stride n for its whole life: the
// active clusters occupy positions 0..k−1 and only the loop
// bounds shrink.
bestI, bestJ := 0, 1
bestD := math.Inf(1)
for i := range k {
for j := i + 1; j < k; j++ {
if dist[i*n+j] < bestD {
bestD = dist[i*n+j]
bestI, bestJ = i, j
}
}
}
ni, nj := actSize[bestI], actSize[bestJ]
idI, idJ := actID[bestI], actID[bestJ]
height := bestD
if squared {
// The centroid and Ward updates carry squared quantities;
// the recorded heights take their root. A rounding descent
// below zero is clamped, the square root refusing it
// otherwise.
height = math.Sqrt(math.Max(0, bestD))
}
dendrogram.Left[step] = min(idI, idJ)
dendrogram.Right[step] = max(idI, idJ)
dendrogram.Heights[step] = height
dendrogram.Sizes[step] = ni + nj
// The merged distances to every surviving cluster, then the
// compaction: the last position's cluster moves into bestJ's
// slot before the working width drops by one.
last := k - 1
for q := range k {
if q == bestI || q == bestJ {
continue
}
dik := dist[bestI*n+q]
djk := dist[bestJ*n+q]
var updated float64
switch method {
case SingleLinkage:
updated = min(dik, djk)
case CompleteLinkage:
updated = max(dik, djk)
case AverageLinkage:
updated = (float64(ni)*dik + float64(nj)*djk) / float64(ni+nj)
case CentroidLinkage:
total := float64(ni + nj)
updated = (float64(ni)*dik+float64(nj)*djk)/total -
float64(ni)*float64(nj)*bestD/(total*total)
case WardLinkage:
total := float64(ni + nj + actSize[q])
updated = ((float64(ni+actSize[q]))*dik + (float64(nj+actSize[q]))*djk -
float64(actSize[q])*bestD) / total
}
if squared {
updated = math.Max(0, updated)
}
dist[bestI*n+q] = updated
dist[q*n+bestI] = updated
}
if bestJ != last {
// The last position's cluster moves into bestJ's slot: its
// row and column transfer whole, the merged cluster's own
// entries against it included, and only the diagonals are
// left alone, zero on both sides.
for q := range k {
switch {
case q == bestJ:
case q == bestI:
dist[bestI*n+bestJ] = dist[bestI*n+last]
dist[bestJ*n+bestI] = dist[bestI*n+last]
default:
dist[bestJ*n+q] = dist[last*n+q]
dist[q*n+bestJ] = dist[q*n+last]
}
}
actID[bestJ] = actID[last]
actSize[bestJ] = actSize[last]
}
actID[bestI] = n + step
actSize[bestI] = ni + nj
k = last
}
return dendrogram, nil
}