350 lines
11 KiB
Go
350 lines
11 KiB
Go
// 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
|
||
}
|