Skip to content
Open
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
Expand Up @@ -30,6 +30,10 @@ For `@model fs(y; kw...)`, `fs(1.0; z=2, w=3).defaults` changes from `(z=2, w=3)

Integer indices into NamedTuples are rejected in binding addresses and LHS variables: `x[1]` on a NamedTuple → `x.a`; Tuples retain integer indices. See [#1501](https://github.com/TuringLang/DynamicPPL.jl/pull/1501).

`LogDensityFunction(model; rng)` now shares the supplied RNG across construction, evaluation, AD preparation, and parameter sampling. Use `rand(__context__.rng, ...)` for model-body draws.

`DynamicPPL.logdensity_internal` takes an optional final `rng` argument, defaulting to the task-local `Random.default_rng()`; append an RNG to the `context` tuple of `AbstractPPL.prepare(DynamicPPL.logdensity_internal, x; context=...)` to control explicit draws in the model body. See [#1504](https://github.com/TuringLang/DynamicPPL.jl/pull/1504).

`DefaultContext()` / `InitContext(...)` → `Context(rng, init_strategy, transform_strategy)`: specify parameter inputs and output transforms explicitly. See [#1503](https://github.com/TuringLang/DynamicPPL.jl/pull/1503).

`contextualize` and `Model.context` are removed: `contextualize(m, ctx); evaluate!!(m, vi)` → `evaluate!!(m, ctx, vi)`. The two-argument `evaluate!!(m, vi)` is removed. See [#1503](https://github.com/TuringLang/DynamicPPL.jl/pull/1503).
Expand Down
3 changes: 2 additions & 1 deletion docs/src/evaluation.md
Original file line number Diff line number Diff line change
Expand Up @@ -118,7 +118,8 @@ record a [`VectorValueAccumulator`](@ref) and pass `get_vector_values(recorded)`

This separation specifies data flow, not purity: evaluation can advance the RNG, and
ordinary Julia mutations in a model body still take effect.
For density evaluation, LHS variables get their values from supplied parameters rather than sampling.
For density evaluation, LHS variables get their values from supplied parameters rather than sampling;
see [Randomness in density evaluation](@ref ldf-rng).

## Accumulators

Expand Down
41 changes: 40 additions & 1 deletion docs/src/ldf/overview.md
Original file line number Diff line number Diff line change
Expand Up @@ -54,7 +54,7 @@ The actual implementation of `LogDensityFunction` is a bit more complex than thi
The main constructor for `LogDensityFunction` is

```julia
LogDensityFunction(model, logdensityfunc, transform_strategy; adtype)
LogDensityFunction(model, logdensityfunc, transform_strategy; adtype, rng)
```

`model` is of course the model itself, but the other arguments deserve more explanation.
Expand Down Expand Up @@ -116,6 +116,45 @@ LogDensityProblems.logdensity_and_gradient(ldf, [3.0, 4.0])

Other functions such as `LogDensityProblems.capabilities` and `LogDensityProblems.dimension` will also work as expected with `LogDensityFunction`.

## [Randomness in density evaluation](@id ldf-rng)

Pass `rng` to `LogDensityFunction` to control model-body randomness. During density
evaluation, `x ~ Normal()` reads `x` from the supplied parameter vector; it does not
sample or advance the RNG. Only explicit random draws in the model body or its callees
advance it. For example:

```@example ldf-rng
using DynamicPPL, Distributions, Random, LogDensityProblems

@model function random_observation(y)
x ~ Normal() # Read from the parameter vector.
i = rand(__context__.rng, eachindex(y)) # Advance the supplied RNG.
return y[i] ~ Normal(x, 1)
end

