From a24d042dbc86fe9da06750d2a9a57400ec3b6b38 Mon Sep 17 00:00:00 2001 From: Shravan Goswami Date: Sun, 18 Jan 2026 00:58:54 +0530 Subject: [PATCH 01/13] Add glue code for AbstractMCMC Callbacks support --- src/mcmc/Inference.jl | 3 + src/mcmc/callbacks.jl | 227 ++++++++++++++++++++++++++++++++++++++++ test/mcmc/callbacks.jl | 232 +++++++++++++++++++++++++++++++++++++++++ test/runtests.jl | 1 + 4 files changed, 463 insertions(+) create mode 100644 src/mcmc/callbacks.jl create mode 100644 test/mcmc/callbacks.jl diff --git a/src/mcmc/Inference.jl b/src/mcmc/Inference.jl index b02d0887f7..7690952c9e 100644 --- a/src/mcmc/Inference.jl +++ b/src/mcmc/Inference.jl @@ -130,4 +130,7 @@ include("prior.jl") include("gibbs.jl") include("gibbs_conditional.jl") +# AbstractMCMC callback interface +include("callbacks.jl") + end # module diff --git a/src/mcmc/callbacks.jl b/src/mcmc/callbacks.jl new file mode 100644 index 0000000000..44b3c8972c --- /dev/null +++ b/src/mcmc/callbacks.jl @@ -0,0 +1,227 @@ +using AbstractMCMC: AbstractMCMC + +function _get_lp(vi::DynamicPPL.AbstractVarInfo) + lp = DynamicPPL.getlogp(vi) + if lp isa NamedTuple + return sum(values(lp)) + end + return lp +end + +function _varinfo_params(vi::DynamicPPL.AbstractVarInfo) + vns = keys(vi) + return Iterators.flatmap(vns) do vn + val = DynamicPPL.getindex_internal(vi, vn) + if val isa AbstractArray + [string(vn, "[", i, "]") => v for (i, v) in enumerate(val)] + else + [string(vn) => val] + end + end +end + +### +### getparams - Extract named parameters from sampler states +### + +# HMCState - used by HMC, HMCDA, NUTS (contains vi field) +function AbstractMCMC.getparams(state::HMCState) + return collect(_varinfo_params(state.vi)) +end + +# MHState - contains varinfo field (different name!) +function AbstractMCMC.getparams(state::MHState) + return collect(_varinfo_params(state.varinfo)) +end + +# PGState - contains vi field +function AbstractMCMC.getparams(state::PGState) + return collect(_varinfo_params(state.vi)) +end + +# SMCState - particles contain VarInfo in their model +function AbstractMCMC.getparams(state::SMCState) + particle = state.particles.vals[state.particleindex] + vi = particle.model.f.varinfo + return collect(_varinfo_params(vi)) +end + +# GibbsState - contains global VarInfo +function AbstractMCMC.getparams(state::GibbsState) + return collect(_varinfo_params(state.vi)) +end + +# ESS uses VarInfo directly as state +function AbstractMCMC.getparams(state::DynamicPPL.AbstractVarInfo) + return collect(_varinfo_params(state)) +end + +# SGHMCState - contains params vector directly (no VarInfo) +function AbstractMCMC.getparams(state::SGHMCState) + return ["θ[$i]" => v for (i, v) in enumerate(state.params)] +end + +# SGLDState - contains params vector directly (no VarInfo) +function AbstractMCMC.getparams(state::SGLDState) + return ["θ[$i]" => v for (i, v) in enumerate(state.params)] +end + +### +### getstats - Extract extra statistics from sampler states +### + +# HMCState - rich stats available +function AbstractMCMC.getstats(state::HMCState) + lp = _get_lp(state.vi) + # Get step size from kernel + ϵ = try + state.kernel.τ.integrator.ϵ + catch + NaN + end + return (lp=lp, step_size=ϵ, iteration=state.i) +end + +# MHState +function AbstractMCMC.getstats(state::MHState) + return (lp=state.logjoint_internal,) +end + +# PGState +function AbstractMCMC.getstats(state::PGState) + lp = _get_lp(state.vi) + return (lp=lp,) +end + +# SMCState +function AbstractMCMC.getstats(state::SMCState) + return (logevidence=state.average_logevidence, particle_index=state.particleindex) +end + +# GibbsState +function AbstractMCMC.getstats(state::GibbsState) + lp = _get_lp(state.vi) + return (lp=lp,) +end + +# ESS (VarInfo as state) +function AbstractMCMC.getstats(state::DynamicPPL.AbstractVarInfo) + lp = _get_lp(state) + return (lp=lp,) +end + +# SGHMCState +function AbstractMCMC.getstats(state::SGHMCState) + lp = try + LogDensityProblems.logdensity(state.logdensity, state.params) + catch + NaN + end + return (lp=lp,) +end + +# SGLDState +function AbstractMCMC.getstats(state::SGLDState) + lp = try + LogDensityProblems.logdensity(state.logdensity, state.params) + catch + NaN + end + return (lp=lp, step=state.step) +end + +### +### getparams/getstats from transitions (ParamsWithStats) +### + +function AbstractMCMC.getparams(transition::DynamicPPL.ParamsWithStats) + # params is OrderedDict{VarName, Any} + return [string(vn) => val for (vn, val) in transition.params] +end + +function AbstractMCMC.getstats(transition::DynamicPPL.ParamsWithStats) + return transition.stats +end + +### +### hyperparam_metrics - Define TensorBoard hyperparam metrics +### + +function AbstractMCMC.hyperparam_metrics(model::DynamicPPL.Model, sampler::NUTS) + return [ + "extras/acceptance_rate/stat/Mean", + "extras/max_hamiltonian_energy_error/stat/Mean", + "extras/lp/stat/Mean", + "extras/n_steps/stat/Mean", + "extras/tree_depth/stat/Mean", + ] +end + +function AbstractMCMC.hyperparam_metrics(model::DynamicPPL.Model, sampler::Hamiltonian) + return [ + "extras/acceptance_rate/stat/Mean", + "extras/lp/stat/Mean", + "extras/n_steps/stat/Mean", + ] +end + +function AbstractMCMC.hyperparam_metrics(model::DynamicPPL.Model, sampler::MH) + return ["extras/lp/stat/Mean"] +end + +function AbstractMCMC.hyperparam_metrics(model::DynamicPPL.Model, sampler::PG) + return ["extras/lp/stat/Mean", "extras/logevidence/stat/Mean"] +end + +### +### _hyperparams_impl - Extract sampler hyperparameters +### + +function AbstractMCMC._hyperparams_impl( + model::DynamicPPL.Model, sampler::HMC, state; kwargs... +) + return ["epsilon" => sampler.ϵ, "n_leapfrog" => sampler.n_leapfrog] +end + +function AbstractMCMC._hyperparams_impl( + model::DynamicPPL.Model, sampler::HMCDA, state; kwargs... +) + return [ + "n_adapts" => sampler.n_adapts, + "delta" => sampler.δ, + "lambda" => sampler.λ, + "epsilon" => sampler.ϵ, + ] +end + +function AbstractMCMC._hyperparams_impl( + model::DynamicPPL.Model, sampler::NUTS, state; kwargs... +) + return [ + "n_adapts" => sampler.n_adapts, + "delta" => sampler.δ, + "max_depth" => sampler.max_depth, + "Delta_max" => sampler.Δ_max, + "epsilon" => sampler.ϵ, + ] +end + +function AbstractMCMC._hyperparams_impl( + model::DynamicPPL.Model, sampler::PG, state; kwargs... +) + return ["nparticles" => sampler.nparticles] +end + +function AbstractMCMC._hyperparams_impl( + model::DynamicPPL.Model, sampler::SGHMC, state; kwargs... +) + return [ + "learning_rate" => sampler.learning_rate, "momentum_decay" => sampler.momentum_decay + ] +end + +function AbstractMCMC._hyperparams_impl( + model::DynamicPPL.Model, sampler::SGLD, state; kwargs... +) + return ["stepsize" => string(sampler.stepsize)] +end diff --git a/test/mcmc/callbacks.jl b/test/mcmc/callbacks.jl new file mode 100644 index 0000000000..fe8f1cc05c --- /dev/null +++ b/test/mcmc/callbacks.jl @@ -0,0 +1,232 @@ +module TuringCallbacksTests + +using Test: @test, @testset +using Turing +using AbstractMCMC: AbstractMCMC +using Random: Random +using DynamicPPL: DynamicPPL + +using ..Models: gdemo_default + +@testset "AbstractMCMC Callbacks Interface" begin + @testset "getparams from states" begin + @testset "HMCState getparams" begin + Random.seed!(42) + chain = sample(gdemo_default, NUTS(100, 0.65), 10; progress=false) + + rng = Random.default_rng() + transition, state = AbstractMCMC.step( + rng, + gdemo_default, + NUTS(100, 0.65); + initial_params=Turing.Inference.init_strategy(NUTS(100, 0.65)), + ) + + params = AbstractMCMC.getparams(state) + @test params isa Vector + @test length(params) >= 2 + @test all(p -> p isa Pair{String,<:Any}, params) + end + + @testset "MHState getparams" begin + Random.seed!(42) + transition, state = AbstractMCMC.step( + Random.default_rng(), + gdemo_default, + MH(); + initial_params=Turing.Inference.init_strategy(MH()), + ) + + params = AbstractMCMC.getparams(state) + @test params isa Vector + @test length(params) >= 2 + @test all(p -> p isa Pair{String,<:Any}, params) + end + + @testset "ESS getparams (VarInfo as state)" begin + # ESS only works with Gaussian priors - need simple model + @model function gaussian_model() + m ~ Normal(0, 1) + return nothing + end + + Random.seed!(42) + transition, state = AbstractMCMC.step( + Random.default_rng(), + gaussian_model(), + ESS(); + initial_params=Turing.Inference.init_strategy(ESS()), + ) + + @test state isa DynamicPPL.AbstractVarInfo + params = AbstractMCMC.getparams(state) + @test params isa Vector + @test length(params) >= 1 + end + end + + @testset "getstats from states" begin + @testset "HMCState getstats" begin + Random.seed!(42) + transition, state = AbstractMCMC.step( + Random.default_rng(), + gdemo_default, + NUTS(100, 0.65); + initial_params=Turing.Inference.init_strategy(NUTS(100, 0.65)), + ) + + stats = AbstractMCMC.getstats(state) + @test stats isa NamedTuple + @test haskey(stats, :lp) + @test stats.lp isa Real + end + + @testset "MHState getstats" begin + Random.seed!(42) + transition, state = AbstractMCMC.step( + Random.default_rng(), + gdemo_default, + MH(); + initial_params=Turing.Inference.init_strategy(MH()), + ) + + stats = AbstractMCMC.getstats(state) + @test stats isa NamedTuple + @test haskey(stats, :lp) + @test stats.lp isa Real + end + + @testset "ESS getstats (VarInfo)" begin + @model function gaussian_model() + m ~ Normal(0, 1) + return nothing + end + + Random.seed!(42) + transition, state = AbstractMCMC.step( + Random.default_rng(), + gaussian_model(), + ESS(); + initial_params=Turing.Inference.init_strategy(ESS()), + ) + + stats = AbstractMCMC.getstats(state) + @test stats isa NamedTuple + @test haskey(stats, :lp) + @test stats.lp isa Real + end + end + + @testset "ParamsWithStats transitions" begin + Random.seed!(42) + transition, _ = AbstractMCMC.step( + Random.default_rng(), + gdemo_default, + NUTS(100, 0.65); + initial_params=Turing.Inference.init_strategy(NUTS(100, 0.65)), + ) + + @test transition isa DynamicPPL.ParamsWithStats + + stats = AbstractMCMC.getstats(transition) + @test stats isa NamedTuple + + params = AbstractMCMC.getparams(transition) + @test params isa Vector + end + + @testset "hyperparam_metrics" begin + @testset "NUTS hyperparam_metrics" begin + sampler = NUTS() + metrics = AbstractMCMC.hyperparam_metrics(gdemo_default, sampler) + @test metrics isa Vector{String} + @test "extras/lp/stat/Mean" in metrics + @test "extras/acceptance_rate/stat/Mean" in metrics + end + + @testset "HMC hyperparam_metrics (via Hamiltonian)" begin + sampler = HMC(0.1, 10) + metrics = AbstractMCMC.hyperparam_metrics(gdemo_default, sampler) + @test metrics isa Vector{String} + @test "extras/lp/stat/Mean" in metrics + end + + @testset "MH hyperparam_metrics" begin + sampler = MH() + metrics = AbstractMCMC.hyperparam_metrics(gdemo_default, sampler) + @test metrics isa Vector{String} + @test "extras/lp/stat/Mean" in metrics + end + + @testset "PG hyperparam_metrics" begin + sampler = PG(10) + metrics = AbstractMCMC.hyperparam_metrics(gdemo_default, sampler) + @test metrics isa Vector{String} + @test "extras/lp/stat/Mean" in metrics + end + end + + @testset "_hyperparams_impl" begin + @testset "HMC hyperparams" begin + sampler = HMC(0.1, 10) + hyperparams = AbstractMCMC._hyperparams_impl(gdemo_default, sampler, nothing) + @test hyperparams isa Vector + + hp_dict = Dict(hyperparams) + @test hp_dict["epsilon"] == 0.1 + @test hp_dict["n_leapfrog"] == 10 + end + + @testset "NUTS hyperparams" begin + sampler = NUTS(200, 0.65) + hyperparams = AbstractMCMC._hyperparams_impl(gdemo_default, sampler, nothing) + @test hyperparams isa Vector + + hp_dict = Dict(hyperparams) + @test hp_dict["n_adapts"] == 200 + @test hp_dict["delta"] == 0.65 + @test haskey(hp_dict, "max_depth") + end + + @testset "HMCDA hyperparams" begin + sampler = HMCDA(200, 0.65, 0.3) + hyperparams = AbstractMCMC._hyperparams_impl(gdemo_default, sampler, nothing) + @test hyperparams isa Vector + + hp_dict = Dict(hyperparams) + @test hp_dict["n_adapts"] == 200 + @test hp_dict["delta"] == 0.65 + @test hp_dict["lambda"] == 0.3 + end + + @testset "PG hyperparams" begin + sampler = PG(10) + hyperparams = AbstractMCMC._hyperparams_impl(gdemo_default, sampler, nothing) + @test hyperparams isa Vector + + hp_dict = Dict(hyperparams) + @test hp_dict["nparticles"] == 10 + end + + @testset "SGHMC hyperparams" begin + sampler = SGHMC(; learning_rate=0.01, momentum_decay=0.1) + hyperparams = AbstractMCMC._hyperparams_impl(gdemo_default, sampler, nothing) + @test hyperparams isa Vector + + hp_dict = Dict(hyperparams) + @test hp_dict["learning_rate"] == 0.01 + @test hp_dict["momentum_decay"] == 0.1 + end + + @testset "SGLD hyperparams" begin + sampler = SGLD() + hyperparams = AbstractMCMC._hyperparams_impl(gdemo_default, sampler, nothing) + @test hyperparams isa Vector + + hp_dict = Dict(hyperparams) + @test haskey(hp_dict, "stepsize") + end + end +end + +end # module diff --git a/test/runtests.jl b/test/runtests.jl index a4a60d8e29..2a6c00c312 100644 --- a/test/runtests.jl +++ b/test/runtests.jl @@ -44,6 +44,7 @@ end @testset "samplers (without AD)" verbose = true begin @timeit_include("mcmc/abstractmcmc.jl") + @timeit_include("mcmc/callbacks.jl") @timeit_include("mcmc/particle_mcmc.jl") @timeit_include("mcmc/emcee.jl") @timeit_include("mcmc/ess.jl") From efb745d4e20c08ef844aa33eca86dee5df22e06c Mon Sep 17 00:00:00 2001 From: Shravan Goswami Date: Sun, 18 Jan 2026 01:04:09 +0530 Subject: [PATCH 02/13] use AbstractMCMC@5.11 --- Project.toml | 2 +- test/Project.toml | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/Project.toml b/Project.toml index 4f489e96fc..38b4a28288 100644 --- a/Project.toml +++ b/Project.toml @@ -47,7 +47,7 @@ TuringDynamicHMCExt = "DynamicHMC" [compat] ADTypes = "1.9" -AbstractMCMC = "5.9" +AbstractMCMC = "5.11" AbstractPPL = "0.11, 0.12, 0.13" Accessors = "0.1" AdvancedHMC = "0.8.3" diff --git a/test/Project.toml b/test/Project.toml index 7db79ea151..be0b4f8bed 100644 --- a/test/Project.toml +++ b/test/Project.toml @@ -40,7 +40,7 @@ TimerOutputs = "a759f4b9-e2f1-59dc-863e-4aeb61b1ea8f" [compat] ADTypes = "1" -AbstractMCMC = "5.9" +AbstractMCMC = "5.11" AbstractPPL = "0.11, 0.12, 0.13" AdvancedMH = "0.8.9" AdvancedPS = "0.7" From 80caf446286000ea39d9a026ca21fa7a2490d7c8 Mon Sep 17 00:00:00 2001 From: Shravan Goswami Date: Fri, 23 Jan 2026 11:45:37 +0530 Subject: [PATCH 03/13] update Callbacks to incorporate suggestion from discussion --- src/mcmc/Inference.jl | 15 ++- src/mcmc/callbacks.jl | 227 ---------------------------------- src/mcmc/ess.jl | 13 ++ src/mcmc/gibbs.jl | 13 ++ src/mcmc/hmc.jl | 18 +++ src/mcmc/mh.jl | 12 ++ src/mcmc/particle_mcmc.jl | 25 ++++ src/mcmc/sghmc.jl | 32 +++++ test/mcmc/callbacks.jl | 250 +++++--------------------------------- 9 files changed, 158 insertions(+), 447 deletions(-) delete mode 100644 src/mcmc/callbacks.jl diff --git a/src/mcmc/Inference.jl b/src/mcmc/Inference.jl index 7690952c9e..859326f724 100644 --- a/src/mcmc/Inference.jl +++ b/src/mcmc/Inference.jl @@ -115,6 +115,18 @@ function mh_accept(logp_current::Real, logp_proposal::Real, log_proposal_ratio:: return log(rand()) + logp_current ≤ logp_proposal + log_proposal_ratio end +# Helper functions for AbstractMCMC callbacks +# Helper to get log probability from VarInfo +function _get_lp(vi::DynamicPPL.AbstractVarInfo) + lp = DynamicPPL.getlogp(vi) + return sum(values(lp)) +end + +# Helper to extract raw parameter values from VarInfo as Vector{<:Real} +function _get_params_vector(vi::DynamicPPL.AbstractVarInfo) + return vi[:] +end + ####################################### # Concrete algorithm implementations. # ####################################### @@ -130,7 +142,4 @@ include("prior.jl") include("gibbs.jl") include("gibbs_conditional.jl") -# AbstractMCMC callback interface -include("callbacks.jl") - end # module diff --git a/src/mcmc/callbacks.jl b/src/mcmc/callbacks.jl deleted file mode 100644 index 44b3c8972c..0000000000 --- a/src/mcmc/callbacks.jl +++ /dev/null @@ -1,227 +0,0 @@ -using AbstractMCMC: AbstractMCMC - -function _get_lp(vi::DynamicPPL.AbstractVarInfo) - lp = DynamicPPL.getlogp(vi) - if lp isa NamedTuple - return sum(values(lp)) - end - return lp -end - -function _varinfo_params(vi::DynamicPPL.AbstractVarInfo) - vns = keys(vi) - return Iterators.flatmap(vns) do vn - val = DynamicPPL.getindex_internal(vi, vn) - if val isa AbstractArray - [string(vn, "[", i, "]") => v for (i, v) in enumerate(val)] - else - [string(vn) => val] - end - end -end - -### -### getparams - Extract named parameters from sampler states -### - -# HMCState - used by HMC, HMCDA, NUTS (contains vi field) -function AbstractMCMC.getparams(state::HMCState) - return collect(_varinfo_params(state.vi)) -end - -# MHState - contains varinfo field (different name!) -function AbstractMCMC.getparams(state::MHState) - return collect(_varinfo_params(state.varinfo)) -end - -# PGState - contains vi field -function AbstractMCMC.getparams(state::PGState) - return collect(_varinfo_params(state.vi)) -end - -# SMCState - particles contain VarInfo in their model -function AbstractMCMC.getparams(state::SMCState) - particle = state.particles.vals[state.particleindex] - vi = particle.model.f.varinfo - return collect(_varinfo_params(vi)) -end - -# GibbsState - contains global VarInfo -function AbstractMCMC.getparams(state::GibbsState) - return collect(_varinfo_params(state.vi)) -end - -# ESS uses VarInfo directly as state -function AbstractMCMC.getparams(state::DynamicPPL.AbstractVarInfo) - return collect(_varinfo_params(state)) -end - -# SGHMCState - contains params vector directly (no VarInfo) -function AbstractMCMC.getparams(state::SGHMCState) - return ["θ[$i]" => v for (i, v) in enumerate(state.params)] -end - -# SGLDState - contains params vector directly (no VarInfo) -function AbstractMCMC.getparams(state::SGLDState) - return ["θ[$i]" => v for (i, v) in enumerate(state.params)] -end - -### -### getstats - Extract extra statistics from sampler states -### - -# HMCState - rich stats available -function AbstractMCMC.getstats(state::HMCState) - lp = _get_lp(state.vi) - # Get step size from kernel - ϵ = try - state.kernel.τ.integrator.ϵ - catch - NaN - end - return (lp=lp, step_size=ϵ, iteration=state.i) -end - -# MHState -function AbstractMCMC.getstats(state::MHState) - return (lp=state.logjoint_internal,) -end - -# PGState -function AbstractMCMC.getstats(state::PGState) - lp = _get_lp(state.vi) - return (lp=lp,) -end - -# SMCState -function AbstractMCMC.getstats(state::SMCState) - return (logevidence=state.average_logevidence, particle_index=state.particleindex) -end - -# GibbsState -function AbstractMCMC.getstats(state::GibbsState) - lp = _get_lp(state.vi) - return (lp=lp,) -end - -# ESS (VarInfo as state) -function AbstractMCMC.getstats(state::DynamicPPL.AbstractVarInfo) - lp = _get_lp(state) - return (lp=lp,) -end - -# SGHMCState -function AbstractMCMC.getstats(state::SGHMCState) - lp = try - LogDensityProblems.logdensity(state.logdensity, state.params) - catch - NaN - end - return (lp=lp,) -end - -# SGLDState -function AbstractMCMC.getstats(state::SGLDState) - lp = try - LogDensityProblems.logdensity(state.logdensity, state.params) - catch - NaN - end - return (lp=lp, step=state.step) -end - -### -### getparams/getstats from transitions (ParamsWithStats) -### - -function AbstractMCMC.getparams(transition::DynamicPPL.ParamsWithStats) - # params is OrderedDict{VarName, Any} - return [string(vn) => val for (vn, val) in transition.params] -end - -function AbstractMCMC.getstats(transition::DynamicPPL.ParamsWithStats) - return transition.stats -end - -### -### hyperparam_metrics - Define TensorBoard hyperparam metrics -### - -function AbstractMCMC.hyperparam_metrics(model::DynamicPPL.Model, sampler::NUTS) - return [ - "extras/acceptance_rate/stat/Mean", - "extras/max_hamiltonian_energy_error/stat/Mean", - "extras/lp/stat/Mean", - "extras/n_steps/stat/Mean", - "extras/tree_depth/stat/Mean", - ] -end - -function AbstractMCMC.hyperparam_metrics(model::DynamicPPL.Model, sampler::Hamiltonian) - return [ - "extras/acceptance_rate/stat/Mean", - "extras/lp/stat/Mean", - "extras/n_steps/stat/Mean", - ] -end - -function AbstractMCMC.hyperparam_metrics(model::DynamicPPL.Model, sampler::MH) - return ["extras/lp/stat/Mean"] -end - -function AbstractMCMC.hyperparam_metrics(model::DynamicPPL.Model, sampler::PG) - return ["extras/lp/stat/Mean", "extras/logevidence/stat/Mean"] -end - -### -### _hyperparams_impl - Extract sampler hyperparameters -### - -function AbstractMCMC._hyperparams_impl( - model::DynamicPPL.Model, sampler::HMC, state; kwargs... -) - return ["epsilon" => sampler.ϵ, "n_leapfrog" => sampler.n_leapfrog] -end - -function AbstractMCMC._hyperparams_impl( - model::DynamicPPL.Model, sampler::HMCDA, state; kwargs... -) - return [ - "n_adapts" => sampler.n_adapts, - "delta" => sampler.δ, - "lambda" => sampler.λ, - "epsilon" => sampler.ϵ, - ] -end - -function AbstractMCMC._hyperparams_impl( - model::DynamicPPL.Model, sampler::NUTS, state; kwargs... -) - return [ - "n_adapts" => sampler.n_adapts, - "delta" => sampler.δ, - "max_depth" => sampler.max_depth, - "Delta_max" => sampler.Δ_max, - "epsilon" => sampler.ϵ, - ] -end - -function AbstractMCMC._hyperparams_impl( - model::DynamicPPL.Model, sampler::PG, state; kwargs... -) - return ["nparticles" => sampler.nparticles] -end - -function AbstractMCMC._hyperparams_impl( - model::DynamicPPL.Model, sampler::SGHMC, state; kwargs... -) - return [ - "learning_rate" => sampler.learning_rate, "momentum_decay" => sampler.momentum_decay - ] -end - -function AbstractMCMC._hyperparams_impl( - model::DynamicPPL.Model, sampler::SGLD, state; kwargs... -) - return ["stepsize" => string(sampler.stepsize)] -end diff --git a/src/mcmc/ess.jl b/src/mcmc/ess.jl index fa02b6222f..3bf27830a8 100644 --- a/src/mcmc/ess.jl +++ b/src/mcmc/ess.jl @@ -118,3 +118,16 @@ function AbstractMCMC.step( "This method is not implemented! If you want to use the ESS sampler in Turing.jl, please use `Turing.ESS()` instead. If you want the default behaviour in EllipticalSliceSampling.jl, wrap your model in a different subtype of `AbstractMCMC.AbstractModel`, and then implement the necessary EllipticalSliceSampling.jl methods on it.", ) end + +##### +##### AbstractMCMC interface +##### + +function AbstractMCMC.getparams(state::DynamicPPL.AbstractVarInfo) + return _get_params_vector(state) +end + +function AbstractMCMC.getstats(state::DynamicPPL.AbstractVarInfo) + lp = _get_lp(state) + return (lp=lp,) +end diff --git a/src/mcmc/gibbs.jl b/src/mcmc/gibbs.jl index a793071a59..00ad2cbfeb 100644 --- a/src/mcmc/gibbs.jl +++ b/src/mcmc/gibbs.jl @@ -627,3 +627,16 @@ function gibbs_step_recursive( kwargs..., ) end + +##### +##### AbstractMCMC interface +##### + +function AbstractMCMC.getparams(state::GibbsState) + return _get_params_vector(state.vi) +end + +function AbstractMCMC.getstats(state::GibbsState) + lp = _get_lp(state.vi) + return (lp=lp,) +end diff --git a/src/mcmc/hmc.jl b/src/mcmc/hmc.jl index e70c421355..1f1fff29df 100644 --- a/src/mcmc/hmc.jl +++ b/src/mcmc/hmc.jl @@ -505,3 +505,21 @@ end function AHMCAdaptor(::Hamiltonian, ::AHMC.AbstractMetric, nadapts::Int; kwargs...) return AHMC.Adaptation.NoAdaptation() end + +##### +##### AbstractMCMC interface +##### + +function AbstractMCMC.getparams(state::HMCState) + return _get_params_vector(state.vi) +end + +function AbstractMCMC.getstats(state::HMCState) + lp = _get_lp(state.vi) + ϵ = try + state.kernel.τ.integrator.ϵ + catch + NaN + end + return (lp=lp, step_size=ϵ, iteration=state.i) +end diff --git a/src/mcmc/mh.jl b/src/mcmc/mh.jl index 270b6327d7..b04608c923 100644 --- a/src/mcmc/mh.jl +++ b/src/mcmc/mh.jl @@ -442,3 +442,15 @@ function DynamicPPL.tilde_observe!!( ) return DynamicPPL.tilde_observe!!(DefaultContext(), right, left, vn, vi) end + +##### +##### AbstractMCMC interface +##### + +function AbstractMCMC.getparams(state::MHState) + return _get_params_vector(state.varinfo) +end + +function AbstractMCMC.getstats(state::MHState) + return (lp=state.logjoint_internal,) +end diff --git a/src/mcmc/particle_mcmc.jl b/src/mcmc/particle_mcmc.jl index 585d906cb6..286f325844 100644 --- a/src/mcmc/particle_mcmc.jl +++ b/src/mcmc/particle_mcmc.jl @@ -519,3 +519,28 @@ Libtask.might_produce(::Type{<:Tuple{typeof(DynamicPPL.tilde_assume!!),Vararg}}) Libtask.might_produce(::Type{<:Tuple{typeof(DynamicPPL.evaluate!!),Vararg}}) = true Libtask.might_produce(::Type{<:Tuple{typeof(DynamicPPL.init!!),Vararg}}) = true Libtask.might_produce(::Type{<:Tuple{<:DynamicPPL.Model,Vararg}}) = true + +##### +##### AbstractMCMC interface +##### + +# SMCState +function AbstractMCMC.getparams(state::SMCState) + particle = state.particles.vals[state.particleindex] + vi = particle.model.f.varinfo + return _get_params_vector(vi) +end + +function AbstractMCMC.getstats(state::SMCState) + return (logevidence=state.average_logevidence, particle_index=state.particleindex) +end + +# PGState +function AbstractMCMC.getparams(state::PGState) + return _get_params_vector(state.vi) +end + +function AbstractMCMC.getstats(state::PGState) + lp = _get_lp(state.vi) + return (lp=lp,) +end diff --git a/src/mcmc/sghmc.jl b/src/mcmc/sghmc.jl index f9d5d4ade4..fc58d5ace7 100644 --- a/src/mcmc/sghmc.jl +++ b/src/mcmc/sghmc.jl @@ -217,3 +217,35 @@ function AbstractMCMC.step( return transition, newstate end + +##### +##### AbstractMCMC interface +##### + +# SGHMCState +function AbstractMCMC.getparams(state::SGHMCState) + return collect(state.params) +end + +function AbstractMCMC.getstats(state::SGHMCState) + lp = try + LogDensityProblems.logdensity(state.logdensity, state.params) + catch + NaN + end + return (lp=lp,) +end + +# SGLDState +function AbstractMCMC.getparams(state::SGLDState) + return collect(state.params) +end + +function AbstractMCMC.getstats(state::SGLDState) + lp = try + LogDensityProblems.logdensity(state.logdensity, state.params) + catch + NaN + end + return (lp=lp, step=state.step) +end diff --git a/test/mcmc/callbacks.jl b/test/mcmc/callbacks.jl index fe8f1cc05c..a0b8c10087 100644 --- a/test/mcmc/callbacks.jl +++ b/test/mcmc/callbacks.jl @@ -1,232 +1,48 @@ -module TuringCallbacksTests +module CallbacksTests -using Test: @test, @testset -using Turing -using AbstractMCMC: AbstractMCMC -using Random: Random -using DynamicPPL: DynamicPPL +using Test, Turing, AbstractMCMC, Random -using ..Models: gdemo_default - -@testset "AbstractMCMC Callbacks Interface" begin - @testset "getparams from states" begin - @testset "HMCState getparams" begin - Random.seed!(42) - chain = sample(gdemo_default, NUTS(100, 0.65), 10; progress=false) - - rng = Random.default_rng() - transition, state = AbstractMCMC.step( - rng, - gdemo_default, - NUTS(100, 0.65); - initial_params=Turing.Inference.init_strategy(NUTS(100, 0.65)), - ) - - params = AbstractMCMC.getparams(state) - @test params isa Vector - @test length(params) >= 2 - @test all(p -> p isa Pair{String,<:Any}, params) - end - - @testset "MHState getparams" begin - Random.seed!(42) - transition, state = AbstractMCMC.step( - Random.default_rng(), - gdemo_default, - MH(); - initial_params=Turing.Inference.init_strategy(MH()), - ) - - params = AbstractMCMC.getparams(state) - @test params isa Vector - @test length(params) >= 2 - @test all(p -> p isa Pair{String,<:Any}, params) - end +if !isdefined(@__MODULE__, :Models) + include(joinpath(@__DIR__, "..", "test_utils", "models.jl")) + using .Models: gdemo_default +end - @testset "ESS getparams (VarInfo as state)" begin - # ESS only works with Gaussian priors - need simple model - @model function gaussian_model() - m ~ Normal(0, 1) - return nothing - end +@model function simple_gaussian() + x ~ Normal(0, 1) +end - Random.seed!(42) +@testset "AbstractMCMC Callbacks Interface" begin + rng = Random.default_rng() + + samplers = [ + ("NUTS", NUTS(10, 0.65), gdemo_default), + ("HMC", HMC(0.1, 5), gdemo_default), + ("MH", MH(), gdemo_default), + ("ESS", ESS(), simple_gaussian()), + ("Gibbs", Gibbs(:m => HMC(0.1, 5), :s => MH()), gdemo_default), + ("SGHMC", SGHMC(learning_rate=0.01, momentum_decay=1e-2), gdemo_default), + ("PG", PG(10), gdemo_default), + ] + + for (name, sampler, model) in samplers + @testset "$name Interface" begin transition, state = AbstractMCMC.step( - Random.default_rng(), - gaussian_model(), - ESS(); - initial_params=Turing.Inference.init_strategy(ESS()), + rng, model, sampler; + initial_params=Turing.Inference.init_strategy(sampler) ) - - @test state isa DynamicPPL.AbstractVarInfo + + # Should return a flat vector of Reals (unconstrained) params = AbstractMCMC.getparams(state) - @test params isa Vector - @test length(params) >= 1 - end - end - - @testset "getstats from states" begin - @testset "HMCState getstats" begin - Random.seed!(42) - transition, state = AbstractMCMC.step( - Random.default_rng(), - gdemo_default, - NUTS(100, 0.65); - initial_params=Turing.Inference.init_strategy(NUTS(100, 0.65)), - ) - + @test params isa Vector{<:Real} + @test !isempty(params) + + # Should return a NamedTuple with at least log probability (:lp) stats = AbstractMCMC.getstats(state) @test stats isa NamedTuple @test haskey(stats, :lp) @test stats.lp isa Real end - - @testset "MHState getstats" begin - Random.seed!(42) - transition, state = AbstractMCMC.step( - Random.default_rng(), - gdemo_default, - MH(); - initial_params=Turing.Inference.init_strategy(MH()), - ) - - stats = AbstractMCMC.getstats(state) - @test stats isa NamedTuple - @test haskey(stats, :lp) - @test stats.lp isa Real - end - - @testset "ESS getstats (VarInfo)" begin - @model function gaussian_model() - m ~ Normal(0, 1) - return nothing - end - - Random.seed!(42) - transition, state = AbstractMCMC.step( - Random.default_rng(), - gaussian_model(), - ESS(); - initial_params=Turing.Inference.init_strategy(ESS()), - ) - - stats = AbstractMCMC.getstats(state) - @test stats isa NamedTuple - @test haskey(stats, :lp) - @test stats.lp isa Real - end - end - - @testset "ParamsWithStats transitions" begin - Random.seed!(42) - transition, _ = AbstractMCMC.step( - Random.default_rng(), - gdemo_default, - NUTS(100, 0.65); - initial_params=Turing.Inference.init_strategy(NUTS(100, 0.65)), - ) - - @test transition isa DynamicPPL.ParamsWithStats - - stats = AbstractMCMC.getstats(transition) - @test stats isa NamedTuple - - params = AbstractMCMC.getparams(transition) - @test params isa Vector - end - - @testset "hyperparam_metrics" begin - @testset "NUTS hyperparam_metrics" begin - sampler = NUTS() - metrics = AbstractMCMC.hyperparam_metrics(gdemo_default, sampler) - @test metrics isa Vector{String} - @test "extras/lp/stat/Mean" in metrics - @test "extras/acceptance_rate/stat/Mean" in metrics - end - - @testset "HMC hyperparam_metrics (via Hamiltonian)" begin - sampler = HMC(0.1, 10) - metrics = AbstractMCMC.hyperparam_metrics(gdemo_default, sampler) - @test metrics isa Vector{String} - @test "extras/lp/stat/Mean" in metrics - end - - @testset "MH hyperparam_metrics" begin - sampler = MH() - metrics = AbstractMCMC.hyperparam_metrics(gdemo_default, sampler) - @test metrics isa Vector{String} - @test "extras/lp/stat/Mean" in metrics - end - - @testset "PG hyperparam_metrics" begin - sampler = PG(10) - metrics = AbstractMCMC.hyperparam_metrics(gdemo_default, sampler) - @test metrics isa Vector{String} - @test "extras/lp/stat/Mean" in metrics - end - end - - @testset "_hyperparams_impl" begin - @testset "HMC hyperparams" begin - sampler = HMC(0.1, 10) - hyperparams = AbstractMCMC._hyperparams_impl(gdemo_default, sampler, nothing) - @test hyperparams isa Vector - - hp_dict = Dict(hyperparams) - @test hp_dict["epsilon"] == 0.1 - @test hp_dict["n_leapfrog"] == 10 - end - - @testset "NUTS hyperparams" begin - sampler = NUTS(200, 0.65) - hyperparams = AbstractMCMC._hyperparams_impl(gdemo_default, sampler, nothing) - @test hyperparams isa Vector - - hp_dict = Dict(hyperparams) - @test hp_dict["n_adapts"] == 200 - @test hp_dict["delta"] == 0.65 - @test haskey(hp_dict, "max_depth") - end - - @testset "HMCDA hyperparams" begin - sampler = HMCDA(200, 0.65, 0.3) - hyperparams = AbstractMCMC._hyperparams_impl(gdemo_default, sampler, nothing) - @test hyperparams isa Vector - - hp_dict = Dict(hyperparams) - @test hp_dict["n_adapts"] == 200 - @test hp_dict["delta"] == 0.65 - @test hp_dict["lambda"] == 0.3 - end - - @testset "PG hyperparams" begin - sampler = PG(10) - hyperparams = AbstractMCMC._hyperparams_impl(gdemo_default, sampler, nothing) - @test hyperparams isa Vector - - hp_dict = Dict(hyperparams) - @test hp_dict["nparticles"] == 10 - end - - @testset "SGHMC hyperparams" begin - sampler = SGHMC(; learning_rate=0.01, momentum_decay=0.1) - hyperparams = AbstractMCMC._hyperparams_impl(gdemo_default, sampler, nothing) - @test hyperparams isa Vector - - hp_dict = Dict(hyperparams) - @test hp_dict["learning_rate"] == 0.01 - @test hp_dict["momentum_decay"] == 0.1 - end - - @testset "SGLD hyperparams" begin - sampler = SGLD() - hyperparams = AbstractMCMC._hyperparams_impl(gdemo_default, sampler, nothing) - @test hyperparams isa Vector - - hp_dict = Dict(hyperparams) - @test haskey(hp_dict, "stepsize") - end end end -end # module +end From 65f545b9cfb2e61fb0e33adfc9c943f36530473e Mon Sep 17 00:00:00 2001 From: Shravan Goswami Date: Fri, 23 Jan 2026 11:48:00 +0530 Subject: [PATCH 04/13] format --- test/mcmc/callbacks.jl | 11 +++++------ 1 file changed, 5 insertions(+), 6 deletions(-) diff --git a/test/mcmc/callbacks.jl b/test/mcmc/callbacks.jl index a0b8c10087..1c6e9179aa 100644 --- a/test/mcmc/callbacks.jl +++ b/test/mcmc/callbacks.jl @@ -8,7 +8,7 @@ if !isdefined(@__MODULE__, :Models) end @model function simple_gaussian() - x ~ Normal(0, 1) + return x ~ Normal(0, 1) end @testset "AbstractMCMC Callbacks Interface" begin @@ -20,22 +20,21 @@ end ("MH", MH(), gdemo_default), ("ESS", ESS(), simple_gaussian()), ("Gibbs", Gibbs(:m => HMC(0.1, 5), :s => MH()), gdemo_default), - ("SGHMC", SGHMC(learning_rate=0.01, momentum_decay=1e-2), gdemo_default), + ("SGHMC", SGHMC(; learning_rate=0.01, momentum_decay=1e-2), gdemo_default), ("PG", PG(10), gdemo_default), ] for (name, sampler, model) in samplers @testset "$name Interface" begin transition, state = AbstractMCMC.step( - rng, model, sampler; - initial_params=Turing.Inference.init_strategy(sampler) + rng, model, sampler; initial_params=Turing.Inference.init_strategy(sampler) ) - + # Should return a flat vector of Reals (unconstrained) params = AbstractMCMC.getparams(state) @test params isa Vector{<:Real} @test !isempty(params) - + # Should return a NamedTuple with at least log probability (:lp) stats = AbstractMCMC.getstats(state) @test stats isa NamedTuple From 73924f50e62651afd29212ab5530e38362c87b6b Mon Sep 17 00:00:00 2001 From: Hong Ge <3279477+yebai@users.noreply.github.com> Date: Fri, 23 Jan 2026 14:37:34 +0000 Subject: [PATCH 05/13] Apply suggestions from code review --- src/mcmc/Inference.jl | 7 +++++++ src/mcmc/ess.jl | 13 ------------- 2 files changed, 7 insertions(+), 13 deletions(-) diff --git a/src/mcmc/Inference.jl b/src/mcmc/Inference.jl index 859326f724..54096cbaa7 100644 --- a/src/mcmc/Inference.jl +++ b/src/mcmc/Inference.jl @@ -126,7 +126,14 @@ end function _get_params_vector(vi::DynamicPPL.AbstractVarInfo) return vi[:] end +function AbstractMCMC.getparams(state::DynamicPPL.AbstractVarInfo) + return _get_params_vector(state) +end +function AbstractMCMC.getstats(state::DynamicPPL.AbstractVarInfo) + lp = _get_lp(state) + return (lp=lp,) +end ####################################### # Concrete algorithm implementations. # ####################################### diff --git a/src/mcmc/ess.jl b/src/mcmc/ess.jl index 3bf27830a8..fa02b6222f 100644 --- a/src/mcmc/ess.jl +++ b/src/mcmc/ess.jl @@ -118,16 +118,3 @@ function AbstractMCMC.step( "This method is not implemented! If you want to use the ESS sampler in Turing.jl, please use `Turing.ESS()` instead. If you want the default behaviour in EllipticalSliceSampling.jl, wrap your model in a different subtype of `AbstractMCMC.AbstractModel`, and then implement the necessary EllipticalSliceSampling.jl methods on it.", ) end - -##### -##### AbstractMCMC interface -##### - -function AbstractMCMC.getparams(state::DynamicPPL.AbstractVarInfo) - return _get_params_vector(state) -end - -function AbstractMCMC.getstats(state::DynamicPPL.AbstractVarInfo) - lp = _get_lp(state) - return (lp=lp,) -end From 354244a224ee4fe0aab6662c0eab6100f66e4871 Mon Sep 17 00:00:00 2001 From: Shravan Goswami Date: Sun, 25 Jan 2026 11:20:20 +0530 Subject: [PATCH 06/13] implement new suggestions --- src/mcmc/Inference.jl | 8 +++++--- src/mcmc/ess.jl | 4 ---- src/mcmc/gibbs.jl | 4 ---- src/mcmc/hmc.jl | 9 +++++---- src/mcmc/mh.jl | 4 ---- src/mcmc/particle_mcmc.jl | 18 +++--------------- src/mcmc/sghmc.jl | 22 ++++++---------------- test/mcmc/callbacks.jl | 36 ++++++++++++++++++------------------ 8 files changed, 37 insertions(+), 68 deletions(-) diff --git a/src/mcmc/Inference.jl b/src/mcmc/Inference.jl index 859326f724..aca3c03ee1 100644 --- a/src/mcmc/Inference.jl +++ b/src/mcmc/Inference.jl @@ -122,9 +122,11 @@ function _get_lp(vi::DynamicPPL.AbstractVarInfo) return sum(values(lp)) end -# Helper to extract raw parameter values from VarInfo as Vector{<:Real} -function _get_params_vector(vi::DynamicPPL.AbstractVarInfo) - return vi[:] +# Consolidated getparams using get_varinfo (defined in each sampler file) +# This covers HMC, MH, PG, Gibbs, and external samplers. +# SGHMC/SGLD have their own implementations since they don't use VarInfo. +function AbstractMCMC.getparams(state) + return get_varinfo(state)[:] end ####################################### diff --git a/src/mcmc/ess.jl b/src/mcmc/ess.jl index 3bf27830a8..6a9ae828a4 100644 --- a/src/mcmc/ess.jl +++ b/src/mcmc/ess.jl @@ -123,10 +123,6 @@ end ##### AbstractMCMC interface ##### -function AbstractMCMC.getparams(state::DynamicPPL.AbstractVarInfo) - return _get_params_vector(state) -end - function AbstractMCMC.getstats(state::DynamicPPL.AbstractVarInfo) lp = _get_lp(state) return (lp=lp,) diff --git a/src/mcmc/gibbs.jl b/src/mcmc/gibbs.jl index 00ad2cbfeb..1f99215d44 100644 --- a/src/mcmc/gibbs.jl +++ b/src/mcmc/gibbs.jl @@ -632,10 +632,6 @@ end ##### AbstractMCMC interface ##### -function AbstractMCMC.getparams(state::GibbsState) - return _get_params_vector(state.vi) -end - function AbstractMCMC.getstats(state::GibbsState) lp = _get_lp(state.vi) return (lp=lp,) diff --git a/src/mcmc/hmc.jl b/src/mcmc/hmc.jl index 1f1fff29df..c1a06b3470 100644 --- a/src/mcmc/hmc.jl +++ b/src/mcmc/hmc.jl @@ -510,12 +510,13 @@ end ##### AbstractMCMC interface ##### -function AbstractMCMC.getparams(state::HMCState) - return _get_params_vector(state.vi) -end - function AbstractMCMC.getstats(state::HMCState) lp = _get_lp(state.vi) + # TODO(penelopeysm): For many Hamiltonian samplers you can get the + # stepsize from the state. However for non-adaptive Hamiltonians the + # info is only in the sampler and thus can't be accessed via getstats. + # HMCState should be modified to include the step size as a field, which + # would allow us to access this information here properly. ϵ = try state.kernel.τ.integrator.ϵ catch diff --git a/src/mcmc/mh.jl b/src/mcmc/mh.jl index b04608c923..64a429c219 100644 --- a/src/mcmc/mh.jl +++ b/src/mcmc/mh.jl @@ -447,10 +447,6 @@ end ##### AbstractMCMC interface ##### -function AbstractMCMC.getparams(state::MHState) - return _get_params_vector(state.varinfo) -end - function AbstractMCMC.getstats(state::MHState) return (lp=state.logjoint_internal,) end diff --git a/src/mcmc/particle_mcmc.jl b/src/mcmc/particle_mcmc.jl index 286f325844..2ebe477ee4 100644 --- a/src/mcmc/particle_mcmc.jl +++ b/src/mcmc/particle_mcmc.jl @@ -524,23 +524,11 @@ Libtask.might_produce(::Type{<:Tuple{<:DynamicPPL.Model,Vararg}}) = true ##### AbstractMCMC interface ##### -# SMCState -function AbstractMCMC.getparams(state::SMCState) - particle = state.particles.vals[state.particleindex] - vi = particle.model.f.varinfo - return _get_params_vector(vi) -end - -function AbstractMCMC.getstats(state::SMCState) - return (logevidence=state.average_logevidence, particle_index=state.particleindex) -end - -# PGState -function AbstractMCMC.getparams(state::PGState) - return _get_params_vector(state.vi) -end +# Note: SMCState getparams/getstats intentionally omitted - SMC is not an MCMC +# method and doesn't fit the conventional AbstractMCMC interface. function AbstractMCMC.getstats(state::PGState) lp = _get_lp(state.vi) return (lp=lp,) end + diff --git a/src/mcmc/sghmc.jl b/src/mcmc/sghmc.jl index fc58d5ace7..5d42b75aaf 100644 --- a/src/mcmc/sghmc.jl +++ b/src/mcmc/sghmc.jl @@ -223,29 +223,19 @@ end ##### # SGHMCState -function AbstractMCMC.getparams(state::SGHMCState) - return collect(state.params) -end +AbstractMCMC.getparams(state::SGHMCState) = state.params function AbstractMCMC.getstats(state::SGHMCState) - lp = try - LogDensityProblems.logdensity(state.logdensity, state.params) - catch - NaN - end + # TODO(penelopeysm): This is inefficient as it requires an extra model evaluation + lp = LogDensityProblems.logdensity(state.logdensity, state.params) return (lp=lp,) end # SGLDState -function AbstractMCMC.getparams(state::SGLDState) - return collect(state.params) -end +AbstractMCMC.getparams(state::SGLDState) = state.params function AbstractMCMC.getstats(state::SGLDState) - lp = try - LogDensityProblems.logdensity(state.logdensity, state.params) - catch - NaN - end + # TODO(penelopeysm): Remove extra evaluation. + lp = LogDensityProblems.logdensity(state.logdensity, state.params) return (lp=lp, step=state.step) end diff --git a/test/mcmc/callbacks.jl b/test/mcmc/callbacks.jl index 1c6e9179aa..2984646245 100644 --- a/test/mcmc/callbacks.jl +++ b/test/mcmc/callbacks.jl @@ -1,30 +1,29 @@ module CallbacksTests -using Test, Turing, AbstractMCMC, Random +using Test, Turing, AbstractMCMC, Random, Distributions, LinearAlgebra -if !isdefined(@__MODULE__, :Models) - include(joinpath(@__DIR__, "..", "test_utils", "models.jl")) - using .Models: gdemo_default -end - -@model function simple_gaussian() - return x ~ Normal(0, 1) +# Simple model that works for all samplers (ESS requires Normal distributions) +@model function test_normals() + x ~ Normal() + y ~ MvNormal(zeros(3), I) end @testset "AbstractMCMC Callbacks Interface" begin rng = Random.default_rng() + model = test_normals() + # All samplers use the same model (4 params: x + y[1:3]) samplers = [ - ("NUTS", NUTS(10, 0.65), gdemo_default), - ("HMC", HMC(0.1, 5), gdemo_default), - ("MH", MH(), gdemo_default), - ("ESS", ESS(), simple_gaussian()), - ("Gibbs", Gibbs(:m => HMC(0.1, 5), :s => MH()), gdemo_default), - ("SGHMC", SGHMC(; learning_rate=0.01, momentum_decay=1e-2), gdemo_default), - ("PG", PG(10), gdemo_default), + ("NUTS", NUTS(10, 0.65)), + ("HMC", HMC(0.1, 5)), + ("MH", MH()), + ("ESS", ESS()), + ("Gibbs", Gibbs(:x => HMC(0.1, 5), :y => MH())), + ("SGHMC", SGHMC(; learning_rate=0.01, momentum_decay=1e-2)), + ("PG", PG(10)), ] - for (name, sampler, model) in samplers + for (name, sampler) in samplers @testset "$name Interface" begin transition, state = AbstractMCMC.step( rng, model, sampler; initial_params=Turing.Inference.init_strategy(sampler) @@ -32,14 +31,15 @@ end # Should return a flat vector of Reals (unconstrained) params = AbstractMCMC.getparams(state) - @test params isa Vector{<:Real} - @test !isempty(params) + @test params isa AbstractVector{<:Real} + @test length(params) == 4 # x (1) + y (3) # Should return a NamedTuple with at least log probability (:lp) stats = AbstractMCMC.getstats(state) @test stats isa NamedTuple @test haskey(stats, :lp) @test stats.lp isa Real + @test isfinite(stats.lp) # Should be a valid log probability end end end From 1c0828e98e37d55077856edc73215efe55ad553d Mon Sep 17 00:00:00 2001 From: Shravan Goswami Date: Mon, 26 Jan 2026 13:30:31 +0530 Subject: [PATCH 07/13] implement all suggestions --- src/mcmc/Inference.jl | 17 +++++++++++++++ test/mcmc/callbacks.jl | 47 +++++++++++++++++++++++++++++++++--------- 2 files changed, 54 insertions(+), 10 deletions(-) diff --git a/src/mcmc/Inference.jl b/src/mcmc/Inference.jl index aca3c03ee1..75b522fef3 100644 --- a/src/mcmc/Inference.jl +++ b/src/mcmc/Inference.jl @@ -129,6 +129,23 @@ function AbstractMCMC.getparams(state) return get_varinfo(state)[:] end +# Override for DynamicPPL.ParamsWithStats: provides named params and full transition metrics +function AbstractMCMC.ParamsWithStats( + model, + sampler, + transition::DynamicPPL.ParamsWithStats, + state; + params::Bool=true, + stats::Bool=false, + extras::Bool=false, +) + # Convert OrderedDict params to Vector{Pair} for AbstractMCMC.ParamsWithStats constructor + p = params ? [string(k) => v for (k, v) in transition.params] : nothing + s = stats ? transition.stats : NamedTuple() + e = extras ? NamedTuple() : NamedTuple() + return AbstractMCMC.ParamsWithStats(p, s, e) +end + ####################################### # Concrete algorithm implementations. # ####################################### diff --git a/test/mcmc/callbacks.jl b/test/mcmc/callbacks.jl index 2984646245..ee3a9391ba 100644 --- a/test/mcmc/callbacks.jl +++ b/test/mcmc/callbacks.jl @@ -2,17 +2,15 @@ module CallbacksTests using Test, Turing, AbstractMCMC, Random, Distributions, LinearAlgebra -# Simple model that works for all samplers (ESS requires Normal distributions) @model function test_normals() x ~ Normal() - y ~ MvNormal(zeros(3), I) + return y ~ MvNormal(zeros(3), I) end @testset "AbstractMCMC Callbacks Interface" begin rng = Random.default_rng() model = test_normals() - # All samplers use the same model (4 params: x + y[1:3]) samplers = [ ("NUTS", NUTS(10, 0.65)), ("HMC", HMC(0.1, 5)), @@ -24,24 +22,53 @@ end ] for (name, sampler) in samplers - @testset "$name Interface" begin - transition, state = AbstractMCMC.step( + @testset "$name" begin + t1, s1 = AbstractMCMC.step( rng, model, sampler; initial_params=Turing.Inference.init_strategy(sampler) ) - # Should return a flat vector of Reals (unconstrained) - params = AbstractMCMC.getparams(state) + # getparams returns flat vector of Reals + params = AbstractMCMC.getparams(s1) @test params isa AbstractVector{<:Real} @test length(params) == 4 # x (1) + y (3) - # Should return a NamedTuple with at least log probability (:lp) - stats = AbstractMCMC.getstats(state) + # getstats returns NamedTuple with :lp + stats = AbstractMCMC.getstats(s1) @test stats isa NamedTuple @test haskey(stats, :lp) @test stats.lp isa Real - @test isfinite(stats.lp) # Should be a valid log probability + @test isfinite(stats.lp) + + # ParamsWithStats returns named params (not θ[i]) + pws = AbstractMCMC.ParamsWithStats( + model, sampler, t1, s1; params=true, stats=true + ) + pairs_dict = Dict(k => v for (k, v) in Base.pairs(pws)) + # Keys are Symbols since ParamsWithStats stores NamedTuple internally + @test haskey(pairs_dict, Symbol("x")) + @test haskey(pairs_dict, Symbol("y")) + @test pairs_dict[Symbol("y")] isa AbstractVector + @test length(pairs_dict[Symbol("y")]) == 3 end end + + # NUTS second step has full AHMC transition metrics + @testset "NUTS Transition Metrics" begin + sampler = NUTS(10, 0.65) + t1, s1 = AbstractMCMC.step( + rng, model, sampler; initial_params=Turing.Inference.init_strategy(sampler) + ) + t2, s2 = AbstractMCMC.step(rng, model, sampler, s1) + + pws = AbstractMCMC.ParamsWithStats(model, sampler, t2, s2; params=true, stats=true) + pairs_dict = Dict(k => v for (k, v) in Base.pairs(pws)) + + # Keys are Symbols from NamedTuple + @test haskey(pairs_dict, :tree_depth) + @test haskey(pairs_dict, :n_steps) + @test haskey(pairs_dict, :acceptance_rate) + @test haskey(pairs_dict, :hamiltonian_energy) + end end end From 03ff92d804df7be63cd653a49504e4c5d6fca605 Mon Sep 17 00:00:00 2001 From: Shravan Goswami Date: Mon, 26 Jan 2026 14:20:56 +0530 Subject: [PATCH 08/13] update tests --- Project.toml | 2 +- src/mcmc/Inference.jl | 3 --- src/mcmc/particle_mcmc.jl | 1 - test/Project.toml | 2 +- test/mcmc/callbacks.jl | 8 ++++++++ 5 files changed, 10 insertions(+), 6 deletions(-) diff --git a/Project.toml b/Project.toml index 56d2211670..64e48c15df 100644 --- a/Project.toml +++ b/Project.toml @@ -47,7 +47,7 @@ TuringDynamicHMCExt = "DynamicHMC" [compat] ADTypes = "1.9" -AbstractMCMC = "5.11" +AbstractMCMC = "5.11, 5.12" AbstractPPL = "0.11, 0.12, 0.13" Accessors = "0.1" AdvancedHMC = "0.8.3" diff --git a/src/mcmc/Inference.jl b/src/mcmc/Inference.jl index 4f299a0d37..1bdf979dd1 100644 --- a/src/mcmc/Inference.jl +++ b/src/mcmc/Inference.jl @@ -145,9 +145,6 @@ function AbstractMCMC.ParamsWithStats( e = extras ? NamedTuple() : NamedTuple() return AbstractMCMC.ParamsWithStats(p, s, e) end -function AbstractMCMC.getparams(state::DynamicPPL.AbstractVarInfo) - return _get_params_vector(state) -end function AbstractMCMC.getstats(state::DynamicPPL.AbstractVarInfo) lp = _get_lp(state) diff --git a/src/mcmc/particle_mcmc.jl b/src/mcmc/particle_mcmc.jl index 5236dbe3d9..2086adc72e 100644 --- a/src/mcmc/particle_mcmc.jl +++ b/src/mcmc/particle_mcmc.jl @@ -551,4 +551,3 @@ function AbstractMCMC.getstats(state::PGState) lp = _get_lp(state.vi) return (lp=lp,) end - diff --git a/test/Project.toml b/test/Project.toml index 15d53797b3..19fc1d8748 100644 --- a/test/Project.toml +++ b/test/Project.toml @@ -40,7 +40,7 @@ TimerOutputs = "a759f4b9-e2f1-59dc-863e-4aeb61b1ea8f" [compat] ADTypes = "1" -AbstractMCMC = "5.11" +AbstractMCMC = "5.11, 5.12" AbstractPPL = "0.11, 0.12, 0.13" AdvancedMH = "0.8.9" AdvancedPS = "0.7.2" diff --git a/test/mcmc/callbacks.jl b/test/mcmc/callbacks.jl index ee3a9391ba..e017946014 100644 --- a/test/mcmc/callbacks.jl +++ b/test/mcmc/callbacks.jl @@ -39,6 +39,14 @@ end @test stats.lp isa Real @test isfinite(stats.lp) + # x ~ Normal(), y ~ MvNormal(zeros(3), I) + # Note: PG uses internal parameterization that may differ from unlinked space + if name != "PG" + expected_lp = + logpdf(Normal(), params[1]) + logpdf(MvNormal(zeros(3), I), params[2:4]) + @test stats.lp ≈ expected_lp atol = 1e-6 + end + # ParamsWithStats returns named params (not θ[i]) pws = AbstractMCMC.ParamsWithStats( model, sampler, t1, s1; params=true, stats=true From 6a70f081f2c6f9e1a09fd9704c2da83dc8fef4e9 Mon Sep 17 00:00:00 2001 From: Shravan Goswami Date: Tue, 27 Jan 2026 16:17:09 +0530 Subject: [PATCH 09/13] fix lp calculation --- src/mcmc/Inference.jl | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/src/mcmc/Inference.jl b/src/mcmc/Inference.jl index 1bdf979dd1..f69898e406 100644 --- a/src/mcmc/Inference.jl +++ b/src/mcmc/Inference.jl @@ -116,9 +116,14 @@ function mh_accept(logp_current::Real, logp_proposal::Real, log_proposal_ratio:: end # Helper functions for AbstractMCMC callbacks -# Helper to get log probability from VarInfo +# Helper to get log probability from VarInfo in unconstrained space. +# The logjac from getlogp is the forward Jacobian (constrained -> unconstrained), +# so we negate it: log q(y) = log p(x) - log|J| function _get_lp(vi::DynamicPPL.AbstractVarInfo) lp = DynamicPPL.getlogp(vi) + if haskey(lp, :logjac) + lp = merge(lp, (; logjac=-lp.logjac)) + end return sum(values(lp)) end From 364df74003b96aa6dbffe73f0a6f07a3d48fc673 Mon Sep 17 00:00:00 2001 From: Shravan Goswami Date: Tue, 27 Jan 2026 16:44:38 +0530 Subject: [PATCH 10/13] pass MCMC states to getparams --- src/mcmc/Inference.jl | 5 ++--- 1 file changed, 2 insertions(+), 3 deletions(-) diff --git a/src/mcmc/Inference.jl b/src/mcmc/Inference.jl index f69898e406..e8bdcc733b 100644 --- a/src/mcmc/Inference.jl +++ b/src/mcmc/Inference.jl @@ -127,10 +127,9 @@ function _get_lp(vi::DynamicPPL.AbstractVarInfo) return sum(values(lp)) end -# Consolidated getparams using get_varinfo (defined in each sampler file) -# This covers HMC, MH, PG, Gibbs, and external samplers. +# Consolidated getparams using get_varinfo (defined in each sampler file). # SGHMC/SGLD have their own implementations since they don't use VarInfo. -function AbstractMCMC.getparams(state) +function AbstractMCMC.getparams(state::Union{HMCState,MHState,PGState,GibbsState}) return get_varinfo(state)[:] end From 8965eef712dc8c29e57cf14b9e306d6105f0d935 Mon Sep 17 00:00:00 2001 From: Shravan Goswami Date: Tue, 27 Jan 2026 18:21:34 +0530 Subject: [PATCH 11/13] nice cleanup --- src/mcmc/Inference.jl | 29 +++++------------------------ src/mcmc/gibbs.jl | 9 --------- src/mcmc/hmc.jl | 19 ------------------- src/mcmc/mh.jl | 8 -------- src/mcmc/particle_mcmc.jl | 12 ------------ src/mcmc/sghmc.jl | 22 ---------------------- test/mcmc/callbacks.jl | 23 +++-------------------- 7 files changed, 8 insertions(+), 114 deletions(-) diff --git a/src/mcmc/Inference.jl b/src/mcmc/Inference.jl index e8bdcc733b..031974f0f3 100644 --- a/src/mcmc/Inference.jl +++ b/src/mcmc/Inference.jl @@ -115,25 +115,11 @@ function mh_accept(logp_current::Real, logp_proposal::Real, log_proposal_ratio:: return log(rand()) + logp_current ≤ logp_proposal + log_proposal_ratio end -# Helper functions for AbstractMCMC callbacks -# Helper to get log probability from VarInfo in unconstrained space. -# The logjac from getlogp is the forward Jacobian (constrained -> unconstrained), -# so we negate it: log q(y) = log p(x) - log|J| -function _get_lp(vi::DynamicPPL.AbstractVarInfo) - lp = DynamicPPL.getlogp(vi) - if haskey(lp, :logjac) - lp = merge(lp, (; logjac=-lp.logjac)) - end - return sum(values(lp)) -end - -# Consolidated getparams using get_varinfo (defined in each sampler file). -# SGHMC/SGLD have their own implementations since they don't use VarInfo. -function AbstractMCMC.getparams(state::Union{HMCState,MHState,PGState,GibbsState}) - return get_varinfo(state)[:] -end - -# Override for DynamicPPL.ParamsWithStats: provides named params and full transition metrics +# Directly overload the constructor of `AbstractMCMC.ParamsWithStats` so that we don't +# hit the default method, which uses `getparams(state)` and `getstats(state)`. For Turing's +# MCMC samplers, the state might contain results that are in linked space. Using the +# outputs of the transition here ensures that parameters and logprobs are provided in +# user space (similar to chains output). function AbstractMCMC.ParamsWithStats( model, sampler, @@ -143,17 +129,12 @@ function AbstractMCMC.ParamsWithStats( stats::Bool=false, extras::Bool=false, ) - # Convert OrderedDict params to Vector{Pair} for AbstractMCMC.ParamsWithStats constructor p = params ? [string(k) => v for (k, v) in transition.params] : nothing s = stats ? transition.stats : NamedTuple() e = extras ? NamedTuple() : NamedTuple() return AbstractMCMC.ParamsWithStats(p, s, e) end -function AbstractMCMC.getstats(state::DynamicPPL.AbstractVarInfo) - lp = _get_lp(state) - return (lp=lp,) -end ####################################### # Concrete algorithm implementations. # ####################################### diff --git a/src/mcmc/gibbs.jl b/src/mcmc/gibbs.jl index 1f99215d44..a793071a59 100644 --- a/src/mcmc/gibbs.jl +++ b/src/mcmc/gibbs.jl @@ -627,12 +627,3 @@ function gibbs_step_recursive( kwargs..., ) end - -##### -##### AbstractMCMC interface -##### - -function AbstractMCMC.getstats(state::GibbsState) - lp = _get_lp(state.vi) - return (lp=lp,) -end diff --git a/src/mcmc/hmc.jl b/src/mcmc/hmc.jl index c1a06b3470..e70c421355 100644 --- a/src/mcmc/hmc.jl +++ b/src/mcmc/hmc.jl @@ -505,22 +505,3 @@ end function AHMCAdaptor(::Hamiltonian, ::AHMC.AbstractMetric, nadapts::Int; kwargs...) return AHMC.Adaptation.NoAdaptation() end - -##### -##### AbstractMCMC interface -##### - -function AbstractMCMC.getstats(state::HMCState) - lp = _get_lp(state.vi) - # TODO(penelopeysm): For many Hamiltonian samplers you can get the - # stepsize from the state. However for non-adaptive Hamiltonians the - # info is only in the sampler and thus can't be accessed via getstats. - # HMCState should be modified to include the step size as a field, which - # would allow us to access this information here properly. - ϵ = try - state.kernel.τ.integrator.ϵ - catch - NaN - end - return (lp=lp, step_size=ϵ, iteration=state.i) -end diff --git a/src/mcmc/mh.jl b/src/mcmc/mh.jl index 64a429c219..270b6327d7 100644 --- a/src/mcmc/mh.jl +++ b/src/mcmc/mh.jl @@ -442,11 +442,3 @@ function DynamicPPL.tilde_observe!!( ) return DynamicPPL.tilde_observe!!(DefaultContext(), right, left, vn, vi) end - -##### -##### AbstractMCMC interface -##### - -function AbstractMCMC.getstats(state::MHState) - return (lp=state.logjoint_internal,) -end diff --git a/src/mcmc/particle_mcmc.jl b/src/mcmc/particle_mcmc.jl index ab8c083654..2512a2f154 100644 --- a/src/mcmc/particle_mcmc.jl +++ b/src/mcmc/particle_mcmc.jl @@ -513,15 +513,3 @@ Libtask.@might_produce(DynamicPPL.tilde_assume!!) Libtask.@might_produce(DynamicPPL.evaluate!!) Libtask.@might_produce(DynamicPPL.init!!) Libtask.might_produce(::Type{<:Tuple{<:DynamicPPL.Model,Vararg}}) = true - -##### -##### AbstractMCMC interface -##### - -# Note: SMCState getparams/getstats intentionally omitted - SMC is not an MCMC -# method and doesn't fit the conventional AbstractMCMC interface. - -function AbstractMCMC.getstats(state::PGState) - lp = _get_lp(state.vi) - return (lp=lp,) -end diff --git a/src/mcmc/sghmc.jl b/src/mcmc/sghmc.jl index 5d42b75aaf..f9d5d4ade4 100644 --- a/src/mcmc/sghmc.jl +++ b/src/mcmc/sghmc.jl @@ -217,25 +217,3 @@ function AbstractMCMC.step( return transition, newstate end - -##### -##### AbstractMCMC interface -##### - -# SGHMCState -AbstractMCMC.getparams(state::SGHMCState) = state.params - -function AbstractMCMC.getstats(state::SGHMCState) - # TODO(penelopeysm): This is inefficient as it requires an extra model evaluation - lp = LogDensityProblems.logdensity(state.logdensity, state.params) - return (lp=lp,) -end - -# SGLDState -AbstractMCMC.getparams(state::SGLDState) = state.params - -function AbstractMCMC.getstats(state::SGLDState) - # TODO(penelopeysm): Remove extra evaluation. - lp = LogDensityProblems.logdensity(state.logdensity, state.params) - return (lp=lp, step=state.step) -end diff --git a/test/mcmc/callbacks.jl b/test/mcmc/callbacks.jl index e017946014..c62e9d133a 100644 --- a/test/mcmc/callbacks.jl +++ b/test/mcmc/callbacks.jl @@ -27,26 +27,6 @@ end rng, model, sampler; initial_params=Turing.Inference.init_strategy(sampler) ) - # getparams returns flat vector of Reals - params = AbstractMCMC.getparams(s1) - @test params isa AbstractVector{<:Real} - @test length(params) == 4 # x (1) + y (3) - - # getstats returns NamedTuple with :lp - stats = AbstractMCMC.getstats(s1) - @test stats isa NamedTuple - @test haskey(stats, :lp) - @test stats.lp isa Real - @test isfinite(stats.lp) - - # x ~ Normal(), y ~ MvNormal(zeros(3), I) - # Note: PG uses internal parameterization that may differ from unlinked space - if name != "PG" - expected_lp = - logpdf(Normal(), params[1]) + logpdf(MvNormal(zeros(3), I), params[2:4]) - @test stats.lp ≈ expected_lp atol = 1e-6 - end - # ParamsWithStats returns named params (not θ[i]) pws = AbstractMCMC.ParamsWithStats( model, sampler, t1, s1; params=true, stats=true @@ -57,6 +37,9 @@ end @test haskey(pairs_dict, Symbol("y")) @test pairs_dict[Symbol("y")] isa AbstractVector @test length(pairs_dict[Symbol("y")]) == 3 + + # Check stats contain lp + @test haskey(pairs_dict, :lp) || haskey(pairs_dict, :logjoint) end end From e5b1b3e5b39a8f92a4b8e702a4304a0fb4c0fb04 Mon Sep 17 00:00:00 2001 From: Shravan Goswami Date: Wed, 28 Jan 2026 16:39:53 +0530 Subject: [PATCH 12/13] version bump: 0.42.6 --> 0.42.7 --- Project.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/Project.toml b/Project.toml index 9ade3f3e39..0503207adc 100644 --- a/Project.toml +++ b/Project.toml @@ -1,6 +1,6 @@ name = "Turing" uuid = "fce5fe82-541a-59a6-adf8-730c64b5f9a0" -version = "0.42.6" +version = "0.42.7" [deps] ADTypes = "47edcb42-4c32-4615-8424-f2b9edc5f35b" From 58a449664961adc151851f2577fa37bdfa145609 Mon Sep 17 00:00:00 2001 From: Hong Ge <3279477+yebai@users.noreply.github.com> Date: Wed, 28 Jan 2026 17:24:22 +0000 Subject: [PATCH 13/13] Update Project.toml --- test/Project.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/test/Project.toml b/test/Project.toml index 19fc1d8748..b2dfb189fa 100644 --- a/test/Project.toml +++ b/test/Project.toml @@ -40,7 +40,7 @@ TimerOutputs = "a759f4b9-e2f1-59dc-863e-4aeb61b1ea8f" [compat] ADTypes = "1" -AbstractMCMC = "5.11, 5.12" +AbstractMCMC = "5.13" AbstractPPL = "0.11, 0.12, 0.13" AdvancedMH = "0.8.9" AdvancedPS = "0.7.2"