// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) // SPDX-License-Identifier: MIT package grad import ( "math" "testing" "sourcedock.dev/petrbalvin/tensor/internal/core" ) // gaussianLogDensity builds the log density of independent standard // normals: log π(q) = −‖q‖²/2. func gaussianLogDensity(q *Tensor) (*Tensor, error) { sq, err := q.Mul(q) if err != nil { return nil, err } total, err := sq.Sum() if err != nil { return nil, err } return total.Scale(-0.5) } // gammaLogDensity builds log π(x) = log x − x, the unnormalised log // density of a Gamma(2, 1) distribution. func gammaLogDensity(q *Tensor) (*Tensor, error) { logq, err := q.Log() if err != nil { return nil, err } shifted, err := logq.Sub(q) if err != nil { return nil, err } return shifted.Sum() } // TestSampleHMCNormal runs the chain on a two-dimensional standard // normal and pins the sample moments: the target's mean is zero, its // variance one and its components independent. A fixed seed makes the // draw deterministic, so the bounds are checked facts about this run, // not hopes about a random one. func TestSampleHMCNormal(t *testing.T) { q0, err := core.FromFloats([]float64{2, -2}, 2) if err != nil { t.Fatalf("FromFloats: %v", err) } samples, err := SampleHMC(gaussianLogDensity, q0, HMCOptions{Step: 0.3, Steps: 20, BurnIn: 500, Samples: 6000, Thin: 1, Seed: 42}) if err != nil { t.Fatalf("SampleHMC: %v", err) } if got := samples.Shape(); got[0] != 6000 || got[1] != 2 { t.Fatalf("samples shape %v, want [6000 2]", got) } rows := samples.Shape()[0] mean := []float64{0, 0} variance := []float64{0, 0} for r := range rows { for c := range 2 { v := samples.FloatAt(r*2 + c) mean[c] += v / float64(rows) } } for r := range rows { for c := range 2 { d := samples.FloatAt(r*2+c) - mean[c] variance[c] += d * d / float64(rows) } } for c := range 2 { if math.Abs(mean[c]) > 0.15 { t.Fatalf("mean[%d] = %.4g, want |mean| ≤ 0.15", c, mean[c]) } if math.Abs(variance[c]-1) > 0.2 { t.Fatalf("variance[%d] = %.4g, want 1 ± 0.2", c, variance[c]) } } // Cross moment of the independent components. cov := 0.0 for r := range rows { cov += (samples.FloatAt(r*2) - mean[0]) * (samples.FloatAt(r*2+1) - mean[1]) / float64(rows) } if math.Abs(cov) > 0.15 { t.Fatalf("cross moment = %.4g, want |cov| ≤ 0.15", cov) } } // TestSampleHMCDeterministic checks the seed contract: the same seed // replays bit-identically, a different seed does not. func TestSampleHMCDeterministic(t *testing.T) { q0, _ := core.FromFloats([]float64{1}, 1) run := func(seed int64) *core.Array { samples, err := SampleHMC(gaussianLogDensity, q0, HMCOptions{Step: 0.4, Steps: 16, BurnIn: 100, Samples: 200, Seed: seed}) if err != nil { t.Fatalf("SampleHMC: %v", err) } return samples } a, b := run(7), run(7) c := run(8) for i := range a.Len() { if a.FloatAt(i) != b.FloatAt(i) { t.Fatalf("the same seed produced different samples at %d", i) } } same := true for i := range a.Len() { if a.FloatAt(i) != c.FloatAt(i) { same = false break } } if same { t.Fatal("different seeds produced identical samples") } } // TestSampleHMCSupportedDensity samples a Gamma(2, 1) target, // log π(x) = log x − x on x > 0. Proposals that overshoot into the // forbidden half-line yield a NaN density and are rejected, so every // kept sample stays positive and the mean approaches 2. func TestSampleHMCSupportedDensity(t *testing.T) { q0, _ := core.FromFloats([]float64{1}, 1) samples, err := SampleHMC(gammaLogDensity, q0, HMCOptions{Step: 0.3, Steps: 10, BurnIn: 500, Samples: 6000, Seed: 3}) if err != nil { t.Fatalf("SampleHMC: %v", err) } sum := 0.0 for i := range samples.Len() { x := samples.FloatAt(i) if x <= 0 { t.Fatalf("sample %d = %g left the support", i, x) } sum += x } mean := sum / float64(samples.Len()) if math.Abs(mean-2) > 0.15 { t.Fatalf("Gamma(2) mean = %.4g, want 2 ± 0.15", mean) } } // TestSampleHMCRejectedByGuard exercises the explicit rejection path: // the density errors outside its support instead of returning NaN, // and the chain still stays inside it. func TestSampleHMCRejectedByGuard(t *testing.T) { q0, _ := core.FromFloats([]float64{0.5}, 1) samples, err := SampleHMC(func(q *Tensor) (*Tensor, error) { x := q.Data().FloatAt(0) if x <= 0 { return nil, errf("outside the support") } return gammaLogDensity(q) }, q0, HMCOptions{Step: 0.5, Steps: 20, BurnIn: 300, Samples: 2000, Seed: 5}) if err != nil { t.Fatalf("SampleHMC: %v", err) } for i := range samples.Len() { if samples.FloatAt(i) <= 0 { t.Fatalf("sample %d = %g left the support", i, samples.FloatAt(i)) } } } // TestSampleHMCErrors pins the validation contract, including a // density that fails at the start, returns a non-scalar, or never // touches the leaf core. func TestSampleHMCErrors(t *testing.T) { q0, _ := core.FromFloats([]float64{1}, 1) if _, err := SampleHMC(nil, q0, HMCOptions{Step: 0.1, Steps: 5, Samples: 1}); err == nil { t.Fatal("expected an error for a nil density") } rank2, _ := core.FromFloats([]float64{1, 1}, 1, 2) if _, err := SampleHMC(gaussianLogDensity, rank2, HMCOptions{Step: 0.1, Steps: 5, Samples: 1}); err == nil { t.Fatal("expected an error for a rank-2 state") } if _, err := SampleHMC(gaussianLogDensity, nil, HMCOptions{Step: 0.1, Steps: 5, Samples: 1}); err == nil { t.Fatal("expected an error for a nil state") } if _, err := SampleHMC(gaussianLogDensity, q0, HMCOptions{Steps: 5, Samples: 1}); err == nil { t.Fatal("expected an error for a non-positive step") } if _, err := SampleHMC(gaussianLogDensity, q0, HMCOptions{Step: 0.1, Samples: 1}); err == nil { t.Fatal("expected an error for a non-positive step count") } if _, err := SampleHMC(gaussianLogDensity, q0, HMCOptions{Step: 0.1, Steps: 5}); err == nil { t.Fatal("expected an error for a non-positive sample count") } if _, err := SampleHMC(gaussianLogDensity, q0, HMCOptions{Step: 0.1, Steps: 5, Samples: 1, BurnIn: -3}); err == nil { t.Fatal("expected an error for a negative BurnIn") } nonScalar := func(q *Tensor) (*Tensor, error) { two, _ := core.FromFloats([]float64{1, 2}, 2) return FromArray(two, false), nil } if _, err := SampleHMC(nonScalar, q0, HMCOptions{Step: 0.1, Steps: 5, Samples: 1}); err == nil { t.Fatal("expected an error for a non-scalar density") } detached := func(q *Tensor) (*Tensor, error) { one, _ := core.FromFloats([]float64{1}, 1) return FromArray(one, false), nil } if _, err := SampleHMC(detached, q0, HMCOptions{Step: 0.1, Steps: 5, Samples: 1}); err == nil { t.Fatal("expected an error for a density disconnected from the leaf") } failsAtStart := func(q *Tensor) (*Tensor, error) { return nil, errf("no density at the start") } if _, err := SampleHMC(failsAtStart, q0, HMCOptions{Step: 0.1, Steps: 5, Samples: 1}); err == nil { t.Fatal("expected the start-time density error to be fatal") } }