feat: initial release
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s

Assisted-by: GLM 5.3 Flash
This commit is contained in:
2026-09-03 10:00:00 +02:00
commit af4ee19703
617 changed files with 191195 additions and 0 deletions
+278
View File
@@ -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
}