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