Files
tensor/stats/pca_test.go
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

442 lines
14 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package stats
import (
"math"
"slices"
"strings"
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// pcaFixture builds a seeded (n, 4) observation array with correlated
// columns of clearly different scales, the ordinary material a PCA
// runs on.
func pcaFixture(t *testing.T, n int, seed int64) *core.Array {
t.Helper()
g := core.NewGenerator(31)
x0 := make([]float64, n)
vals := make([]float64, 0, 4*n)
for i := range n {
x0[i] = g.NormalUnit()
}
for i := range n {
x1 := 0.8*x0[i] + 0.6*g.NormalUnit()
x2 := -0.5*x0[i] + g.NormalUnit()
x3 := 0.3 * g.NormalUnit()
vals = append(vals, x0[i], x1, x2, x3)
}
return mustFromFloats(t, vals, n, 4)
}
// TestPCALongAxis puts a two-cluster anisotropic cloud under the
// decomposition: two Gaussian blobs strung along a thirty-degree axis
// must come back with the first component along that axis, the
// measured angle against the truth, and nearly all the variance on it.
func TestPCALongAxis(t *testing.T) {
const n = 80
g := core.NewGenerator(29)
const theta = math.Pi / 6
cos, sin := math.Cos(theta), math.Sin(theta)
vals := make([]float64, 0, 4*n)
for shift := 0.0; shift <= 10; shift += 10 {
for range n {
u := 5 * g.NormalUnit()
v := 0.5 * g.NormalUnit()
vals = append(vals,
u*cos-v*sin+shift*cos,
u*sin+v*cos+shift*sin)
}
}
design := mustFromFloats(t, vals, 2*n, 2)
res, err := PCA(design)
if err != nil {
t.Fatalf("PCA: %v", err)
}
// The angle of a component's axis, read off its two loadings and
// folded into (−π/2, π/2], where the fixed sign convention leaves
// it. Loadings entry (j, k) sits at j·p+k.
angleOf := func(k int) float64 {
a := math.Atan2(res.Loadings.FloatAt(1*2+k), res.Loadings.FloatAt(0*2+k))
if a > math.Pi/2 {
a -= math.Pi
}
if a <= -math.Pi/2 {
a += math.Pi
}
return a
}
// Axis directions are defined only up to a half turn, so the
// short axis's angle is folded like the first before comparing.
first := angleOf(0)
second := angleOf(1)
secondWant := theta + math.Pi/2
if secondWant > math.Pi/2 {
secondWant -= math.Pi
}
t.Logf("first axis at %.4f rad against %.4f, second at %.4f against %.4f, explaining %.4f of the variance",
first, theta, second, secondWant, res.ExplainedVarianceRatio[0])
if math.Abs(first-theta) > 0.05 {
t.Fatalf("the first component sits at %.4f rad, want the long axis %.4f", first, theta)
}
if math.Abs(second-secondWant) > 0.05 {
t.Fatalf("the second component sits at %.4f rad, want the short axis %.4f", second, secondWant)
}
if res.ExplainedVarianceRatio[0] < 0.98 {
t.Fatalf("the long axis explains only %.4f of the variance", res.ExplainedVarianceRatio[0])
}
}
// TestPCAVariancesAndLoadings pins the spectral accounting: the
// eigenvalues sum to the covariance's trace, the ratios sum to one,
// they arrive in falling order, the loadings are orthonormal as
// columns, and the loadings and eigenvalues rebuild the covariance
// they came from.
func TestPCAVariancesAndLoadings(t *testing.T) {
const n, p = 60, 4
a := pcaFixture(t, n, 31)
res, err := PCA(a)
if err != nil {
t.Fatalf("PCA: %v", err)
}
// The trace, computed here straight from the data.
means := make([]float64, p)
for j := range p {
s := 0.0
for i := range n {
s += a.FloatAt(i*p + j)
}
means[j] = s / float64(n)
}
trace := 0.0
for j := range p {
s := 0.0
for i := range n {
d := a.FloatAt(i*p+j) - means[j]
s += d * d
}
trace += s / float64(n-1)
}
total := 0.0
for k := range p {
total += res.ExplainedVariance[k]
}
if math.Abs(total-trace) > 1e-10*math.Max(1, trace) {
t.Fatalf("the eigenvalues sum to %.12g, want the trace %.12g", total, trace)
}
ratioSum := 0.0
for k := range p {
ratioSum += res.ExplainedVarianceRatio[k]
if k > 0 && res.ExplainedVariance[k] > res.ExplainedVariance[k-1] {
t.Fatalf("the variances are not descending at %d", k)
}
}
if math.Abs(ratioSum-1) > 1e-12 {
t.Fatalf("the ratios sum to %.16g, want 1", ratioSum)
}
// Orthonormal columns: LᵀL is the identity.
for j := range p {
for k := j; k < p; k++ {
dot := 0.0
for i := range p {
dot += res.Loadings.FloatAt(i*p+j) * res.Loadings.FloatAt(i*p+k)
}
want := 0.0
if j == k {
want = 1
}
if math.Abs(dot-want) > 1e-10 {
t.Fatalf("loadings %d and %d have inner product %.4g, want %.4g", j, k, dot, want)
}
}
}
// The covariance rebuilds: L·D·Lᵀ against the entries the package
// computed.
cov, err := CovarianceMatrix(a)
if err != nil {
t.Fatalf("CovarianceMatrix: %v", err)
}
for i := range p {
for j := range p {
s := 0.0
for k := range p {
s += res.ExplainedVariance[k] * res.Loadings.FloatAt(i*p+k) * res.Loadings.FloatAt(j*p+k)
}
if math.Abs(s-cov.FloatAt(i*p+j)) > 1e-9*math.Max(1, math.Abs(cov.FloatAt(i*p+j))) {
t.Fatalf("the rebuilt covariance entry (%d, %d) is %.12g, want %.12g", i, j, s, cov.FloatAt(i*p+j))
}
}
}
}
// TestPCAScoresWhitenRoundTrip pins the transforms: the scores are
// the centred observations times the loadings, the whitened data is
// the scores scaled by the components' standard deviations, the
// covariance of the whitened data is the identity, and unwhitening
// returns the centred observations.
func TestPCAScoresWhitenRoundTrip(t *testing.T) {
const n, p = 60, 4
a := pcaFixture(t, n, 31)
res, err := PCA(a)
if err != nil {
t.Fatalf("PCA: %v", err)
}
// Scores against their definition.
for i := range n {
for k := range p {
s := 0.0
for j := range p {
s += (a.FloatAt(i*p+j) - res.Mean[j]) * res.Loadings.FloatAt(j*p+k)
}
if math.Abs(s-res.Scores.FloatAt(i*p+k)) > 1e-9 {
t.Fatalf("score (%d, %d) is %.12g, want the centred row times the loading %.12g",
i, k, res.Scores.FloatAt(i*p+k), s)
}
}
}
// Whitened against the scores scaled by the component scales, and
// with identity covariance.
z, err := res.Whiten(a)
if err != nil {
t.Fatalf("Whiten: %v", err)
}
for i := range n {
for k := range p {
want := res.Scores.FloatAt(i*p+k) / math.Sqrt(res.ExplainedVariance[k])
if math.Abs(z.FloatAt(i*p+k)-want) > 1e-8 {
t.Fatalf("whitened (%d, %d) is %.12g, want the scaled score %.12g",
i, k, z.FloatAt(i*p+k), want)
}
}
}
zcov, err := CovarianceMatrix(z)
if err != nil {
t.Fatalf("CovarianceMatrix of the whitened data: %v", err)
}
for i := range p {
for j := range p {
want := 0.0
if i == j {
want = 1
}
if math.Abs(zcov.FloatAt(i*p+j)-want) > 1e-9 {
t.Fatalf("the whitened covariance (%d, %d) is %.6g, want %.6g", i, j, zcov.FloatAt(i*p+j), want)
}
}
}
// Unwhitening returns the observations themselves: the round trip
// closes exactly.
back, err := res.Unwhiten(z)
if err != nil {
t.Fatalf("Unwhiten: %v", err)
}
for i := range n {
for j := range p {
if math.Abs(back.FloatAt(i*p+j)-a.FloatAt(i*p+j)) > 1e-9 {
t.Fatalf("the round trip returned %.12g at (%d, %d), want the observation %.12g",
back.FloatAt(i*p+j), i, j, a.FloatAt(i*p+j))
}
}
}
}
// TestPCATransposedSpectrum pins the transform consistency across the
// transposed problem: for a square matrix centred along both axes the
// covariance of the rows and the covariance of the columns are the
// Gram pair XXᵀ and XᵀX, which share their spectrum exactly, and the
// decomposition must see the same eigenvalues from either side.
func TestPCATransposedSpectrum(t *testing.T) {
const n = 6
g := core.NewGenerator(37)
vals := make([]float64, 0, n*n)
for range n * n {
vals = append(vals, g.NormalUnit())
}
// Centre along the columns and then along the rows, so both
// readings of the matrix describe the same centred scatter.
for i := range n {
mean := 0.0
for j := range n {
mean += vals[i*n+j]
}
mean /= float64(n)
for j := range n {
vals[i*n+j] -= mean
}
}
for j := range n {
mean := 0.0
for i := range n {
mean += vals[i*n+j]
}
mean /= float64(n)
for i := range n {
vals[i*n+j] -= mean
}
}
a := mustFromFloats(t, vals, n, n)
transposed := make([]float64, 0, n*n)
for i := range n {
for j := range n {
transposed = append(transposed, vals[j*n+i])
}
}
at := mustFromFloats(t, transposed, n, n)
res, err := PCA(a)
if err != nil {
t.Fatalf("PCA: %v", err)
}
resT, err := PCA(at)
if err != nil {
t.Fatalf("PCA of the transposed data: %v", err)
}
for k := range n {
if math.Abs(res.ExplainedVariance[k]-resT.ExplainedVariance[k]) > 1e-8 {
t.Fatalf("eigenvalue %d: %.10g from the rows, %.10g from the columns",
k, res.ExplainedVariance[k], resT.ExplainedVariance[k])
}
}
}
// TestPCASignConvention pins the orientation rule directly on the
// helper: every eigenvector row turns so its largest-magnitude entry
// is positive, the first index winning a tie, and a row already
// oriented stays untouched.
func TestPCASignConvention(t *testing.T) {
// Row 0 ties at 0.6 across indices 0 and 1, index 0 negative: the
// first index wins, so the row flips. Row 1's largest entry is
// −0.9: it flips. Row 2's largest entry is 0.7: it stays.
v := [][]float64{
{-0.6, 0.6, 0.1},
{0.2, -0.9, 0.4},
{0.1, 0.7, -0.2},
}
fixEigenSigns(v)
want := [][]float64{
{0.6, -0.6, -0.1},
{-0.2, 0.9, -0.4},
{0.1, 0.7, -0.2},
}
for i := range 3 {
if !slices.Equal(v[i], want[i]) {
t.Fatalf("orientation wrong at row %d: %v, want %v", i, v[i], want[i])
}
}
// And on a real fit: every component's heaviest loading positive.
a := pcaFixture(t, 40, 31)
res, err := PCA(a)
if err != nil {
t.Fatalf("PCA: %v", err)
}
for k := range 4 {
worst, index := 0.0, 0
for i := range 4 {
if magnitude := math.Abs(res.Loadings.FloatAt(i*4 + k)); magnitude > worst {
worst = magnitude
index = i
}
}
if res.Loadings.FloatAt(index*4+k) < 0 {
t.Fatalf("component %d is oriented against the convention", k)
}
}
}
// TestPCAValidation refuses the inputs without a decomposition and
// withholds the whitening transforms where they do not exist.
func TestPCAValidation(t *testing.T) {
good := pcaFixture(t, 20, 31)
if _, err := PCA(mustFromFloats(t, []float64{1, 2, 3}, 3)); err == nil || !strings.Contains(err.Error(), "2-D array") {
t.Fatalf("a rank 1 array: got %v, want the rank refusal", err)
}
if _, err := PCA(mustFromFloats(t, []float64{1, 2}, 1, 2)); err == nil || !strings.Contains(err.Error(), "at least two observations") {
t.Fatalf("a single observation: got %v, want the observation floor refusal", err)
}
if _, err := PCA(core.New(core.Complex, 4, 2)); err == nil || !strings.Contains(err.Error(), "complex observations") {
t.Fatalf("complex observations: got %v, want the complex refusal", err)
}
if _, err := PCA(mustFromFloats(t, []float64{1, 2, math.NaN(), 4, 5, 6, 7, 8}, 4, 2)); err == nil || !strings.Contains(err.Error(), "non-finite") {
t.Fatalf("non-finite observations: got %v, want the non-finite refusal", err)
}
constant := mustFromFloats(t, []float64{1, 2, 1, 2, 1, 2, 1, 2}, 4, 2)
if _, err := PCA(constant); err == nil || !strings.Contains(err.Error(), "no variance") {
t.Fatalf("data with no variance: got %v, want the variance refusal", err)
}
// A duplicated column: the decomposition stands, the whitening
// transforms do not exist and are withheld.
singular := mustFromFloats(t, []float64{
1, 1, 2,
2, 2, 1,
3, 3, 0,
4, 4, 1,
5, 5, 2,
6, 6, 3,
}, 6, 3)
res, err := PCA(singular)
if err != nil {
t.Fatalf("PCA on a singular covariance: %v", err)
}
if res.Whitening != nil || res.Unwhitening != nil {
t.Fatalf("a rank-deficient fit published whitening transforms")
}
if _, err := res.Whiten(singular); err == nil || !strings.Contains(err.Error(), "rank deficient") {
t.Fatalf("Whiten on a rank-deficient fit: got %v, want the rank-deficiency refusal", err)
}
if _, err := res.Unwhiten(singular); err == nil || !strings.Contains(err.Error(), "rank deficient") {
t.Fatalf("Unwhiten on a rank-deficient fit: got %v, want the rank-deficiency refusal", err)
}
// Shape and content checks on the transforms.
fit, err := PCA(good)
if err != nil {
t.Fatalf("PCA: %v", err)
}
if _, err := fit.Whiten(mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6}, 3, 2)); err == nil || !strings.Contains(err.Error(), "columns, the fit") {
t.Fatalf("Whiten with a column mismatch: got %v, want the column refusal", err)
}
if _, err := fit.Whiten(mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6, 7, math.NaN(), 9, 10, 11, 12, 13, 14, 15, 16}, 4, 4)); err == nil || !strings.Contains(err.Error(), "non-finite") {
t.Fatalf("Whiten with non-finite observations: got %v, want the non-finite refusal", err)
}
if _, err := fit.Unwhiten(mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6}, 3, 2)); err == nil || !strings.Contains(err.Error(), "columns, the fit") {
t.Fatalf("Unwhiten with a column mismatch: got %v, want the column refusal", err)
}
var noFit *PCAResult
if _, err := noFit.Whiten(good); err == nil || !strings.Contains(err.Error(), "no fit to whiten") {
t.Fatalf("a nil fit whitened: got %v, want the nil-fit refusal", err)
}
if _, err := noFit.Unwhiten(good); err == nil || !strings.Contains(err.Error(), "no fit to unwhiten") {
t.Fatalf("a nil fit unwhitened: got %v, want the nil-fit refusal", err)
}
complexData := core.New(core.Complex, 4, 4)
if _, err := fit.Whiten(complexData); err == nil || !strings.Contains(err.Error(), "complex observations") {
t.Fatalf("Whiten with complex observations: got %v, want the complex refusal", err)
}
if _, err := fit.Unwhiten(complexData); err == nil || !strings.Contains(err.Error(), "complex observations") {
t.Fatalf("Unwhiten with complex observations: got %v, want the complex refusal", err)
}
// Integer observations reach the transforms through the widening
// accessor instead of a raw float payload, and whiten back out
// identically. Five rows of four generic columns keep the
// covariance full rank.
integers := mustFromInts(t, []int64{
1, 2, 3, 4,
2, 4, 6, 3,
3, 6, 2, 9,
4, 3, 8, 1,
5, 7, 1, 2,
}, 5, 4)
intFit, err := PCA(integers)
if err != nil {
t.Fatalf("PCA on integer observations: %v", err)
}
z, err := intFit.Whiten(integers)
if err != nil {
t.Fatalf("Whiten on integer observations: %v", err)
}
if _, err := intFit.Unwhiten(z); err != nil {
t.Fatalf("Unwhiten on integer observations: %v", err)
}
}