Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions HISTORY.md
Original file line number Diff line number Diff line change
@@ -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.
Expand Down
4 changes: 2 additions & 2 deletions Project.toml
Original file line number Diff line number Diff line change
@@ -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"
Expand Down Expand Up @@ -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"
Expand Down
5 changes: 2 additions & 3 deletions docs/src/api.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
29 changes: 29 additions & 0 deletions ext/DynamicPPLMCMCChainsExt.jl
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
module DynamicPPLMCMCChainsExt

using DynamicPPL: DynamicPPL, AbstractPPL, AbstractMCMC, Random
using DynamicPPL: ParamsWithStats, VarNamedTuple
using BangBang: setindex!!
using MCMCChains: MCMCChains

Expand Down Expand Up @@ -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)

Expand Down
139 changes: 139 additions & 0 deletions src/chains.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
49 changes: 49 additions & 0 deletions test/chains.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
14 changes: 14 additions & 0 deletions test/ext/DynamicPPLMCMCChainsExt.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down
Loading