// Copyright (c) 2026 Petr Balvín (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 }