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
|
|||
|
|
}
|