generated from mathias/template-go-web
internal/eval: three pure-Go diagnostics on frozen embeddings: LinearProbe(emb, y, λ) → val_vol_r2 (OOS R², closed-form ridge, Cholesky) Silhouette(emb, labels) → mean silhouette (Euclidean, multi-label, errors on <2 classes) EffectiveRank(emb) → Roy effective rank (Jacobi eigenvalues → entropy → exp(H)) cmd/eval/main.go: CLI driver reading embeddings.json (exported by train.py with EXPORT_EMBEDDINGS=1), standardises per-dim, dispatches to -metric flag. task eval:probe / eval:silhouette / eval:collapse wired in Taskfile. 8/8 tests pass (red-green: perfect clusters, rank-1, full-rank, noise, constant target, single-label error). Pure stdlib, no external deps. Closes #4. Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
139 lines
3.4 KiB
Go
139 lines
3.4 KiB
Go
package eval_test
|
|
|
|
import (
|
|
"math"
|
|
"math/rand"
|
|
"testing"
|
|
|
|
"gitea.d-ma.be/mathias/jepa-fx-risk/internal/eval"
|
|
)
|
|
|
|
func seededRNG(seed int64) *rand.Rand {
|
|
return rand.New(rand.NewSource(seed))
|
|
}
|
|
|
|
// ── LinearProbe (val_vol_r2) ──────────────────────────────────────────────────
|
|
|
|
func TestLinearProbe_Perfect(t *testing.T) {
|
|
n := 50
|
|
emb := make([][]float64, n)
|
|
y := make([]float64, n)
|
|
for i := range emb {
|
|
emb[i] = []float64{float64(i)}
|
|
y[i] = float64(i)
|
|
}
|
|
r2 := eval.LinearProbe(emb, y, 1e-3)
|
|
if r2 < 0.99 {
|
|
t.Fatalf("perfect predictor: want R²≥0.99, got %.4f", r2)
|
|
}
|
|
}
|
|
|
|
func TestLinearProbe_ConstantTarget(t *testing.T) {
|
|
n := 40
|
|
emb := make([][]float64, n)
|
|
y := make([]float64, n)
|
|
for i := range emb {
|
|
emb[i] = []float64{float64(i), float64(i * i)}
|
|
y[i] = 3.0
|
|
}
|
|
r2 := eval.LinearProbe(emb, y, 1e-3)
|
|
if r2 > 0.01 {
|
|
t.Fatalf("constant target: want R²≤0.01, got %.4f", r2)
|
|
}
|
|
}
|
|
|
|
func TestLinearProbe_NoiseEmbedding(t *testing.T) {
|
|
rng := seededRNG(42)
|
|
n := 80
|
|
emb := make([][]float64, n)
|
|
y := make([]float64, n)
|
|
for i := range emb {
|
|
emb[i] = []float64{rng.NormFloat64(), rng.NormFloat64()}
|
|
y[i] = float64(i)
|
|
}
|
|
r2 := eval.LinearProbe(emb, y, 1e-3)
|
|
if r2 > 0.10 {
|
|
t.Fatalf("noise embedding: want R²<0.10, got %.4f", r2)
|
|
}
|
|
}
|
|
|
|
// ── Silhouette ────────────────────────────────────────────────────────────────
|
|
|
|
func TestSilhouette_PerfectClusters(t *testing.T) {
|
|
emb := make([][]float64, 40)
|
|
labels := make([]int, 40)
|
|
for i := range emb {
|
|
if i < 20 {
|
|
emb[i] = []float64{0.0, 0.0}
|
|
labels[i] = 0
|
|
} else {
|
|
emb[i] = []float64{1000.0, 1000.0}
|
|
labels[i] = 1
|
|
}
|
|
}
|
|
sil, err := eval.Silhouette(emb, labels)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if sil < 0.95 {
|
|
t.Fatalf("perfect clusters: want sil≥0.95, got %.4f", sil)
|
|
}
|
|
}
|
|
|
|
func TestSilhouette_SingleLabel(t *testing.T) {
|
|
emb := [][]float64{{1, 2}, {3, 4}, {5, 6}}
|
|
labels := []int{0, 0, 0}
|
|
_, err := eval.Silhouette(emb, labels)
|
|
if err == nil {
|
|
t.Fatal("expected error for single-label input")
|
|
}
|
|
}
|
|
|
|
func TestSilhouette_RandomClusters(t *testing.T) {
|
|
rng := seededRNG(7)
|
|
n := 60
|
|
emb := make([][]float64, n)
|
|
labels := make([]int, n)
|
|
for i := range emb {
|
|
emb[i] = []float64{rng.NormFloat64(), rng.NormFloat64()}
|
|
labels[i] = i % 2
|
|
}
|
|
sil, err := eval.Silhouette(emb, labels)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if math.Abs(sil) > 0.30 {
|
|
t.Fatalf("random clusters: want |sil|≤0.30, got %.4f", sil)
|
|
}
|
|
}
|
|
|
|
// ── EffectiveRank ─────────────────────────────────────────────────────────────
|
|
|
|
func TestEffectiveRank_Rank1(t *testing.T) {
|
|
emb := make([][]float64, 30)
|
|
for i := range emb {
|
|
emb[i] = []float64{1.0, 2.0, 3.0, 4.0}
|
|
}
|
|
er := eval.EffectiveRank(emb)
|
|
if er > 1.5 {
|
|
t.Fatalf("rank-1 matrix: want erank≤1.5, got %.4f", er)
|
|
}
|
|
}
|
|
|
|
func TestEffectiveRank_FullRank(t *testing.T) {
|
|
rng := seededRNG(99)
|
|
dim := 8
|
|
emb := make([][]float64, 200)
|
|
for i := range emb {
|
|
row := make([]float64, dim)
|
|
for j := range row {
|
|
row[j] = rng.NormFloat64()
|
|
}
|
|
emb[i] = row
|
|
}
|
|
er := eval.EffectiveRank(emb)
|
|
if er < float64(dim)*0.7 {
|
|
t.Fatalf("full-rank: want erank≥%.1f, got %.4f", float64(dim)*0.7, er)
|
|
}
|
|
}
|