rng = Xoshiro(42)
ldf = LogDensityFunction(random_observation([1.0, 2.0, 3.0]); rng)
LogDensityProblems.logdensity(ldf, [0.5])
```

The same RNG reaches nested submodels and AD evaluation. It is shared with the caller,
not copied or reset, so repeated calls can choose different observations and return
different densities at identical parameters. Ordinary `rand(...)` without an RNG still
uses Julia's default RNG; use `rand(__context__.rng, ...)` for evaluation-controlled draws.

Construction is distinct from density evaluation: it may sample parameters to determine
the vector layout. AD preparation may also execute the model and advance the RNG.
`rand(ldf)` samples parameters using `ldf.rng`; `rand(other_rng, ldf)` uses `other_rng`.

RNG control does not make stochastic densities suitable for ordinary HMC or NUTS,
which require a deterministic target. AD backends may evaluate a model multiple times
for one gradient, drawing different values, or replay a compiled tape without drawing
again. RNG consumption therefore depends on the backend, and gradients may not
correspond to a single realisation. Random operations also need backend support;
passing an RNG does not supply missing differentiation rules. For deterministic
inference, select random data or other auxiliary randomness outside density evaluation
and supply it to the model.

## Is `LogDensityFunction` less powerful than model evaluation?

Given that the core purpose of `LogDensityFunction` is to evaluate a function mapping from vectors to log-densities, it might seem that it contains less information than the original model itself.
Expand Down
4 changes: 4 additions & 0 deletions ext/DynamicPPLMooncakeExt.jl
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,10 @@ module DynamicPPLMooncakeExt
using DynamicPPL: DynamicPPL, is_transformed
using AbstractPPL: AbstractPPL
using Mooncake: Mooncake
using Random: AbstractRNG

# RNG state is not a differentiable input to model evaluation.
Mooncake.tangent_type(::Type{<:AbstractRNG}) = Mooncake.NoTangent

Mooncake.@is_primitive Mooncake.MinimalCtx Tuple{
DynamicPPL._StanDifferentiableFunction,<:AbstractArray{<:Real}
Expand Down
13 changes: 6 additions & 7 deletions src/accumulators/fixed_transforms.jl
Original file line number Diff line number Diff line change
Expand Up @@ -55,26 +55,25 @@ end
"""
get_fixed_transforms(
model::DynamicPPL.Model,
transform_strategy::AbstractTransformStrategy
transform_strategy::AbstractTransformStrategy;
rng::Random.AbstractRNG=Random.default_rng(),
)

Extract the fixed transforms for all variables in a model by running the model with the
given transform strategy.

Note that, even though this method evaluates the model once, this method does *not* accept
an RNG argument to control that evaluation. This is because the fixed transforms are
supposed to be *fixed*, i.e., they should not depend on random choices made during model
execution!
The RNG controls this evaluation, but the transforms must not depend on its random choices.

If you are unsure about whether the transforms for your model are fixed, you can use
[`DynamicPPL.DebugUtils.has_static_constraints`](@ref). Note though that this relies on
executing the model multiple times and is thus not foolproof (it may be that your transforms
just happen to be the same each time).
"""
function get_fixed_transforms(
model::DynamicPPL.Model, transform_strategy::AbstractTransformStrategy
model::DynamicPPL.Model,
transform_strategy::AbstractTransformStrategy;
rng::Random.AbstractRNG=Random.default_rng(),
)
rng = Random.default_rng()
accs = VarInfo(FixedTransformAccumulator())
_, accs = init!!(rng, model, accs, InitFromPrior(), transform_strategy)
return get_fixed_transforms(accs)
Expand Down
21 changes: 14 additions & 7 deletions src/chains.jl
Original file line number Diff line number Diff line change
Expand Up @@ -49,10 +49,10 @@ function ParamsWithStats(
DynamicPPL.LogLikelihoodAccumulator(),
DynamicPPL.RawValueAccumulator(include_colon_eq),
)
_pws_eval(model, accs, init_strategy, stats, true)
_pws_eval(Random.default_rng(), model, accs, init_strategy, stats, true)
else
accs = (DynamicPPL.RawValueAccumulator(include_colon_eq),)
_pws_eval(model, accs, init_strategy, stats, false)
_pws_eval(Random.default_rng(), model, accs, init_strategy, stats, false)
end
end

Expand Down Expand Up @@ -103,7 +103,7 @@ end
)

Generate a `ParamsWithStats` by re-evaluating the given `ldf` with the provided
`param_vector`.
`param_vector` and `ldf.rng`.

This method obtains parameter values and statistics in one model evaluation, without
constructing an intermediate value trace.
Expand Down Expand Up @@ -142,19 +142,26 @@ function pws_with_eval(
DynamicPPL.LogLikelihoodAccumulator(),
DynamicPPL.RawValueAccumulator(include_colon_eq),
)
_pws_eval(ldf.model, accs, strategy, stats, true)
_pws_eval(ldf.rng, ldf.model, accs, strategy, stats, true)
else
accs = (DynamicPPL.RawValueAccumulator(include_colon_eq),)
_pws_eval(ldf.model, accs, strategy, stats, false)
_pws_eval(ldf.rng, ldf.model, accs, strategy, stats, false)
end
end
@noinline function _pws_eval(
model::Model, accs::Tuple, strategy, stats::NamedTuple, include_log_probs::Bool
rng::Random.AbstractRNG,
model::Model,
accs::Tuple,
strategy,
stats::NamedTuple,
include_log_probs::Bool,
)
# UnlinkAll() actually doesn't have any impact here, because there isn't even a
# LogJacobianAccumulator; consequently, it doesn't matter whether we interpret the
# parameters as being in linked space or not. However, we just include it for clarity.
_, vi = DynamicPPL.init!!(model, VarInfo(AccumulatorTuple(accs)), strategy, UnlinkAll())
_, vi = DynamicPPL.init!!(
rng, model, VarInfo(AccumulatorTuple(accs)), strategy, UnlinkAll()
)
params = densify!!(get_raw_values(vi))
if include_log_probs
stats = merge(
Expand Down
2 changes: 2 additions & 0 deletions src/compiler.jl
Original file line number Diff line number Diff line change
@@ -1,3 +1,5 @@
# `__context__` is internal although the docs show `rand(__context__.rng, ...)`; a public
# accessor for the evaluation RNG is still to be added.
const INTERNALNAMES = (:__model__, :__context__, :__varinfo__)

drop_escape(x) = x
Expand Down
Loading
Loading