279 lines
8.3 KiB
Go
279 lines
8.3 KiB
Go
// 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
|
||
}
|