diff --git a/HISTORY.md b/HISTORY.md index d29e71fc3..7636941e1 100644 --- a/HISTORY.md +++ b/HISTORY.md @@ -1,3 +1,7 @@ +# 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). + # 0.42.13 Model bodies no longer contain a `try` block, so Libtask can tape them again. diff --git a/Project.toml b/Project.toml index 2933c24f9..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" @@ -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/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 diff --git a/ext/DynamicPPLMCMCChainsExt.jl b/ext/DynamicPPLMCMCChainsExt.jl index e16bf8824..e44467130 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, Int.(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 + """ chunk_ranges(n::Int, nchunks::Int) diff --git a/src/chains.jl b/src/chains.jl index 219ccb444..0b9de86c3 100644 --- a/src/chains.jl +++ b/src/chains.jl @@ -264,3 +264,142 @@ function InitFromParams( ) return InitFromParams(ps.params, fallback) 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`. 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. + +```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} + 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 + +""" + 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 + +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=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=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=false, +) + 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() + # 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, + ) + 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, + 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/chains.jl b/test/chains.jl index 66362f0e1..525ab060c 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,53 @@ end @test isempty(offending) end +@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 / 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(m(), chain) ≈ map(p -> f(m(), p), params) + end + 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 +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 + return y[2] ~ dist + 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; 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)) + end +end @info "Completed $(@__FILE__) in $(now() - __now__)." end # module diff --git a/test/ext/DynamicPPLMCMCChainsExt.jl b/test/ext/DynamicPPLMCMCChainsExt.jl index e207ad6bf..33ac84b6d 100644 --- a/test/ext/DynamicPPLMCMCChainsExt.jl +++ b/test/ext/DynamicPPLMCMCChainsExt.jl @@ -21,6 +21,20 @@ 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) + # Exercise conversion to the native Int indices required by MCMCChains. + output = AbstractMCMC.SamplingOutput( + samples; iterations=big.(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()