Files
tensor/stats/rankcorr_test.go
T

128 lines
6.1 KiB
Go
Raw Normal View History

2026-09-03 10:00:00 +02:00
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package stats
import (
"math"
"strings"
"testing"
)
// TestSpearmanRho pins ρ on exact cases: a perfect monotone pairing is
// exactly ±1, the classic d² example x = (1..5) against
// y = (3, 1, 4, 2, 5) has rank differences (−2, 1, −1, 2, 0), so
// Σd² = 10 and ρ = 1 − 60/120 = 0.5, and the tied pairing
// x = (1, 2, 2, 4) against y = (1, 2, 3, 4) works out through the
// mid-ranks to √0.9.
func TestSpearmanRho(t *testing.T) {
x := mustFloats(t, []float64{1, 2, 3, 4, 5})
y := mustFloats(t, []float64{10, 20, 30, 40, 50})
rho, err := SpearmanRho(x, y)
if err != nil || rho != 1 {
t.Fatalf("perfect monotone ρ = %v, %v, want exactly 1", rho, err)
}
rho, err = SpearmanRho(x, mustFloats(t, []float64{50, 40, 30, 20, 10}))
if err != nil || rho != -1 {
t.Fatalf("reversed ρ = %v, %v, want exactly −1", rho, err)
}
rho, err = SpearmanRho(x, mustFloats(t, []float64{3, 1, 4, 2, 5}))
if err != nil || math.Abs(rho-0.5) > 1e-14 {
t.Fatalf("classic example ρ = %v, %v, want 0.5", rho, err)
}
rho, err = SpearmanRho(
mustFloats(t, []float64{1, 2, 2, 4}),
mustFloats(t, []float64{1, 2, 3, 4}))
if err != nil || math.Abs(rho-math.Sqrt(0.9)) > 1e-14 {
t.Fatalf("tied ρ = %.16g, %v, want √0.9 = %.16g", rho, err, math.Sqrt(0.9))
}
// The tie correction is honest: the same pairing without the tie in
// x ranks higher.
untied, _ := SpearmanRho(mustFloats(t, []float64{1, 2, 3, 4}),
mustFloats(t, []float64{1, 2, 3, 4}))
if untied != 1 || rho >= 1 {
t.Fatalf("ties did not lower ρ: tied %v, untied %v", rho, untied)
}
if _, err := SpearmanRho(x, mustFloats(t, []float64{1, 2, 3})); err == nil || !strings.Contains(err.Error(), "the samples have") {
t.Fatalf("length mismatch: got %v, want the length refusal", err)
}
if _, err := SpearmanRho(mustFloats(t, []float64{1}), mustFloats(t, []float64{1})); err == nil || !strings.Contains(err.Error(), "at least two paired") {
t.Fatalf("one pair: got %v, want the pairing floor refusal", err)
}
if _, err := SpearmanRho(mustFloats(t, []float64{1, 2, 3}), mustFloats(t, []float64{5, 5, 5})); err == nil || !strings.Contains(err.Error(), "ranks all agree") {
t.Fatalf("a constant sample: got %v, want the constant-sample refusal", err)
}
if _, err := SpearmanRho(mustFloats(t, []float64{1, math.NaN(), 3}), mustFloats(t, []float64{1, 2, 3})); err == nil || !strings.Contains(err.Error(), "non-finite") {
t.Fatalf("a NaN observation: got %v, want the non-finite refusal", err)
}
if _, err := SpearmanRho(mustComplexes(t, []complex128{1, 2, 3}, 3), mustFloats(t, []float64{1, 2, 3})); err == nil || !strings.Contains(err.Error(), "complex") {
t.Fatalf("a complex sample: got %v, want the complex refusal", err)
}
}
// TestKendallTau pins τ-b on hand-counted cases. For x = (1, 2, 3, 4)
// against y = (1, 3, 2, 4) only the pair (2, 3) against (3, 2) is
// discordant, so τ = (5−1)/6 = 2/3. With the tie x = (1, 2, 2, 4) the
// denominator shrinks to √((6−1)(6−0)) and τ-b = 5/√30; the naive n₀
// denominator would report 5/6 and miss the attainable maximum. The
// all-tied pairing x = (1, 1, 2, 2) against y = (5, 5, 7, 7) is
// perfectly monotone and must give exactly ±1, where the naive
// normalisation would report 4/6.
func TestKendallTau(t *testing.T) {
tau, err := KendallTau(
mustFloats(t, []float64{1, 2, 3, 4}),
mustFloats(t, []float64{1, 3, 2, 4}))
if err != nil || math.Abs(tau-2.0/3) > 1e-15 {
t.Fatalf("hand case τ = %.16g, %v, want 2/3", tau, err)
}
tau, err = KendallTau(
mustFloats(t, []float64{1, 2, 2, 4}),
mustFloats(t, []float64{1, 2, 3, 4}))
if err != nil || math.Abs(tau-5/math.Sqrt(30)) > 1e-14 {
t.Fatalf("tied τ-b = %.16g, %v, want 5/√30 = %.16g", tau, err, 5/math.Sqrt(30))
}
tau, err = KendallTau(
mustFloats(t, []float64{1, 1, 2, 2}),
mustFloats(t, []float64{5, 5, 7, 7}))
if err != nil || tau != 1 {
t.Fatalf("perfectly monotone with ties τ-b = %v, %v, want exactly 1", tau, err)
}
tau, err = KendallTau(
mustFloats(t, []float64{1, 1, 2, 2}),
mustFloats(t, []float64{7, 7, 5, 5}))
if err != nil || tau != -1 {
t.Fatalf("perfectly reversed with ties τ-b = %v, %v, want exactly −1", tau, err)
}
// A tie-heavy pairing: 3 concordant and 3 discordant cross pairs
// cancel exactly, so τ-b is exactly 0 with the denominator
// √((15−6)(15−3)) well defined. Shifting the second block up turns
// it into 6 concordant, 1 discordant, τ-b = 5/√117.
tau, err = KendallTau(
mustFloats(t, []float64{1, 1, 1, 2, 2, 2}),
mustFloats(t, []float64{1, 2, 3, 1, 2, 3}))
if err != nil || tau != 0 {
t.Fatalf("tie-heavy τ-b = %.16g, %v, want exactly 0", tau, err)
}
tau, err = KendallTau(
mustFloats(t, []float64{1, 1, 1, 2, 2, 2}),
mustFloats(t, []float64{1, 2, 3, 2, 3, 4}))
if err != nil || math.Abs(tau-5/math.Sqrt(117)) > 1e-14 {
t.Fatalf("tie-heavy shifted τ-b = %.16g, %v, want 5/√117", tau, err)
}
if _, err := KendallTau(mustFloats(t, []float64{1, 2, 3}), mustFloats(t, []float64{1, 2})); err == nil || !strings.Contains(err.Error(), "the samples have") {
t.Fatalf("length mismatch: got %v, want the length refusal", err)
}
if _, err := KendallTau(mustFloats(t, []float64{1}), mustFloats(t, []float64{1})); err == nil || !strings.Contains(err.Error(), "at least two paired") {
t.Fatalf("one pair: got %v, want the pairing floor refusal", err)
}
if _, err := KendallTau(mustFloats(t, []float64{4, 4, 4}), mustFloats(t, []float64{1, 2, 3})); err == nil || !strings.Contains(err.Error(), "denominator zero") {
t.Fatalf("a constant sample: got %v, want the constant-sample refusal", err)
}
if _, err := KendallTau(mustFloats(t, []float64{1, math.Inf(1), 3}), mustFloats(t, []float64{1, 2, 3})); err == nil || !strings.Contains(err.Error(), "non-finite") {
t.Fatalf("a non-finite observation: got %v, want the non-finite refusal", err)
}
if _, err := KendallTau(mustFloats(t, []float64{1, 2, 3}), mustComplexes(t, []complex128{1, 2, 3}, 3)); err == nil || !strings.Contains(err.Error(), "complex") {
t.Fatalf("a complex second sample: got %v, want the complex refusal", err)
}
}