// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) // SPDX-License-Identifier: MIT package stats import ( "sourcedock.dev/petrbalvin/tensor/internal/base" "sourcedock.dev/petrbalvin/tensor/internal/core" ) import ( "cmp" "math" "slices" ) // Rank correlations over paired observations: Spearman's ρ and // Kendall's τ-b. Both refuse length mismatches, fewer than two pairs, // complex or non-finite input, the contract every inference entry // point of the package shares: a NaN would otherwise flow silently // through the ranks and the pair counts. // SpearmanRho returns Spearman's rank correlation of the paired // samples x and y: Pearson's correlation evaluated on the mid-ranks, // the averaged ranks ties share. Computing it as the Pearson of the // mid-ranks is exactly the tie-corrected form, so tied observations // lower ρ honestly instead of pretending a finer ordering than the // data carry. A sample whose ranks have zero variance (all values // equal) is refused: the correlation is undefined, not zero. func SpearmanRho(x, y *core.Array) (float64, error) { const name = "SpearmanRho" rx, ry, err := pairedSamples(name, x, y) if err != nil { return 0, err } rankX := midRanks(rx) rankY := midRanks(ry) n := float64(len(rx)) mx, my := 0.0, 0.0 for i := range rx { mx += rankX[i] my += rankY[i] } mx /= n my /= n sxx, syy, sxy := 0.0, 0.0, 0.0 for i := range rx { dx := rankX[i] - mx dy := rankY[i] - my sxx += dx * dx syy += dy * dy sxy += dx * dy } if sxx == 0 || syy == 0 { return 0, base.Errf("%s: a sample whose ranks all agree has no rank correlation", name) } return sxy / math.Sqrt(sxx*syy), nil } // KendallTau returns Kendall's τ-b for the paired samples x and y: the // concordant minus discordant pair count over the tie-corrected // denominator √((n₀−n₁)(n₀−n₂)), with n₀ the pair total and n₁, n₂ the // within-sample tied pair counts. The normalisation keeps τ inside // [−1, 1] under ties and reaches ±1 on every perfectly monotone // pairing, which the naive n₀ denominator cannot; pairs tied in both // samples count for neither side. func KendallTau(x, y *core.Array) (float64, error) { const name = "KendallTau" xs, ys, err := pairedSamples(name, x, y) if err != nil { return 0, err } // The pair counts are exact well below 2⁵³ pairs, so the walks run // in float64 without an overflow thought. n0 := float64(len(xs)) * float64(len(xs)-1) / 2 tally := tallyPairs(xs, ys) n1 := float64(tally.tiedX) n2 := float64(tally.tiedY) if n0 == n1 || n0 == n2 { return 0, base.Errf("%s: a constant sample leaves the denominator zero", name) } concordant := float64(tally.concordant) discordant := float64(tally.discordant) return (concordant - discordant) / math.Sqrt((n0-n1)*(n0-n2)), nil } // pairedSamples extracts two equal-length real samples, the shared // gate of both rank correlations: complex input, a length mismatch, // fewer than two pairs and non-finite values are all refused here, // under the entry point's own name. func pairedSamples(name string, x, y *core.Array) ([]float64, []float64, error) { if x.Dtype() == core.Complex || y.Dtype() == core.Complex { return nil, nil, base.Errf("%s: complex samples are not supported", name) } n := x.Len() if y.Len() != n { return nil, nil, base.Errf("%s: the samples have %d and %d observations", name, n, y.Len()) } if n < 2 { return nil, nil, base.Errf("%s: at least two paired observations are needed, got %d", name, n) } if err := checkFinite(name, "the first sample", x); err != nil { return nil, nil, err } if err := checkFinite(name, "the second sample", y); err != nil { return nil, nil, err } xs := make([]float64, n) if fs := rawFloats(x); fs != nil { copy(xs, fs) } else { for i := range n { xs[i] = x.FloatAt(i) } } ys := make([]float64, n) if fs := rawFloats(y); fs != nil { copy(ys, fs) } else { for i := range n { ys[i] = y.FloatAt(i) } } return xs, ys, nil } // midRanks returns the mid-ranks of vals: 1-based ranks with ties // sharing the average of the ranks they span, the ordering both // correlations count on. func midRanks(vals []float64) []float64 { order := make([]int, len(vals)) for i := range order { order[i] = i } slices.SortFunc(order, func(a, b int) int { return cmp.Compare(vals[a], vals[b]) }) ranks := make([]float64, len(vals)) for i := 0; i < len(order); { j := i for j < len(order) && vals[order[j]] == vals[order[i]] { j++ } // The average of the 1-based ranks i+1 through j. mid := float64(i+j+1) / 2 for k := i; k < j; k++ { ranks[order[k]] = mid } i = j } return ranks } // rankPair is one paired observation, the sort key of the pair walk. type rankPair struct{ x, y float64 } // pairTally holds the exact pair counts of a paired sample: the // concordant and discordant pairs and the pairs tied within the first // and within the second sample. Every field is the integer the O(n²) // enumeration would produce, so the τ-b numerator and denominator built // from them carry the same bits. type pairTally struct { concordant int64 discordant int64 tiedX int64 tiedY int64 } // tallyPairs counts the pairs of two paired samples in O(n log n). The // pairs are first ordered by (x, y): a pair is then discordant exactly // when the reordered y sequence descends, which the merge sort of // sortTally counts as it sorts, and every remaining pair is either // ascending in y or tied in y. Ascending pairs inside a block of equal // x are tied in the predictor and count for neither side either, so the // block's ordered pairs, its ties and its pairs tied in both are // removed by one scan of the sorted pairs. A pair tied in either sample // counts for neither side, which is the definition the enumeration // applies by testing the product of the two differences against zero. func tallyPairs(xs, ys []float64) pairTally { n := len(xs) pairs := make([]rankPair, n) for i := range n { pairs[i] = rankPair{x: xs[i], y: ys[i]} } slices.SortFunc(pairs, func(a, b rankPair) int { if c := cmp.Compare(a.x, b.x); c != 0 { return c } return cmp.Compare(a.y, b.y) }) ordered := make([]float64, n) for i := range pairs { ordered[i] = pairs[i].y } var tally pairTally tally.discordant, tally.tiedY = sortTally(ordered) total := int64(n) * int64(n-1) / 2 // The blocks of equal x hold their y values ascending, so no pair // inside one is an inversion: the block's pairs are tied in x, and // the concordant count is what is left of the ascending pairs once // the block's own ordered pairs are taken out. tiedBoth := int64(0) for lo := 0; lo < n; { hi := lo + 1 for hi < n && pairs[hi].x == pairs[lo].x { hi++ } m := int64(hi - lo) tally.tiedX += m * (m - 1) / 2 for i := lo; i < hi; { j := i + 1 for j < hi && pairs[j].y == pairs[i].y { j++ } c := int64(j - i) tiedBoth += c * (c - 1) / 2 i = j } lo = hi } ascending := total - tally.discordant - tally.tiedY tally.concordant = ascending - (tally.tiedX - tiedBoth) return tally } // sortTally sorts vals by merge sort and returns the number of pairs // i < j with vals[i] > vals[j] and the number of pairs i < j with // vals[i] == vals[j]. Both are exact integers, so the counts are the // ones a brute-force enumeration of the pairs produces, at O(n log n) // cost: an element taken from the right run descends past every left // element still waiting, and the equal pairs are the groups the // finished sort leaves. vals is consumed as scratch and left sorted. func sortTally(vals []float64) (descending, equal int64) { n := len(vals) if n < 2 { return 0, 0 } buf := make([]float64, n) src, dst := vals, buf for width := 1; width < n; width *= 2 { for lo := 0; lo < n; lo += 2 * width { mid := min(lo+width, n) hi := min(lo+2*width, n) i, j, k := lo, mid, lo for i < mid && j < hi { if src[i] <= src[j] { dst[k] = src[i] i++ } else { descending += int64(mid - i) dst[k] = src[j] j++ } k++ } for i < mid { dst[k] = src[i] i++ k++ } for j < hi { dst[k] = src[j] j++ k++ } } // Every position was written, so src holds the merged sequence. src, dst = dst, src } for i := 0; i < n; { j := i + 1 for j < n && src[j] == src[i] { j++ } c := int64(j - i) equal += c * (c - 1) / 2 i = j } return descending, equal }