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