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

279 lines
8.3 KiB
Go
Raw 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 (
"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
}