// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) // SPDX-License-Identifier: MIT package stats import ( "math" "slices" "strings" "testing" "sourcedock.dev/petrbalvin/tensor/internal/base" "sourcedock.dev/petrbalvin/tensor/internal/core" ) // orthonormalDesign builds an (n, p) design whose columns are // orthogonal in the (1/n)·XᵀX = I sense: every column has mean zero // and population standard deviation one. The Gram-Schmidt walk starts // from the constant vector and throws that direction away again: it // stands for the intercept the design's columns must be orthogonal // to, and what is left is exactly the space the slopes live in. func orthonormalDesign(t *testing.T, n, p int, seed int64) *core.Array { t.Helper() g := core.NewGenerator(seed) cols := make([][]float64, 0, p+1) for j := range p + 1 { v := make([]float64, n) if j == 0 { for i := range n { v[i] = 1 } } else { for i := range n { v[i] = g.NormalUnit() } } for _, c := range cols { dot := 0.0 for i := range n { dot += v[i] * c[i] } for i := range n { v[i] -= dot / float64(n) * c[i] } } norm := 0.0 for i := range n { norm += v[i] * v[i] } factor := math.Sqrt(float64(n) / norm) for i := range n { v[i] *= factor } cols = append(cols, v) } cols = cols[1:] vals := make([]float64, 0, n*p) for i := range n { for j := range p { vals = append(vals, cols[j][i]) } } return mustFromFloats(t, vals, n, p) } // TestLassoOrthonormalClosedForm pins the coordinate descent against // the closed form the orthonormal case admits: the columns decouple, // and every slope is the soft-thresholded projection of the centred // response on its own column, sign(ρ)·max(|ρ| − λ, 0), independent of // the other coordinates and of the iteration. func TestLassoOrthonormalClosedForm(t *testing.T) { n, p := 8, 4 design := orthonormalDesign(t, n, p, 5) // The precondition itself, checked rather than assumed: (1/n)XᵀX // is the identity and the column means are zero. for j := range p { mean := 0.0 for i := range n { mean += design.FloatAt(i*p + j) } if math.Abs(mean/float64(n)) > 1e-12 { t.Fatalf("column %d has mean %.3g, want 0", j, mean/float64(n)) } for k := j; k < p; k++ { dot := 0.0 for i := range n { dot += design.FloatAt(i*p+j) * design.FloatAt(i*p+k) } want := 0.0 if j == k { want = 1 } if math.Abs(dot/float64(n)-want) > 1e-12 { t.Fatalf("columns %d and %d have inner product %.12f, want %.12f", j, k, dot/float64(n), want) } } } y := mustFromFloats(t, []float64{3, -1, 4, 1, 5, -2, 0, 2}, n) res, err := Lasso(design, y, 0.6) if err != nil { t.Fatalf("Lasso: %v", err) } // The closed form, evaluated on the standardised system the fit // documents in its own result. yMean := 0.0 for i := range n { yMean += y.FloatAt(i) } yMean /= float64(n) for j := range p { mean := res.ColumnMeans[j] scale := res.ColumnScales[j] rho := 0.0 for i := range n { z := (design.FloatAt(i*p+j) - mean) / scale rho += z * (y.FloatAt(i) - yMean) } rho /= float64(n) // The threshold is spelled inline, not borrowed from // softThreshold: a helper on both sides of the comparison would // be wrong together. sign(ρ)·max(|ρ|−λ, 0) is the closed form. want := 0.0 if rho > 0.6 { want = rho - 0.6 } else if rho < -0.6 { want = rho + 0.6 } if math.Abs(res.Coefficients[j]-want) > 1e-9 { t.Fatalf("coefficient %d = %.12f, want the soft threshold %.12f", j, res.Coefficients[j], want) } } // The threshold itself, on literals: the shrinkage keeps the sign, // shrinks by exactly λ and floors at zero. for _, c := range []struct { v, lambda, want float64 }{ {0.0, 0.5, 0.0}, {0.25, 0.5, 0.0}, {0.5, 0.5, 0.0}, {1.5, 0.5, 1.0}, {-0.5, 0.5, 0.0}, {-1.5, 0.5, -1.0}, {2.5, 0.5, 2.0}, {-2.5, 0.5, -2.0}, } { if got := softThreshold(c.v, c.lambda); got != c.want { t.Fatalf("softThreshold(%g, %g) = %g, want %g", c.v, c.lambda, got, c.want) } } // The intercept keeps the unpenalised identity: the fit passes // through the column means. predicted := res.Intercept for j := range p { predicted += res.Coefficients[j] * res.ColumnMeans[j] } if math.Abs(predicted-yMean) > 1e-12 { t.Fatalf("the intercept identity broke: %.12g at the column means, want ȳ = %.12g", predicted, yMean) } if !res.Converged { t.Fatalf("the orthonormal fit did not report convergence") } } // TestElasticNetRidgeAgreement pins the alpha = 0 degeneration: with // the L1 term gone the coordinate descent is solving the Tikhonov // ridge system (ZᵀZ/n + λI)β = Zᵀ(y − ȳ)/n on the standardised // design, and the fit must agree with the answer the shared LU solve // delivers for that system, the same solve LinearRegression runs. func TestElasticNetRidgeAgreement(t *testing.T) { const ( n = 40 p = 5 lambda = 0.7 ) g := core.NewGenerator(11) x := make([]float64, 0, n*p) yv := make([]float64, 0, n) betaTrue := []float64{2, -1, 0.5, 0, 3} for range n { factor := g.NormalUnit() row := make([]float64, p) fitted := 1.0 for j := range p { row[j] = factor + 0.3*g.NormalUnit() fitted += betaTrue[j] * row[j] x = append(x, row[j]) } yv = append(yv, fitted+0.25*g.NormalUnit()) } design := mustFromFloats(t, x, n, p) y := mustFromFloats(t, yv, n) res, err := ElasticNet(design, y, lambda, 0) if err != nil { t.Fatalf("ElasticNet: %v", err) } // The standardised system, rebuilt in the test from the // standardisation the result records. yMean := 0.0 for _, v := range yv { yMean += v } yMean /= float64(n) z := make([][]float64, p) zty := make([]float64, p) for j := range p { z[j] = make([]float64, n) for i := range n { z[j][i] = (x[i*p+j] - res.ColumnMeans[j]) / res.ColumnScales[j] zty[j] += z[j][i] * (yv[i] - yMean) } zty[j] /= float64(n) } normal := make([][]float64, p) for j := range p { normal[j] = make([]float64, p) for k := range p { for i := range n { normal[j][k] += z[j][i] * z[k][i] } normal[j][k] /= float64(n) } normal[j][j] += lambda } solved, err := base.SolveSystem("ridgeReference", normal, [][]float64{zty}) if err != nil { t.Fatalf("the reference ridge solve failed: %v", err) } worst := 0.0 for j := range p { want := solved[0][j] / res.ColumnScales[j] if math.Abs(res.Coefficients[j]-want) > 1e-7 { t.Fatalf("ridge coefficient %d = %.10f, want the Tikhonov answer %.10f", j, res.Coefficients[j], want) } if d := math.Abs(res.Coefficients[j] - want); d > worst { worst = d } } wantIntercept := yMean for j := range p { wantIntercept -= solved[0][j] / res.ColumnScales[j] * res.ColumnMeans[j] } if math.Abs(res.Intercept-wantIntercept) > 1e-7 { t.Fatalf("ridge intercept = %.10f, want %.10f", res.Intercept, wantIntercept) } t.Logf("alpha = 0 agrees with the LU ridge solve to %.3g", worst) } // TestLassoPathSameGrid pins the shared path: the alpha = 0 ridge // walks exactly the lambda grid the pure lasso defines, so the two // fits are comparable coefficient for coefficient along it. func TestLassoPathSameGrid(t *testing.T) { g := core.NewGenerator(3) const n, p = 30, 4 x := make([]float64, 0, n*p) yv := make([]float64, 0, n) for range n { fitted := 2.0 for j := range p { v := g.NormalUnit() fitted += float64(p-j) * v x = append(x, v) } yv = append(yv, fitted+0.5*g.NormalUnit()) } design := mustFromFloats(t, x, n, p) y := mustFromFloats(t, yv, n) lasso, err := LassoPath(design, y, 1) if err != nil { t.Fatalf("LassoPath: %v", err) } ridge, err := LassoPath(design, y, 0) if err != nil { t.Fatalf("LassoPath: %v", err) } if !slices.Equal(lasso.Lambdas, ridge.Lambdas) { t.Fatalf("the ridge path left the lasso grid") } // Grid shape: descending, log-spaced, spanning the documented // three decades. for k := 1; k < len(lasso.Lambdas); k++ { if lasso.Lambdas[k] >= lasso.Lambdas[k-1] { t.Fatalf("the grid is not descending at %d: %g then %g", k, lasso.Lambdas[k-1], lasso.Lambdas[k]) } } ratio := lasso.Lambdas[1] / lasso.Lambdas[0] if math.Abs(ratio-math.Pow(lassoGridRatio, 1.0/float64(lassoGridSteps-1))) > 1e-12 { t.Fatalf("the grid is not log-spaced: consecutive ratio %.12g", ratio) } if math.Abs(lasso.Lambdas[len(lasso.Lambdas)-1]/lasso.Lambdas[0]-lassoGridRatio) > 1e-9 { t.Fatalf("the grid spans %.6g decades of ratio, want %.6g", lasso.Lambdas[len(lasso.Lambdas)-1]/lasso.Lambdas[0], lassoGridRatio) } } // TestLassoSparseSupportRecovery recovers the support of a sparse // true model from seeded generated data: along the path there is a // lambda interval where the nonzero set is exactly the true one and // the estimates sit close to the truth, while the top of the grid is // exactly the all-zero answer it promises. func TestLassoSparseSupportRecovery(t *testing.T) { const n, p = 300, 15 g := core.NewGenerator(7) betaTrue := make([]float64, p) betaTrue[2], betaTrue[7], betaTrue[11] = 1.5, -2.0, 0.8 x := make([]float64, 0, n*p) yv := make([]float64, 0, n) for range n { fitted := 3.0 for j := range p { v := g.NormalUnit() fitted += betaTrue[j] * v x = append(x, v) } yv = append(yv, fitted+0.5*g.NormalUnit()) } design := mustFromFloats(t, x, n, p) y := mustFromFloats(t, yv, n) path, err := LassoPath(design, y, 1) if err != nil { t.Fatalf("LassoPath: %v", err) } if !path.Converged { t.Fatalf("the path did not converge within its budget") } // The top of the grid: every slope exactly zero, the documented // meaning of lambdaMax. for j := range p { if path.Coefficients[0][j] != 0 { t.Fatalf("coefficient %d = %g at the top of the grid, want exactly 0", j, path.Coefficients[0][j]) } } // Somewhere along the path the support is the true one. trueSet := []int{2, 7, 11} recovered := false for k := range path.Lambdas { support := []int{} for j := range p { if path.Coefficients[k][j] != 0 { support = append(support, j) } } if !slices.Equal(support, trueSet) { continue } recovered = true closeEnough := true for _, j := range trueSet { if math.Abs(path.Coefficients[k][j]-betaTrue[j]) > 0.2 { closeEnough = false } } if closeEnough { t.Logf("support recovered at lambda[%d] = %.4g, coefficients within %.3f of the truth", k, path.Lambdas[k], 0.2) break } recovered = false } if !recovered { t.Fatalf("no lambda on the path recovered the true support {2, 7, 11}") } } // TestLassoPathWarmStart measures the warm start: the same grid fitted // cold, one ElasticNet call per lambda from zero every time, must // spend more coordinate cycles than the warm path, and by the bottom // of the grid the difference is the whole point of the path. func TestLassoPathWarmStart(t *testing.T) { const n, p = 100, 10 g := core.NewGenerator(13) betaTrue := make([]float64, p) betaTrue[1], betaTrue[4], betaTrue[8] = 1.2, -1.6, 0.9 x := make([]float64, 0, n*p) yv := make([]float64, 0, n) for range n { factor := g.NormalUnit() fitted := 1.0 for j := range p { v := factor + 0.5*g.NormalUnit() fitted += betaTrue[j] * v x = append(x, v) } yv = append(yv, fitted+0.4*g.NormalUnit()) } design := mustFromFloats(t, x, n, p) y := mustFromFloats(t, yv, n) warm, err := LassoPath(design, y, 1) if err != nil { t.Fatalf("LassoPath: %v", err) } cold := make([]int, len(warm.Lambdas)) totalWarm, totalCold := 0, 0 for k, lambda := range warm.Lambdas { fit, err := ElasticNet(design, y, lambda, 1) if err != nil { t.Fatalf("ElasticNet: %v", err) } cold[k] = fit.Iterations totalWarm += warm.Iterations[k] totalCold += cold[k] if warm.Iterations[k] > cold[k] { t.Fatalf("the warm start lost to the cold fit at lambda[%d]: %d cycles against %d", k, warm.Iterations[k], cold[k]) } } t.Logf("warm path %d cycles against cold %d; at the smallest lambda %d against %d", totalWarm, totalCold, warm.Iterations[len(cold)-1], cold[len(cold)-1]) if totalWarm >= totalCold { t.Fatalf("the warm path spent %d cycles, the cold path only %d", totalWarm, totalCold) } if warm.Iterations[len(cold)-1] >= cold[len(cold)-1] { t.Fatalf("at the smallest lambda the warm start spent %d cycles against the cold %d", warm.Iterations[len(cold)-1], cold[len(cold)-1]) } } // TestLassoDuplicateColumnResolve pins the collinear resolution: two // identical columns make the optimum non-unique, the effect free to // sit anywhere on the face the pair spans. The coordinate descent // settles on that face deterministically, so identical inputs give // bit-identical coefficients, and the fit is the single-copy answer // in everything a consumer can measure: the fitted values, the // summed effect of the pair and the residual structure. func TestLassoDuplicateColumnResolve(t *testing.T) { const n = 60 g := core.NewGenerator(17) x0 := make([]float64, n) x2 := make([]float64, n) yv := make([]float64, n) for i := range n { x0[i] = g.NormalUnit() x2[i] = g.NormalUnit() yv[i] = 1 + 2*x0[i] + 0.5*x2[i] + 0.3*g.NormalUnit() } dupVals := make([]float64, 0, 3*n) singleVals := make([]float64, 0, 2*n) for i := range n { dupVals = append(dupVals, x0[i], x0[i], x2[i]) singleVals = append(singleVals, x0[i], x2[i]) } dup := mustFromFloats(t, dupVals, n, 3) single := mustFromFloats(t, singleVals, n, 2) y := mustFromFloats(t, yv, n) res, err := ElasticNet(dup, y, 0.02, 1) if err != nil { t.Fatalf("ElasticNet: %v", err) } if !res.Converged { t.Fatalf("the duplicated design did not converge") } for j := range 3 { if math.IsNaN(res.Coefficients[j]) || math.IsInf(res.Coefficients[j], 0) { t.Fatalf("coefficient %d diverged to %g", j, res.Coefficients[j]) } } // Determinism: a second, identical fit lands on the same bits. again, err := ElasticNet(dup, y, 0.02, 1) if err != nil { t.Fatalf("the repeated ElasticNet failed: %v", err) } if !slices.Equal(res.Coefficients, again.Coefficients) { t.Fatalf("the duplicated design settled differently on a repeat: %v against %v", res.Coefficients, again.Coefficients) } // The degenerate face carries the single-copy effect: the pair // sums to it, and the third column agrees with its own fit. ref, err := ElasticNet(single, y, 0.02, 1) if err != nil { t.Fatalf("ElasticNet on the single-copy design: %v", err) } if math.Abs(res.Coefficients[0]+res.Coefficients[1]-ref.Coefficients[0]) > 1e-6 { t.Fatalf("the duplicate pair summed to %g, the single copy took %g", res.Coefficients[0]+res.Coefficients[1], ref.Coefficients[0]) } if math.Abs(res.Coefficients[2]-ref.Coefficients[1]) > 1e-6 { t.Fatalf("the independent column moved from %g to %g under the duplicate", ref.Coefficients[1], res.Coefficients[2]) } for i := range n { if math.Abs(res.Fitted[i]-ref.Fitted[i]) > 1e-6 { t.Fatalf("the duplicate changed fitted value %d: %.10g against %.10g", i, res.Fitted[i], ref.Fitted[i]) } } t.Logf("duplicate pair resolved deterministically as (%g, %g), the single-copy effect being %g", res.Coefficients[0], res.Coefficients[1], ref.Coefficients[0]) } // TestLassoRefitsSmoke walks the mixed alphas through one small fit // each, so the whole alpha range shares one code path and none of it // is only exercised by the pins above. func TestLassoAlphaRangeSmoke(t *testing.T) { design := orthonormalDesign(t, 12, 3, 21) y := mustFromFloats(t, []float64{2, -1, 3, 0, 1, -2, 4, 1, 0, -1, 2, 3}, 12) for _, alpha := range []float64{0, 0.25, 0.5, 0.75, 1} { res, err := ElasticNet(design, y, 0.3, alpha) if err != nil { t.Fatalf("ElasticNet alpha %g: %v", alpha, err) } if !res.Converged { t.Fatalf("ElasticNet alpha %g did not converge", alpha) } for j := range 3 { if math.IsNaN(res.Coefficients[j]) { t.Fatalf("ElasticNet alpha %g produced a NaN at %d", alpha, j) } } } } // TestLassoOnIntegerArrays exercises the widening accessor's fallback // paths in the standardisation and the final sweep: an integer design // and response reach the fit through FloatAt rather than a raw float // payload, and the answer matches the widened floats exactly. func TestLassoOnIntegerArrays(t *testing.T) { design := mustFromInts(t, []int64{0, 0, 1, 1, 2, 0, 3, 1, 4, 0, 5, 1, 6, 0}, 7, 2) y := mustFromInts(t, []int64{2, 4, 5, 7, 8, 10, 11}, 7) res, err := Lasso(design, y, 1e-5) if err != nil { t.Fatalf("Lasso on integer input: %v", err) } if !res.Converged { t.Fatalf("the integer-input fit did not converge") } if math.Abs(res.Coefficients[0]-1.5) > 1e-4 || math.Abs(res.Coefficients[1]-0.5) > 1e-4 { t.Fatalf("the integer-input fit is (%.6f, %.6f), want (1.5, 0.5)", res.Coefficients[0], res.Coefficients[1]) } widened := mustFromFloats(t, []float64{0, 0, 1, 1, 2, 0, 3, 1, 4, 0, 5, 1, 6, 0}, 7, 2) yw := mustFromFloats(t, []float64{2, 4, 5, 7, 8, 10, 11}, 7) reference, err := Lasso(widened, yw, 1e-5) if err != nil { t.Fatalf("Lasso on the widened floats: %v", err) } if !slices.Equal(res.Coefficients, reference.Coefficients) { t.Fatalf("the integer input fit %v against the float input %v", res.Coefficients, reference.Coefficients) } } // TestLassoInputValidation refuses every malformed input the fit // cannot answer, naming the condition in each case. func TestLassoInputValidation(t *testing.T) { good := mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6, 7, 8}, 4, 2) resp := mustFromFloats(t, []float64{1, 2, 3, 4}, 4) constant := mustFromFloats(t, []float64{1, 1, 1, 1, 1, 2, 1, 3}, 4, 2) nanY := mustFromFloats(t, []float64{1, 2, math.NaN(), 4}, 4) if _, err := ElasticNet(mustFromFloats(t, []float64{1, 2, 3, 4}, 4), resp, 0.1, 1); err == nil || !strings.Contains(err.Error(), "must be rank 2") { t.Fatalf("a rank 1 design: got %v, want the rank refusal", err) } if _, err := ElasticNet(good, mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6}, 3, 2), 0.1, 1); err == nil || !strings.Contains(err.Error(), "must be rank 1") { t.Fatalf("a rank 2 response: got %v, want the rank refusal", err) } if _, err := ElasticNet(good, mustFromFloats(t, []float64{1, 2, 3}, 3), 0.1, 1); err == nil || !strings.Contains(err.Error(), "rows but the response") { t.Fatalf("a row mismatch: got %v, want the length refusal", err) } if _, err := ElasticNet(mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6, 7, 8}, 4, 2), nanY, 0.1, 1); err == nil || !strings.Contains(err.Error(), "non-finite") { t.Fatalf("a non-finite response: got %v, want the non-finite refusal", err) } if _, err := ElasticNet(constant, resp, 0.1, 1); err == nil || !strings.Contains(err.Error(), "cannot be standardised") { t.Fatalf("a constant design column: got %v, want the standardisation refusal", err) } if _, err := ElasticNet(good, resp, -0.5, 1); err == nil || !strings.Contains(err.Error(), "lambda") { t.Fatalf("a negative lambda: got %v, want the lambda refusal", err) } if _, err := ElasticNet(good, resp, 0.1, 1.5); err == nil || !strings.Contains(err.Error(), "alpha") { t.Fatalf("an alpha above 1: got %v, want the alpha refusal", err) } if _, err := ElasticNet(good, resp, 0.1, -0.1); err == nil || !strings.Contains(err.Error(), "alpha") { t.Fatalf("a negative alpha: got %v, want the alpha refusal", err) } if _, err := ElasticNet(mustFromFloats(t, []float64{1, 2}, 1, 2), mustFromFloats(t, []float64{1}, 1), 0.1, 1); err == nil || !strings.Contains(err.Error(), "at least two observations") { t.Fatalf("a single observation: got %v, want the observation floor refusal", err) } complexDesign := core.New(core.Complex, 4, 2) if _, err := ElasticNet(complexDesign, resp, 0.1, 1); err == nil || !strings.Contains(err.Error(), "complex") { t.Fatalf("complex input: got %v, want the complex refusal", err) } if _, err := LassoPath(good, resp, 2); err == nil || !strings.Contains(err.Error(), "alpha") { t.Fatalf("LassoPath, an alpha above 1: got %v, want the alpha refusal", err) } }