From 59c2e3e3e793e9ef52f386266a744e29dc019e6e Mon Sep 17 00:00:00 2001 From: Hong Ge Date: Sun, 20 Sep 2026 20:03:25 +0100 Subject: [PATCH 01/11] chains: support model evaluation on `SamplingOutput` Assisted-by: Codex --- Project.toml | 2 +- src/chains.jl | 78 ++++++++++++++++++++++++++++++++++++++++++ test/runtests.jl | 1 + test/samplingoutput.jl | 44 ++++++++++++++++++++++++ 4 files changed, 124 insertions(+), 1 deletion(-) create mode 100644 test/samplingoutput.jl diff --git a/Project.toml b/Project.toml index 57284a8bc..b7861dbf7 100644 --- a/Project.toml +++ b/Project.toml @@ -55,7 +55,7 @@ DynamicPPLReverseDiffExt = ["ReverseDiff"] [compat] ADTypes = "1" -AbstractMCMC = "5.14" +AbstractMCMC = "5.17" AbstractPPL = "0.15" Accessors = "0.1" BangBang = "0.4.1" diff --git a/src/chains.jl b/src/chains.jl index 219ccb444..fe81f6143 100644 --- a/src/chains.jl +++ b/src/chains.jl @@ -264,3 +264,81 @@ function InitFromParams( ) return InitFromParams(ps.params, fallback) end + +_sampling_output_params(draw::ParamsWithStats) = draw.params +_sampling_output_params(draw::VarNamedTuple) = draw + +""" + returned(model::Model, chain::AbstractMCMC.SamplingOutput) + +Return a matrix of model return values evaluated at each draw's parameters. +""" +function returned( + model::Model, chain::AbstractMCMC.SamplingOutput{<:Union{ParamsWithStats,VarNamedTuple}} +) + return map(draw -> returned(model, _sampling_output_params(draw)), chain.samples) +end + +for f in (:logjoint, :logprior, :(Distributions.loglikelihood)) + @eval function $f( + model::Model, + chain::AbstractMCMC.SamplingOutput{<:Union{ParamsWithStats,VarNamedTuple}}, + ) + return map(draw -> $f(model, _sampling_output_params(draw)), chain.samples) + end +end + +""" + predict([rng::AbstractRNG,] model::Model, chain::AbstractMCMC.SamplingOutput; include_all=true) + +Sample predictions using each draw's parameters, drawing absent variables from their priors. + +Return a `SamplingOutput` with the input's iteration indices and freshly evaluated log +probabilities. Set `include_all=false` to omit parameters supplied by each input draw. +Sampling times and sampler states are not carried over to the predictions. +""" +function predict( + rng::Random.AbstractRNG, + model::Model, + chain::AbstractMCMC.SamplingOutput{<:Union{ParamsWithStats,VarNamedTuple}}; + include_all::Bool=true, +) + predictions = map(chain.samples) do draw + params = _sampling_output_params(draw) + vi = OnlyAccsVarInfo( + AccumulatorTuple( + LogPriorAccumulator(), + LogLikelihoodAccumulator(), + RawValueAccumulator(true), + ), + ) + _, vi = init!!(rng, model, vi, InitFromParams(params), UnlinkAll()) + prediction = ParamsWithStats(vi) + if include_all + prediction + else + predicted_params = VarNamedTuple() + for (vn, value) in pairs(prediction.params) + for (leaf, leaf_value) in AbstractPPL.varname_and_value_leaves(vn, value) + if !haskey(params, leaf) || ismissing(params[leaf]) + predicted_params = templated_setindex!!( + predicted_params, + leaf_value, + leaf, + prediction.params.data[AbstractPPL.getsym(leaf)], + ) + end + end + end + ParamsWithStats(densify!!(predicted_params), prediction.stats) + end + end + return AbstractMCMC.SamplingOutput(predictions; iterations=chain.iterations) +end +function predict( + model::Model, + chain::AbstractMCMC.SamplingOutput{<:Union{ParamsWithStats,VarNamedTuple}}; + kwargs..., +) + return predict(Random.default_rng(), model, chain; kwargs...) +end diff --git a/test/runtests.jl b/test/runtests.jl index 569391f53..4908857f5 100644 --- a/test/runtests.jl +++ b/test/runtests.jl @@ -44,6 +44,7 @@ Random.seed!(100) include("debug_utils.jl") include("submodels.jl") include("chains.jl") + include("samplingoutput.jl") end if GROUP in [TEST_GROUP_ALL, TEST_GROUP_GROUP2] diff --git a/test/samplingoutput.jl b/test/samplingoutput.jl new file mode 100644 index 000000000..a3011d6b9 --- /dev/null +++ b/test/samplingoutput.jl @@ -0,0 +1,44 @@ +module DynamicPPLSamplingOutputTests +using AbstractMCMC, DynamicPPL, Distributions, Random, Test + +@model function model(y) + x ~ Normal() + y ~ Normal(x, 1) + return x + y +end + +@testset "SamplingOutput model evaluation" begin + params = [VarNamedTuple(; x=i + j / 10) for i in 1:2, j in 1:2] + for draws in (params, map(p -> DynamicPPL.ParamsWithStats(p, (;)), params)) + chain = SamplingOutput(draws; iterations=3:2:5) + @test returned(model(1), chain) == map(p -> p[@varname(x)] + 1, params) + for f in (logjoint, logprior, loglikelihood) + @test f(model(1), chain) ≈ map(p -> f(model(1), p), params) + end + predictions = predict(Xoshiro(1), model(missing), chain) + @test size(predictions) == size(chain) + @test predictions.iterations == chain.iterations + @test all(ismissing, predictions.sampler_states) + @test map(p -> p.params[@varname(x)], predictions.samples) == + map(p -> p[@varname(x)], params) + again = predict(Xoshiro(1), model(missing), chain; include_all=false) + @test all(p -> !haskey(p.params, @varname(x)), again.samples) + @test map(p -> p.params[@varname(y)], again.samples) == + map(p -> p.params[@varname(y)], predictions.samples) + @test map(p -> p.stats, again.samples) == map(p -> p.stats, predictions.samples) + @test predict(model(missing), chain) isa SamplingOutput + @test_throws ErrorException returned(model(missing), chain) + end + @model function indexed_model() + x = zeros(2) + for i in eachindex(x) + x[i] ~ Normal() + end + end + params = DynamicPPL.templated_setindex!!(VarNamedTuple(), 4.0, @varname(x[1]), zeros(2)) + chain = SamplingOutput(fill(params, 1, 1)) + prediction = predict(Xoshiro(1), indexed_model(), chain; include_all=false)[1, 1] + @test !haskey(prediction.params, @varname(x[1])) + @test haskey(prediction.params, @varname(x[2])) +end +end From ec7d6f8463ba35e38f3676136ed7ca0535ca4bcc Mon Sep 17 00:00:00 2001 From: Hong Ge Date: Sun, 20 Sep 2026 21:02:16 +0100 Subject: [PATCH 02/11] tests: consolidate sampling output coverage in chains Assisted-by: Codex --- test/chains.jl | 43 +++++++++++++++++++++++++++++++++++++++++ test/runtests.jl | 1 - test/samplingoutput.jl | 44 ------------------------------------------ 3 files changed, 43 insertions(+), 45 deletions(-) delete mode 100644 test/samplingoutput.jl diff --git a/test/chains.jl b/test/chains.jl index 66362f0e1..4cc9883e5 100644 --- a/test/chains.jl +++ b/test/chains.jl @@ -4,9 +4,11 @@ using Dates: now @info "Testing $(@__FILE__)..." __now__ = now() +using AbstractMCMC: SamplingOutput using DynamicPPL using Distributions using LinearAlgebra +using Random: Xoshiro using Test @testset "ParamsWithStats from VarInfo" begin @@ -209,6 +211,47 @@ end @test isempty(offending) end +@testset "SamplingOutput model evaluation" begin + @model function model(y) + x ~ Normal() + y ~ Normal(x, 1) + return x + y + end + + params = [VarNamedTuple(; x=i + j / 10) for i in 1:2, j in 1:2] + for draws in (params, map(p -> DynamicPPL.ParamsWithStats(p, (;)), params)) + chain = SamplingOutput(draws; iterations=3:2:5) + @test returned(model(1), chain) == map(p -> p[@varname(x)] + 1, params) + for f in (logjoint, logprior, loglikelihood) + @test f(model(1), chain) ≈ map(p -> f(model(1), p), params) + end + predictions = predict(Xoshiro(1), model(missing), chain) + @test size(predictions) == size(chain) + @test predictions.iterations == chain.iterations + @test all(ismissing, predictions.sampler_states) + @test map(p -> p.params[@varname(x)], predictions.samples) == + map(p -> p[@varname(x)], params) + again = predict(Xoshiro(1), model(missing), chain; include_all=false) + @test all(p -> !haskey(p.params, @varname(x)), again.samples) + @test map(p -> p.params[@varname(y)], again.samples) == + map(p -> p.params[@varname(y)], predictions.samples) + @test map(p -> p.stats, again.samples) == map(p -> p.stats, predictions.samples) + @test predict(model(missing), chain) isa SamplingOutput + @test_throws ErrorException returned(model(missing), chain) + end + @model function indexed_model() + x = zeros(2) + for i in eachindex(x) + x[i] ~ Normal() + end + end + params = DynamicPPL.templated_setindex!!(VarNamedTuple(), 4.0, @varname(x[1]), zeros(2)) + chain = SamplingOutput(fill(params, 1, 1)) + prediction = predict(Xoshiro(1), indexed_model(), chain; include_all=false)[1, 1] + @test !haskey(prediction.params, @varname(x[1])) + @test haskey(prediction.params, @varname(x[2])) +end + @info "Completed $(@__FILE__) in $(now() - __now__)." end # module diff --git a/test/runtests.jl b/test/runtests.jl index 4908857f5..569391f53 100644 --- a/test/runtests.jl +++ b/test/runtests.jl @@ -44,7 +44,6 @@ Random.seed!(100) include("debug_utils.jl") include("submodels.jl") include("chains.jl") - include("samplingoutput.jl") end if GROUP in [TEST_GROUP_ALL, TEST_GROUP_GROUP2] diff --git a/test/samplingoutput.jl b/test/samplingoutput.jl deleted file mode 100644 index a3011d6b9..000000000 --- a/test/samplingoutput.jl +++ /dev/null @@ -1,44 +0,0 @@ -module DynamicPPLSamplingOutputTests -using AbstractMCMC, DynamicPPL, Distributions, Random, Test - -@model function model(y) - x ~ Normal() - y ~ Normal(x, 1) - return x + y -end - -@testset "SamplingOutput model evaluation" begin - params = [VarNamedTuple(; x=i + j / 10) for i in 1:2, j in 1:2] - for draws in (params, map(p -> DynamicPPL.ParamsWithStats(p, (;)), params)) - chain = SamplingOutput(draws; iterations=3:2:5) - @test returned(model(1), chain) == map(p -> p[@varname(x)] + 1, params) - for f in (logjoint, logprior, loglikelihood) - @test f(model(1), chain) ≈ map(p -> f(model(1), p), params) - end - predictions = predict(Xoshiro(1), model(missing), chain) - @test size(predictions) == size(chain) - @test predictions.iterations == chain.iterations - @test all(ismissing, predictions.sampler_states) - @test map(p -> p.params[@varname(x)], predictions.samples) == - map(p -> p[@varname(x)], params) - again = predict(Xoshiro(1), model(missing), chain; include_all=false) - @test all(p -> !haskey(p.params, @varname(x)), again.samples) - @test map(p -> p.params[@varname(y)], again.samples) == - map(p -> p.params[@varname(y)], predictions.samples) - @test map(p -> p.stats, again.samples) == map(p -> p.stats, predictions.samples) - @test predict(model(missing), chain) isa SamplingOutput - @test_throws ErrorException returned(model(missing), chain) - end - @model function indexed_model() - x = zeros(2) - for i in eachindex(x) - x[i] ~ Normal() - end - end - params = DynamicPPL.templated_setindex!!(VarNamedTuple(), 4.0, @varname(x[1]), zeros(2)) - chain = SamplingOutput(fill(params, 1, 1)) - prediction = predict(Xoshiro(1), indexed_model(), chain; include_all=false)[1, 1] - @test !haskey(prediction.params, @varname(x[1])) - @test haskey(prediction.params, @varname(x[2])) -end -end From 7d12c976d74cd8f542eece754ab0ea21b3588865 Mon Sep 17 00:00:00 2001 From: Hong Ge Date: Sun, 20 Sep 2026 21:02:37 +0100 Subject: [PATCH 03/11] chains: preserve structured values when filtering predictions Assisted-by: Codex --- src/chains.jl | 13 ++++++++++--- test/chains.jl | 19 +++++++++++++++++++ 2 files changed, 29 insertions(+), 3 deletions(-) diff --git a/src/chains.jl b/src/chains.jl index fe81f6143..32f45e2d2 100644 --- a/src/chains.jl +++ b/src/chains.jl @@ -318,9 +318,16 @@ function predict( prediction else predicted_params = VarNamedTuple() - for (vn, value) in pairs(prediction.params) - for (leaf, leaf_value) in AbstractPPL.varname_and_value_leaves(vn, value) - if !haskey(params, leaf) || ismissing(params[leaf]) + # Raw values retain sampled-variable boundaries before arrays are densified. + for (vn, value) in pairs(get_raw_values(vi)) + leaves = AbstractPPL.varname_and_value_leaves(vn, value) + keep_all = all( + p -> !haskey(params, first(p)) || ismissing(params[first(p)]), + leaves, + ) + retained = keep_all ? ((vn, value),) : leaves + for (leaf, leaf_value) in retained + if keep_all || !haskey(params, leaf) || ismissing(params[leaf]) predicted_params = templated_setindex!!( predicted_params, leaf_value, diff --git a/test/chains.jl b/test/chains.jl index 4cc9883e5..81c5cf04b 100644 --- a/test/chains.jl +++ b/test/chains.jl @@ -252,6 +252,25 @@ end @test haskey(prediction.params, @varname(x[2])) end +@testset "Prediction filtering preserves structured values" begin + dist = product_distribution((; a=Normal(), b=Bernoulli(), c=MvNormal(zeros(2), 1))) + @model function m(y=Any[missing, missing]) + y[1] ~ dist + y[2] ~ dist + return (y, Float32(1)) + end + partial = DynamicPPL.templated_setindex!!( + VarNamedTuple(), rand(Xoshiro(2), dist), @varname(y[1]), zeros(2) + ) + for params in (VarNamedTuple(), partial) + chain = SamplingOutput(fill(params, 1, 1)) + full = predict(Xoshiro(1), m(), chain) + filtered = predict(Xoshiro(1), m(), chain; include_all=false)[1, 1].params + @test haskey(filtered, @varname(y[1])) == !haskey(params, @varname(y[1])) + restored = densify!!(merge(params, filtered)) + @test logjoint(m(), restored) ≈ only(logjoint(m(), full)) + end +end @info "Completed $(@__FILE__) in $(now() - __now__)." end # module From 9c6784aeb4babd24a2de694f7b57c0cc579541a9 Mon Sep 17 00:00:00 2001 From: Hong Ge Date: Sun, 20 Sep 2026 21:17:36 +0100 Subject: [PATCH 04/11] chains: support pointwise log densities for sampling output Assisted-by: Codex --- src/chains.jl | 24 ++++++++++++++++++++++++ test/chains.jl | 24 ++++++++++++++++++++++++ 2 files changed, 48 insertions(+) diff --git a/src/chains.jl b/src/chains.jl index 32f45e2d2..19dcffdff 100644 --- a/src/chains.jl +++ b/src/chains.jl @@ -288,6 +288,30 @@ for f in (:logjoint, :logprior, :(Distributions.loglikelihood)) end end +for f in (:pointwise_logdensities, :pointwise_loglikelihoods, :pointwise_prior_logdensities) + @eval begin + """ + $($f)(model::Model, chain::AbstractMCMC.SamplingOutput; factorize=false) + + Evaluate `$($f)` at each draw and return a `SamplingOutput` of `VarNamedTuple`s. + Preserve iteration indices; omit sampling statistics and sampler states. + All model parameters must be supplied in each draw. + + $(_FACTORIZE_KWARG_DOC) + """ + function $f( + model::Model, + chain::AbstractMCMC.SamplingOutput{<:Union{ParamsWithStats,VarNamedTuple}}; + factorize=false, + ) + densities = map(chain.samples) do draw + $f(model, InitFromParams(_sampling_output_params(draw), nothing); factorize) + end + return AbstractMCMC.SamplingOutput(densities; iterations=chain.iterations) + end + end +end + """ predict([rng::AbstractRNG,] model::Model, chain::AbstractMCMC.SamplingOutput; include_all=true) diff --git a/test/chains.jl b/test/chains.jl index 81c5cf04b..2c2e97951 100644 --- a/test/chains.jl +++ b/test/chains.jl @@ -271,6 +271,30 @@ end @test logjoint(m(), restored) ≈ only(logjoint(m(), full)) end end +@testset "SamplingOutput pointwise log densities" begin + @model function pointwise_model(y) + x ~ MvNormal(zeros(2), I) + return y ~ MvNormal(x, I) + end + model = pointwise_model([0.5, -0.5]) + params = [VarNamedTuple(; x=[i / 10, j / 10]) for i in 1:2, j in 1:2] + fs = (pointwise_logdensities, pointwise_loglikelihoods, pointwise_prior_logdensities) + for draws in (params, map(p -> ParamsWithStats(p, (; ignored=1)), params)) + chain = SamplingOutput(draws; iterations=3:2:5, sampler_states=[:a, :b]) + for f in fs, factorize in (false, true) + result = f(model, chain; factorize) + @test result isa SamplingOutput{<:VarNamedTuple} + @test size(result) == size(chain) + @test result.iterations == chain.iterations + @test all(ismissing, [result.sampler_states; result.sampling_stats]) + for i in eachindex(params) + @test result[i] == f(model, InitFromParams(params[i], nothing); factorize) + end + incomplete = SamplingOutput(fill(VarNamedTuple(), 1, 1)) + @test_throws ErrorException f(model, incomplete) + end + end +end @info "Completed $(@__FILE__) in $(now() - __now__)." end # module From 2ad64ce3ab9e2a461cb4088c6b03af1c18c29a7f Mon Sep 17 00:00:00 2001 From: Hong Ge Date: Sun, 20 Sep 2026 21:20:54 +0100 Subject: [PATCH 05/11] chains: support conversion and handle mixed and empty draws Assisted-by: Codex --- HISTORY.md | 4 ++++ Project.toml | 2 +- ext/DynamicPPLMCMCChainsExt.jl | 29 +++++++++++++++++++++++++++++ src/chains.jl | 24 ++++++++++++++++++++++++ test/ext/DynamicPPLMCMCChainsExt.jl | 13 +++++++++++++ 5 files changed, 71 insertions(+), 1 deletion(-) diff --git a/HISTORY.md b/HISTORY.md index 327c7bab7..054574357 100644 --- a/HISTORY.md +++ b/HISTORY.md @@ -1,3 +1,7 @@ +# 0.42.13 + +Added `SamplingOutput` support for `pointwise_logdensities`, `pointwise_loglikelihoods`, and `pointwise_prior_logdensities`, plus conversion to `MCMCChains.Chains`. + # 0.42.12 `check_model` now warns when a latent tilde statement overwrites a value computed from a model input. diff --git a/Project.toml b/Project.toml index b7861dbf7..5a0b4b4fb 100644 --- a/Project.toml +++ b/Project.toml @@ -1,6 +1,6 @@ name = "DynamicPPL" uuid = "366bfd00-2699-11ea-058f-f148b4cae6d8" -version = "0.42.12" +version = "0.42.13" [deps] ADTypes = "47edcb42-4c32-4615-8424-f2b9edc5f35b" diff --git a/ext/DynamicPPLMCMCChainsExt.jl b/ext/DynamicPPLMCMCChainsExt.jl index a765119c2..379006ca8 100644 --- a/ext/DynamicPPLMCMCChainsExt.jl +++ b/ext/DynamicPPLMCMCChainsExt.jl @@ -1,6 +1,7 @@ module DynamicPPLMCMCChainsExt using DynamicPPL: DynamicPPL, AbstractPPL, AbstractMCMC, Random +using DynamicPPL: ParamsWithStats, VarNamedTuple using BangBang: setindex!! using MCMCChains: MCMCChains @@ -207,6 +208,34 @@ function AbstractMCMC.bundle_samples( return sort_chain ? sort(chain) : chain end +mcmcchains_sample(draw::VarNamedTuple) = ParamsWithStats(draw, (;)) +function mcmcchains_sample(draw::ParamsWithStats) + leaves = AbstractPPL.varname_and_value_leaves(draw.stats) + return ParamsWithStats(draw.params, (; (Symbol(vn) => v for (vn, v) in leaves)...)) +end + +# Prefer convert to direct bundling: array-valued statistics become scalar columns +# instead of missing; iteration indices, timing, and saved states are retained. +function Base.convert( + ::Type{MCMCChains.Chains}, + output::AbstractMCMC.SamplingOutput{<:Union{ParamsWithStats,VarNamedTuple}}, +) + c = AbstractMCMC.from_samples(MCMCChains.Chains, map(mcmcchains_sample, output.samples)) + c = MCMCChains.setrange(c, output.iterations) + info = c.info + for (key, f) in ((:start_time, :start), (:stop_time, :stop), (:samplerstate, nothing)) + values = if f === nothing + output.sampler_states + else + map(s -> ismissing(s) ? missing : getproperty(s, f), output.sampling_stats) + end + any(!ismissing, values) || continue + value = size(output, 2) == 1 ? only(values) : values + info = merge(info, NamedTuple{(key,)}((value,))) + end + return MCMCChains.setinfo(c, info) +end + """ reevaluate_with_chain( rng::AbstractRNG, diff --git a/src/chains.jl b/src/chains.jl index 19dcffdff..e384b1026 100644 --- a/src/chains.jl +++ b/src/chains.jl @@ -268,6 +268,29 @@ end _sampling_output_params(draw::ParamsWithStats) = draw.params _sampling_output_params(draw::VarNamedTuple) = draw +""" + convert(::Type{T}, output::AbstractMCMC.SamplingOutput) + +Convert structured `SamplingOutput` draws to an `AbstractMCMC.AbstractChains` type using +`AbstractMCMC.from_samples`. This fallback converts only the draws and drops chain-level +metadata. Chain packages that preserve metadata should overload this method, or extend +`from_samples` to accept `iterations`, `sampling_stats`, and `sampler_states` and forward +those fields from their overload. + +```julia +AbstractMCMC.from_samples(::Type{T}, draws; iterations, sampling_stats, sampler_states) where {T} = + Chain(draws; iterations, sampling_stats, sampler_states) # package-specific constructor +Base.convert(::Type{T}, o::AbstractMCMC.SamplingOutput) where {T<:AbstractMCMC.AbstractChains} = + AbstractMCMC.from_samples(T, o.samples; iterations=o.iterations, + sampling_stats=o.sampling_stats, sampler_states=o.sampler_states) +``` +""" +function Base.convert( + ::Type{T}, output::AbstractMCMC.SamplingOutput{<:Union{ParamsWithStats,VarNamedTuple}} +) where {T<:AbstractChains} + return AbstractMCMC.from_samples(T, output.samples) +end + """ returned(model::Model, chain::AbstractMCMC.SamplingOutput) @@ -345,6 +368,7 @@ function predict( # Raw values retain sampled-variable boundaries before arrays are densified. for (vn, value) in pairs(get_raw_values(vi)) leaves = AbstractPPL.varname_and_value_leaves(vn, value) + isempty(leaves) && haskey(params, vn) && !ismissing(params[vn]) && continue keep_all = all( p -> !haskey(params, first(p)) || ismissing(params[first(p)]), leaves, diff --git a/test/ext/DynamicPPLMCMCChainsExt.jl b/test/ext/DynamicPPLMCMCChainsExt.jl index e207ad6bf..b08b5f1ba 100644 --- a/test/ext/DynamicPPLMCMCChainsExt.jl +++ b/test/ext/DynamicPPLMCMCChainsExt.jl @@ -21,6 +21,19 @@ function make_chain_from_prior(model::Model, n_iters::Int) end @testset "DynamicPPLMCMCChainsExt" begin + @testset "SamplingOutput conversion" begin + samples = fill(ParamsWithStats(VarNamedTuple(; x=1), (; ess=[2, 3])), 1, 1) + output = AbstractMCMC.SamplingOutput( + samples; iterations=3:2:3, sampler_states=[:saved] + ) + chain = convert(MCMCChains.Chains, output) + @test output[1, 1].stats.ess == [2, 3] + @test only(chain[Symbol("ess[2]")]) == 3 + @test only(AbstractMCMC.from_samples(MCMCChains.Chains, samples)[:ess]) === missing + @test range(chain) == output.iterations + @test chain.info.samplerstate === :saved + end + @testset "from_samples" begin @model function f(z) x ~ Normal() From 2f1bcdcc46a3d4b8bbc9e06f3f27e52cc91ba7cc Mon Sep 17 00:00:00 2001 From: Hong Ge Date: Sun, 20 Sep 2026 22:25:10 +0100 Subject: [PATCH 06/11] chains: preserve identity and supported conversion metadata Assisted-by: Codex --- src/chains.jl | 12 +++++++++--- test/chains.jl | 6 ++++++ 2 files changed, 15 insertions(+), 3 deletions(-) diff --git a/src/chains.jl b/src/chains.jl index e384b1026..a9110642b 100644 --- a/src/chains.jl +++ b/src/chains.jl @@ -272,8 +272,8 @@ _sampling_output_params(draw::VarNamedTuple) = draw convert(::Type{T}, output::AbstractMCMC.SamplingOutput) Convert structured `SamplingOutput` draws to an `AbstractMCMC.AbstractChains` type using -`AbstractMCMC.from_samples`. This fallback converts only the draws and drops chain-level -metadata. Chain packages that preserve metadata should overload this method, or extend +`AbstractMCMC.from_samples`. Unsupported metadata is intentionally omitted. +Chain packages can overload this method, or support metadata by extending `from_samples` to accept `iterations`, `sampling_stats`, and `sampler_states` and forward those fields from their overload. @@ -288,7 +288,13 @@ Base.convert(::Type{T}, o::AbstractMCMC.SamplingOutput) where {T<:AbstractMCMC.A function Base.convert( ::Type{T}, output::AbstractMCMC.SamplingOutput{<:Union{ParamsWithStats,VarNamedTuple}} ) where {T<:AbstractChains} - return AbstractMCMC.from_samples(T, output.samples) + output isa T && return output + metadata = (; output.iterations, output.sampling_stats, output.sampler_states) + signature = Tuple{Type{T},typeof(output.samples)} + supported = filter(keys(metadata)) do key + hasmethod(AbstractMCMC.from_samples, signature, (key,)) + end + return AbstractMCMC.from_samples(T, output.samples; NamedTuple{supported}(metadata)...) end """ diff --git a/test/chains.jl b/test/chains.jl index 2c2e97951..4fef88461 100644 --- a/test/chains.jl +++ b/test/chains.jl @@ -281,6 +281,12 @@ end fs = (pointwise_logdensities, pointwise_loglikelihoods, pointwise_prior_logdensities) for draws in (params, map(p -> ParamsWithStats(p, (; ignored=1)), params)) chain = SamplingOutput(draws; iterations=3:2:5, sampler_states=[:a, :b]) + for T in (SamplingOutput, typeof(chain), DynamicPPL.AbstractChains) + @test convert(T, chain) === chain + end + widened = convert(SamplingOutput{Any}, chain) + @test widened.iterations == chain.iterations + @test widened.sampler_states == chain.sampler_states for f in fs, factorize in (false, true) result = f(model, chain; factorize) @test result isa SamplingOutput{<:VarNamedTuple} From 6244fc16ea68e0586cf41692dc50717e6f1b9ab3 Mon Sep 17 00:00:00 2001 From: Hong Ge Date: Sun, 20 Sep 2026 22:39:24 +0100 Subject: [PATCH 07/11] tests: consolidate sampling output coverage Assisted-by: Codex --- test/chains.jl | 89 +++++++++++++------------------------------------- 1 file changed, 23 insertions(+), 66 deletions(-) diff --git a/test/chains.jl b/test/chains.jl index 4fef88461..5cf054e6f 100644 --- a/test/chains.jl +++ b/test/chains.jl @@ -211,53 +211,40 @@ end @test isempty(offending) end -@testset "SamplingOutput model evaluation" begin - @model function model(y) - x ~ Normal() - y ~ Normal(x, 1) +@testset "SamplingOutput evaluation and conversion" begin + @model function m(y=ones(2)) + x ~ MvNormal(zeros(2), I) + y ~ MvNormal(x, I) return x + y end - - params = [VarNamedTuple(; x=i + j / 10) for i in 1:2, j in 1:2] - for draws in (params, map(p -> DynamicPPL.ParamsWithStats(p, (;)), params)) - chain = SamplingOutput(draws; iterations=3:2:5) - @test returned(model(1), chain) == map(p -> p[@varname(x)] + 1, params) + params = [VarNamedTuple(; x=[i / 10, j / 10]) for i in 1:2, j in 1:2] + fs = (pointwise_logdensities, pointwise_loglikelihoods, pointwise_prior_logdensities) + for draws in (params, map(p -> ParamsWithStats(p, (;)), params)) + chain = SamplingOutput(draws; iterations=3:2:5, sampler_states=[:a, :b]) + @test returned(m(), chain) == map(p -> p[@varname(x)] + ones(2), params) + for T in (SamplingOutput, typeof(chain), DynamicPPL.AbstractChains) + @test convert(T, chain) === chain + end + widened = convert(SamplingOutput{Any}, chain) + @test (widened.iterations, widened.sampler_states) == + (chain.iterations, chain.sampler_states) for f in (logjoint, logprior, loglikelihood) - @test f(model(1), chain) ≈ map(p -> f(model(1), p), params) + @test f(m(), chain) ≈ map(p -> f(m(), p), params) end - predictions = predict(Xoshiro(1), model(missing), chain) - @test size(predictions) == size(chain) - @test predictions.iterations == chain.iterations - @test all(ismissing, predictions.sampler_states) - @test map(p -> p.params[@varname(x)], predictions.samples) == - map(p -> p[@varname(x)], params) - again = predict(Xoshiro(1), model(missing), chain; include_all=false) - @test all(p -> !haskey(p.params, @varname(x)), again.samples) - @test map(p -> p.params[@varname(y)], again.samples) == - map(p -> p.params[@varname(y)], predictions.samples) - @test map(p -> p.stats, again.samples) == map(p -> p.stats, predictions.samples) - @test predict(model(missing), chain) isa SamplingOutput - @test_throws ErrorException returned(model(missing), chain) - end - @model function indexed_model() - x = zeros(2) - for i in eachindex(x) - x[i] ~ Normal() + for f in fs, factorize in (false, true) + result = f(m(), chain; factorize) + @test (size(result), result.iterations) == (size(chain), chain.iterations) + @test result.samples == + map(p -> f(m(), InitFromParams(p, nothing); factorize), params) + @test_throws ErrorException f(m(), SamplingOutput(fill(VarNamedTuple(), 1, 1))) end end - params = DynamicPPL.templated_setindex!!(VarNamedTuple(), 4.0, @varname(x[1]), zeros(2)) - chain = SamplingOutput(fill(params, 1, 1)) - prediction = predict(Xoshiro(1), indexed_model(), chain; include_all=false)[1, 1] - @test !haskey(prediction.params, @varname(x[1])) - @test haskey(prediction.params, @varname(x[2])) end - @testset "Prediction filtering preserves structured values" begin dist = product_distribution((; a=Normal(), b=Bernoulli(), c=MvNormal(zeros(2), 1))) @model function m(y=Any[missing, missing]) y[1] ~ dist - y[2] ~ dist - return (y, Float32(1)) + return y[2] ~ dist end partial = DynamicPPL.templated_setindex!!( VarNamedTuple(), rand(Xoshiro(2), dist), @varname(y[1]), zeros(2) @@ -271,36 +258,6 @@ end @test logjoint(m(), restored) ≈ only(logjoint(m(), full)) end end -@testset "SamplingOutput pointwise log densities" begin - @model function pointwise_model(y) - x ~ MvNormal(zeros(2), I) - return y ~ MvNormal(x, I) - end - model = pointwise_model([0.5, -0.5]) - params = [VarNamedTuple(; x=[i / 10, j / 10]) for i in 1:2, j in 1:2] - fs = (pointwise_logdensities, pointwise_loglikelihoods, pointwise_prior_logdensities) - for draws in (params, map(p -> ParamsWithStats(p, (; ignored=1)), params)) - chain = SamplingOutput(draws; iterations=3:2:5, sampler_states=[:a, :b]) - for T in (SamplingOutput, typeof(chain), DynamicPPL.AbstractChains) - @test convert(T, chain) === chain - end - widened = convert(SamplingOutput{Any}, chain) - @test widened.iterations == chain.iterations - @test widened.sampler_states == chain.sampler_states - for f in fs, factorize in (false, true) - result = f(model, chain; factorize) - @test result isa SamplingOutput{<:VarNamedTuple} - @test size(result) == size(chain) - @test result.iterations == chain.iterations - @test all(ismissing, [result.sampler_states; result.sampling_stats]) - for i in eachindex(params) - @test result[i] == f(model, InitFromParams(params[i], nothing); factorize) - end - incomplete = SamplingOutput(fill(VarNamedTuple(), 1, 1)) - @test_throws ErrorException f(model, incomplete) - end - end -end @info "Completed $(@__FILE__) in $(now() - __now__)." end # module From 5976df2eb23141dfa16985f71bb2a07ec7200fd5 Mon Sep 17 00:00:00 2001 From: Hong Ge Date: Sun, 20 Sep 2026 22:39:24 +0100 Subject: [PATCH 08/11] chains: normalize iteration indices for MCMCChains Assisted-by: Codex --- ext/DynamicPPLMCMCChainsExt.jl | 2 +- test/ext/DynamicPPLMCMCChainsExt.jl | 3 ++- 2 files changed, 3 insertions(+), 2 deletions(-) diff --git a/ext/DynamicPPLMCMCChainsExt.jl b/ext/DynamicPPLMCMCChainsExt.jl index 379006ca8..6ef92f4d1 100644 --- a/ext/DynamicPPLMCMCChainsExt.jl +++ b/ext/DynamicPPLMCMCChainsExt.jl @@ -221,7 +221,7 @@ function Base.convert( output::AbstractMCMC.SamplingOutput{<:Union{ParamsWithStats,VarNamedTuple}}, ) c = AbstractMCMC.from_samples(MCMCChains.Chains, map(mcmcchains_sample, output.samples)) - c = MCMCChains.setrange(c, output.iterations) + c = MCMCChains.setrange(c, Int.(output.iterations)) info = c.info for (key, f) in ((:start_time, :start), (:stop_time, :stop), (:samplerstate, nothing)) values = if f === nothing diff --git a/test/ext/DynamicPPLMCMCChainsExt.jl b/test/ext/DynamicPPLMCMCChainsExt.jl index b08b5f1ba..33ac84b6d 100644 --- a/test/ext/DynamicPPLMCMCChainsExt.jl +++ b/test/ext/DynamicPPLMCMCChainsExt.jl @@ -23,8 +23,9 @@ end @testset "DynamicPPLMCMCChainsExt" begin @testset "SamplingOutput conversion" begin samples = fill(ParamsWithStats(VarNamedTuple(; x=1), (; ess=[2, 3])), 1, 1) + # Exercise conversion to the native Int indices required by MCMCChains. output = AbstractMCMC.SamplingOutput( - samples; iterations=3:2:3, sampler_states=[:saved] + samples; iterations=big.(3:2:3), sampler_states=[:saved] ) chain = convert(MCMCChains.Chains, output) @test output[1, 1].stats.ess == [2, 3] From f55853142fe37a887f2d081af7c7789bb516f252 Mon Sep 17 00:00:00 2001 From: Hong Ge Date: Sun, 20 Sep 2026 22:50:23 +0100 Subject: [PATCH 09/11] docs: link sampling output history entry to PR Assisted-by: Codex --- HISTORY.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/HISTORY.md b/HISTORY.md index 054574357..7b36abd97 100644 --- a/HISTORY.md +++ b/HISTORY.md @@ -1,6 +1,6 @@ # 0.42.13 -Added `SamplingOutput` support for `pointwise_logdensities`, `pointwise_loglikelihoods`, and `pointwise_prior_logdensities`, plus conversion to `MCMCChains.Chains`. +Added `SamplingOutput` support for `pointwise_logdensities`, `pointwise_loglikelihoods`, and `pointwise_prior_logdensities`, plus conversion to `MCMCChains.Chains`. See [#1506](https://github.com/TuringLang/DynamicPPL.jl/pull/1506). # 0.42.12 From 772f7e6e10768b50308a9cabaf3c6516054a6d02 Mon Sep 17 00:00:00 2001 From: Hong Ge Date: Sun, 20 Sep 2026 23:11:56 +0100 Subject: [PATCH 10/11] release: bump to 0.42.14 and fix API documentation coverage Assisted-by: Codex --- Project.toml | 2 +- docs/src/api.md | 5 ++--- 2 files changed, 3 insertions(+), 4 deletions(-) diff --git a/Project.toml b/Project.toml index 5a0b4b4fb..51b88788f 100644 --- a/Project.toml +++ b/Project.toml @@ -1,6 +1,6 @@ name = "DynamicPPL" uuid = "366bfd00-2699-11ea-058f-f148b4cae6d8" -version = "0.42.13" +version = "0.42.14" [deps] ADTypes = "47edcb42-4c32-4615-8424-f2b9edc5f35b" diff --git a/docs/src/api.md b/docs/src/api.md index 8745477b4..7d024a1ad 100644 --- a/docs/src/api.md +++ b/docs/src/api.md @@ -179,11 +179,10 @@ It is possible to manually increase (or decrease) the accumulated log likelihood @addlogprob! ``` -Return values of the model function can be obtained with [`returned(model, sample)`](@ref), where `sample` is either a `MCMCChains.Chains` object (which represents a collection of samples), or a single sample represented as a `NamedTuple` or a dictionary of VarNames. +Return values of the model function can be obtained with [`returned(model, sample)`](@ref), where `sample` is a collection of samples in an `AbstractMCMC.SamplingOutput` or `MCMCChains.Chains`, or a single sample represented as a `VarNamedTuple`, `NamedTuple`, or dictionary of VarNames. ```@docs -returned(::DynamicPPL.Model, ::MCMCChains.Chains) -returned(::DynamicPPL.Model, ::Union{NamedTuple,AbstractDict{<:VarName}}) +returned ``` For a chain of samples, one can compute the pointwise log-likelihoods of each observed random variable with [`pointwise_loglikelihoods`](@ref). Similarly, the log-densities of the priors using From 7101ad6e376b41b6828ccf304bcad1838b23006e Mon Sep 17 00:00:00 2001 From: Hong Ge Date: Mon, 21 Sep 2026 00:51:03 +0100 Subject: [PATCH 11/11] chains: align prediction defaults and correct release history Assisted-by: Codex --- HISTORY.md | 2 +- src/chains.jl | 6 +++--- test/chains.jl | 4 ++-- 3 files changed, 6 insertions(+), 6 deletions(-) diff --git a/HISTORY.md b/HISTORY.md index 6b398e43b..7636941e1 100644 --- a/HISTORY.md +++ b/HISTORY.md @@ -1,10 +1,10 @@ # 0.42.14 Added `SamplingOutput` support for `pointwise_logdensities`, `pointwise_loglikelihoods`, and `pointwise_prior_logdensities`, plus conversion to `MCMCChains.Chains`. See [#1506](https://github.com/TuringLang/DynamicPPL.jl/pull/1506). -Model bodies no longer contain a `try` block, so Libtask can tape them again. # 0.42.13 +Model bodies no longer contain a `try` block, so Libtask can tape them again. Particle samplers such as `SMC`, `PG` and `CSMC` threw while building a `TapedTask` on 0.42.12, whether or not model checking was enabled. See [#1487](https://github.com/TuringLang/DynamicPPL.jl/issues/1487). diff --git a/src/chains.jl b/src/chains.jl index a9110642b..0b9de86c3 100644 --- a/src/chains.jl +++ b/src/chains.jl @@ -342,19 +342,19 @@ for f in (:pointwise_logdensities, :pointwise_loglikelihoods, :pointwise_prior_l end """ - predict([rng::AbstractRNG,] model::Model, chain::AbstractMCMC.SamplingOutput; include_all=true) + predict([rng::AbstractRNG,] model::Model, chain::AbstractMCMC.SamplingOutput; include_all=false) Sample predictions using each draw's parameters, drawing absent variables from their priors. Return a `SamplingOutput` with the input's iteration indices and freshly evaluated log -probabilities. Set `include_all=false` to omit parameters supplied by each input draw. +probabilities. Set `include_all=true` to retain parameters supplied by each input draw. Sampling times and sampler states are not carried over to the predictions. """ function predict( rng::Random.AbstractRNG, model::Model, chain::AbstractMCMC.SamplingOutput{<:Union{ParamsWithStats,VarNamedTuple}}; - include_all::Bool=true, + include_all::Bool=false, ) predictions = map(chain.samples) do draw params = _sampling_output_params(draw) diff --git a/test/chains.jl b/test/chains.jl index 5cf054e6f..525ab060c 100644 --- a/test/chains.jl +++ b/test/chains.jl @@ -251,8 +251,8 @@ end ) for params in (VarNamedTuple(), partial) chain = SamplingOutput(fill(params, 1, 1)) - full = predict(Xoshiro(1), m(), chain) - filtered = predict(Xoshiro(1), m(), chain; include_all=false)[1, 1].params + full = predict(Xoshiro(1), m(), chain; include_all=true) + filtered = predict(Xoshiro(1), m(), chain)[1, 1].params @test haskey(filtered, @varname(y[1])) == !haskey(params, @varname(y[1])) restored = densify!!(merge(params, filtered)) @test logjoint(m(), restored) ≈ only(logjoint(m(), full))