diff --git a/HISTORY.md b/HISTORY.md index ebfd95e9c..3fd28d13d 100644 --- a/HISTORY.md +++ b/HISTORY.md @@ -24,6 +24,8 @@ Indexed prefixes accept a prefix template: `prefix(m, @varname(a[2]))` → `pref `check_model` accepts explicit argument bindings and warns about binding names absent from the model and reached unprefixed submodels. See [#1501](https://github.com/TuringLang/DynamicPPL.jl/pull/1501). +Density evaluation accepts an explicit RNG: `DynamicPPL.logdensity_internal(args...)` → `DynamicPPL.logdensity_internal(args..., rng)`; append `rng` to `AbstractPPL.prepare`’s `context` tuple. See [#1504](https://github.com/TuringLang/DynamicPPL.jl/pull/1504). + ## Breaking changes `marginalize` and the `DynamicPPLMarginalLogDensitiesExt` extension are removed in DynamicPPL 0.43. Users requiring the existing Turing/MLD integration can remain on DynamicPPL 0.42.x with a compatible Turing release—for example, Turing 0.49.0—and MarginalLogDensities 0.4.3–0.4.x. See the [0.42 marginalisation documentation](https://turinglang.org/DynamicPPL.jl/v0.42/api/#Marginalisation). @@ -120,6 +122,10 @@ Custom `AbstractContext`/`AbstractParentContext` subtyping is unsupported, and b `DynamicPPL.evaluate_nowarn!!(m, vi)` is removed → `evaluate!!(m, ctx, vi)`. See [#1503](https://github.com/TuringLang/DynamicPPL.jl/pull/1503). +`LogDensityFunction` shares RNG state across construction, evaluation, AD preparation, and sampling: implicit draws → `LogDensityFunction(m; rng)` and model-body `rand(__context__.rng, ...)`. See [#1504](https://github.com/TuringLang/DynamicPPL.jl/pull/1504), [#721](https://github.com/TuringLang/DynamicPPL.jl/issues/721). + +`rand(ldf)` uses `ldf.rng`: relying on the task-local RNG → `rand(Random.default_rng(), ldf)`. See [#1504](https://github.com/TuringLang/DynamicPPL.jl/pull/1504). + `@vnt` is no longer exported: `@vnt` → `DynamicPPL.@vnt` or explicitly import it. See [#1501](https://github.com/TuringLang/DynamicPPL.jl/pull/1501). `CondFixContext` is removed: context-stored observations/fixed bindings → `condition(model, values)`/`fix(model, values)`. See [#1501](https://github.com/TuringLang/DynamicPPL.jl/pull/1501). diff --git a/docs/src/evaluation.md b/docs/src/evaluation.md index ee257e05d..8627d3bcd 100644 --- a/docs/src/evaluation.md +++ b/docs/src/evaluation.md @@ -120,7 +120,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, latent LHS variables get their values from supplied parameters rather than sampling. +For density evaluation, latent LHS variables get their values from supplied parameters rather than sampling; +see [Randomness in density evaluation](@ref ldf-rng). ## Accumulators diff --git a/docs/src/ldf/overview.md b/docs/src/ldf/overview.md index d74e1a466..e57feca68 100644 --- a/docs/src/ldf/overview.md +++ b/docs/src/ldf/overview.md @@ -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. @@ -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. diff --git a/docs/src/migration.md b/docs/src/migration.md index fbe2cf3c0..32a51dada 100644 --- a/docs/src/migration.md +++ b/docs/src/migration.md @@ -245,6 +245,20 @@ _, vi = init!!(Xoshiro(468), model, vi, InitFromParams(params, nothing), UnlinkA vi ``` +## Random number generators + +`LogDensityFunction` now shares its supplied RNG across construction, evaluation, AD +preparation, and parameter sampling. Replace implicit model-body draws such as `rand()` +with `rand(__context__.rng)` and supply the RNG with `LogDensityFunction(model; rng)`. +Replace `rand(ldf)` with `rand(Random.default_rng(), ldf)` if sampling should use the +task-local RNG; `rand(ldf)` now uses `ldf.rng`. + +For low-level density calls, replace `DynamicPPL.logdensity_internal(args...)` with +`DynamicPPL.logdensity_internal(args..., rng)` to select an RNG explicitly. For +`AbstractPPL.prepare`, replace `context=(model, getlogdensity, ranges, strategy, accs)` +with `context=(model, getlogdensity, ranges, strategy, accs, rng)`. +See [Randomness in density evaluation](@ref ldf-rng) for the evaluation contract. + ## Binding arguments and data Arguments on the LHS are observed by default. Replace placeholder-based latent data, diff --git a/ext/DynamicPPLMooncakeExt.jl b/ext/DynamicPPLMooncakeExt.jl index affb32c6c..2caa64f25 100644 --- a/ext/DynamicPPLMooncakeExt.jl +++ b/ext/DynamicPPLMooncakeExt.jl @@ -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} diff --git a/src/accumulators/fixed_transforms.jl b/src/accumulators/fixed_transforms.jl index eec029586..bb15ef3ea 100644 --- a/src/accumulators/fixed_transforms.jl +++ b/src/accumulators/fixed_transforms.jl @@ -55,16 +55,14 @@ 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 @@ -72,9 +70,10 @@ executing the model multiple times and is thus not foolproof (it may be that you 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) diff --git a/src/chains.jl b/src/chains.jl index 9f3f09e15..0b6911780 100644 --- a/src/chains.jl +++ b/src/chains.jl @@ -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 @@ -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. @@ -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( diff --git a/src/compiler.jl b/src/compiler.jl index c34ddc28c..71e5d7736 100644 --- a/src/compiler.jl +++ b/src/compiler.jl @@ -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 diff --git a/src/logdensityfunction.jl b/src/logdensityfunction.jl index 5113a9ecf..a87a9bd7d 100644 --- a/src/logdensityfunction.jl +++ b/src/logdensityfunction.jl @@ -32,6 +32,7 @@ using Random: Random x::AbstractVector{<:Real}, accs::Union{NTuple{<:Any,AbstractAccumulator},AccumulatorTuple}=ldf_accs(getlogdensity); adtype::Union{ADTypes.AbstractADType,Nothing}=nothing, + rng::Random.AbstractRNG=Random.default_rng(), ) A struct which contains a model, along with all the information necessary to: @@ -105,6 +106,13 @@ along with other information for efficient calculation of the gradient of the lo Note that preparing a `LogDensityFunction` with an AD type `AutoBackend()` requires the AD backend itself to have been loaded (e.g. with `import Backend`). +The `rng` keyword supplies the RNG for construction and evaluation. During density +evaluation, `x ~ Normal()` reads `x` from the parameter vector and draws nothing. +Explicit model-body calls such as `rand(__context__.rng)` advance the supplied RNG; +the RNG is shared, not copied or reset. Construction may sample parameters, and AD +preparation may execute the model. See [Randomness in density evaluation](@ref ldf-rng) +for stochastic-density limitations. + ## Fields Note that it is undefined behaviour to access any of a `LogDensityFunction`'s fields, apart @@ -115,6 +123,7 @@ from: type was provided. - `ldf.transform_strategy`: The transform strategy that specifies the transforms for all variables in the model. +- `ldf.rng`: The supplied RNG, also available inside the model as `__context__.rng`. For all other fields, please use the corresponding getter functions provided in the API: @@ -154,6 +163,7 @@ struct LogDensityFunction{ AC<:AccumulatorTuple, # whether all transforms are FixedTransforms AllFixed, + R<:Random.AbstractRNG, } model::M adtype::AD @@ -164,6 +174,7 @@ struct LogDensityFunction{ _dim::Int _x::X _accs::AC + rng::R function LogDensityFunction( model::Model, @@ -174,6 +185,7 @@ struct LogDensityFunction{ getlogdensity ); adtype::Union{ADTypes.AbstractADType,Nothing}=nothing, + rng::Random.AbstractRNG=Random.default_rng(), ) dim = length(x) # Determine LDF transform strategy. @@ -192,7 +204,12 @@ struct LogDensityFunction{ # Make backend-specific tweaks to the adtype adtype = DynamicPPL.tweak_adtype(adtype, model, x) context = ( - model, getlogdensity, ranges_and_transforms, transform_strategy, accs + model, + getlogdensity, + ranges_and_transforms, + transform_strategy, + accs, + rng, ) # `x` was just constructed from the same range metadata stored in `context`, # so the AD wrapper can skip its hot-path dimension validation. @@ -210,6 +227,7 @@ struct LogDensityFunction{ typeof(x), typeof(accs), all_fixed, + typeof(rng), }( model, adtype, @@ -220,6 +238,7 @@ struct LogDensityFunction{ dim, x, accs, + rng, ) end end @@ -232,6 +251,7 @@ end accs::Union{NTuple{<:Any,AbstractAccumulator},AccumulatorTuple}=ldf_accs(getlogdensity); adtype::Union{ADTypes.AbstractADType,Nothing}=nothing, fix_transforms::Bool=false, + rng::Random.AbstractRNG=Random.default_rng(), ) Most users of LogDensityFunction should use this constructor, which does **not** require @@ -272,6 +292,9 @@ You can pass either: The `adtype` keyword argument allows you to specify an AD type for gradient preparation and calculation. +The `rng` keyword is forwarded through construction, AD preparation, and evaluation; +see [Randomness in density evaluation](@ref ldf-rng). + The `fix_transforms` keyword argument allows you to specify whether the transforms used in the `LogDensityFunction` should be cached at the time of construction. If so, the model is evaluated once using the provided transform strategy, and the transforms used for each @@ -287,6 +310,7 @@ function LogDensityFunction( accs::Union{NTuple{<:Any,AbstractAccumulator},AccumulatorTuple}=ldf_accs(getlogdensity); adtype::Union{ADTypes.AbstractADType,Nothing}=nothing, fix_transforms::Bool=false, + rng::Random.AbstractRNG=Random.default_rng(), ) # Handle fixed transforms flag. if fix_transforms @@ -298,13 +322,13 @@ function LogDensityFunction( # tolerable since this isn't something that is in a performance-sensitive code # path. dynamic_transform_strategy = infer_transform_strategy_from_values(vecvals) - transforms_vnt = get_fixed_transforms(model, dynamic_transform_strategy) + transforms_vnt = get_fixed_transforms(model, dynamic_transform_strategy; rng) vecvals = update_transforms!!(vecvals, transforms_vnt) end end ranges_and_transforms, x = get_rat_and_samplevec(vecvals) return LogDensityFunction( - model, getlogdensity, ranges_and_transforms, x, accs; adtype=adtype + model, getlogdensity, ranges_and_transforms, x, accs; adtype, rng ) end function LogDensityFunction( @@ -314,6 +338,7 @@ function LogDensityFunction( accs::Union{NTuple{<:Any,AbstractAccumulator},AccumulatorTuple}=ldf_accs(getlogdensity); adtype::Union{ADTypes.AbstractADType,Nothing}=nothing, fix_transforms::Bool=false, + rng::Random.AbstractRNG=Random.default_rng(), ) if !hasacc(vi, Val(VECTORVAL_ACCNAME)) error( @@ -321,9 +346,7 @@ function LogDensityFunction( ) end vnt = getacc(vi, Val(VECTORVAL_ACCNAME)).values - return LogDensityFunction( - model, getlogdensity, vnt, accs; adtype=adtype, fix_transforms=fix_transforms - ) + return LogDensityFunction(model, getlogdensity, vnt, accs; adtype, fix_transforms, rng) end function LogDensityFunction( model::Model, @@ -332,13 +355,14 @@ function LogDensityFunction( accs::Union{NTuple{<:Any,AbstractAccumulator},AccumulatorTuple}=ldf_accs(getlogdensity); adtype::Union{ADTypes.AbstractADType,Nothing}=nothing, fix_transforms::Bool=false, + rng::Random.AbstractRNG=Random.default_rng(), ) # note that this reevaluates the model vi = VarInfo(VectorValueAccumulator()) - _, vi = DynamicPPL.init!!(model, vi, InitFromPrior(), transform_strategy) + _, vi = DynamicPPL.init!!(rng, model, vi, InitFromPrior(), transform_strategy) vecvals = getacc(vi, Val(VECTORVAL_ACCNAME)).values return LogDensityFunction( - model, getlogdensity, vecvals, accs; adtype=adtype, fix_transforms=fix_transforms + model, getlogdensity, vecvals, accs; adtype, fix_transforms, rng ) end @@ -418,11 +442,13 @@ ldf_accs(::typeof(getloglikelihood)) = AccumulatorTuple((LogLikelihoodAccumulato varname_ranges::VarNamedTuple, transform_strategy::AbstractTransformStrategy, accs::AccumulatorTuple, + rng::Random.AbstractRNG=Random.default_rng(), ) Calculate the log density at the given `params`, using the provided information extracted from a `LogDensityFunction`. This is the internal implementation behind -`LogDensityProblems.logdensity(ldf, params)`. +`LogDensityProblems.logdensity(ldf, params)`. `rng` drives explicit draws in the model +body, such as `rand(__context__.rng)`; it defaults to the task-local RNG. """ function logdensity_internal( params::AbstractVector{<:Real}, @@ -431,9 +457,10 @@ function logdensity_internal( varname_ranges::VarNamedTuple, transform_strategy::AbstractTransformStrategy, accs::AccumulatorTuple, + rng::Random.AbstractRNG=Random.default_rng(), ) init_strategy = InitFromVector(params, varname_ranges, transform_strategy) - _, vi = DynamicPPL.init!!(model, VarInfo(accs), init_strategy, transform_strategy) + _, vi = DynamicPPL.init!!(rng, model, VarInfo(accs), init_strategy, transform_strategy) return getlogdensity(vi) end @@ -460,7 +487,8 @@ function LogDensityAt( getlogdensity, varname_ranges::VarNamedTuple, transform_strategy::AbstractTransformStrategy, - accs::AccumulatorTuple, + accs::AccumulatorTuple; + rng::Random.AbstractRNG=Random.default_rng(), ) Base.depwarn( "`DynamicPPL.LogDensityAt` is deprecated; call " * @@ -469,7 +497,7 @@ function LogDensityAt( :LogDensityAt, ) dim = mapreduce(rat -> length(rat.range), +, values(varname_ranges); init=0) - context = (model, getlogdensity, varname_ranges, transform_strategy, accs) + context = (model, getlogdensity, varname_ranges, transform_strategy, accs, rng) return AbstractPPL.prepare( logdensity_internal, zeros(dim); check_dims=false, context=context ) @@ -492,6 +520,7 @@ end ldf._varname_ranges, ldf.transform_strategy, ldf._accs, + ldf.rng, ) end @@ -716,6 +745,7 @@ end Generate a random vector of parameters that is consistent with the given `LogDensityFunction`, using the provided initialisation strategy. +If `rng` is omitted, use `ldf.rng`. Note that this function only generates parameters, and does not return the log density. If you also need the log density, instead of calling `rand` and then @@ -737,5 +767,5 @@ end function Base.rand( ldf::LogDensityFunction, init_strategy::AbstractInitStrategy=InitFromPrior() ) - return rand(Random.default_rng(), ldf, init_strategy) + return rand(ldf.rng, ldf, init_strategy) end diff --git a/src/subsample.jl b/src/subsample.jl index 5803e7c95..a512c3dae 100644 --- a/src/subsample.jl +++ b/src/subsample.jl @@ -609,7 +609,7 @@ function _probe_independent_model( vi = VarInfo(accumulators) _, vi = init!!(rng, model, vi, init_strategy, transform_strategy) IndependentLogJoint()(vi) - return LogDensityFunction(model, getlogjoint_internal, get_vector_values(vi)) + return LogDensityFunction(model, getlogjoint_internal, get_vector_values(vi); rng) end function _strict_independent_ldf( @@ -620,7 +620,8 @@ function _strict_independent_ldf( scale::Real, expected_nobs::Int, population_size::Int, - indices=nothing, + indices=nothing; + rng::Random.AbstractRNG, ) accumulators = AccumulatorTuple(( VectorValueAccumulator(), @@ -629,7 +630,7 @@ function _strict_independent_ldf( )..., )) return LogDensityFunction( - model, IndependentLogJoint(ranges), ranges, sample, accumulators + model, IndependentLogJoint(ranges), ranges, sample, accumulators; rng ) end @@ -874,7 +875,7 @@ function subsample( indices = _sample_without_replacement( rng, size(problem.full_data, ndims(problem.full_data)), Int(batch_size) ) - return _batch_ldf(problem, indices) + return _batch_ldf(problem, indices; rng) end function subsample( @@ -885,7 +886,7 @@ function subsample( transform_strategy::T=UnlinkAll(), ) where {T<:AbstractTransformStrategy} problem = _subsampling_problem(model, dataset_size, transform_strategy) - return _batch_ldf(problem, indices) + return _batch_ldf(problem, indices; rng) end function subsample( @@ -903,7 +904,7 @@ function subsample( "the resampler must return a vector of integer indices; got $(typeof(indices))", ), ) - return _batch_ldf(problem, indices) + return _batch_ldf(problem, indices; rng) end function subsample( @@ -950,7 +951,11 @@ function _sample_without_replacement( return indices end -function _batch_ldf(problem::SubsamplingState, batch::AbstractVector{<:Integer}) +function _batch_ldf( + problem::SubsamplingState, + batch::AbstractVector{<:Integer}; + rng::Random.AbstractRNG=problem.ldf.rng, +) batch_data = _select_batch(problem.full_data, batch) batch = collect(Int, batch) batch_model = condition( @@ -984,6 +989,7 @@ function _batch_ldf(problem::SubsamplingState, batch::AbstractVector{<:Integer}) scale, batch_size, full_size, - batch, + batch; + rng, ) end diff --git a/src/test_utils/ad.jl b/src/test_utils/ad.jl index 7b879f3ac..e9e164a6a 100644 --- a/src/test_utils/ad.jl +++ b/src/test_utils/ad.jl @@ -333,7 +333,7 @@ function run_ad( verbose && @info "Running AD on $(model.f) with $(adtype)\n" # Generate initial parameters - ldf = LogDensityFunction(model, getlogdensity, transform_strategy; adtype=adtype) + ldf = LogDensityFunction(model, getlogdensity, transform_strategy; adtype, rng) if isnothing(params) params = rand(rng, ldf, InitFromPrior()) end @@ -358,7 +358,7 @@ function run_ad( grad_true = test.grad elseif test isa WithBackend ldf_reference = LogDensityFunction( - model, getlogdensity, transform_strategy; adtype=test.adtype + model, getlogdensity, transform_strategy; adtype=test.adtype, rng ) value_true, grad_true = logdensity_and_gradient(ldf_reference, params) grad_true = collect(grad_true) diff --git a/test/logdensityfunction.jl b/test/logdensityfunction.jl index e3d5bd0c1..0021444e3 100644 --- a/test/logdensityfunction.jl +++ b/test/logdensityfunction.jl @@ -10,13 +10,79 @@ using DynamicPPL.TestUtils.AD: run_ad, WithExpectedResult, NoTest using LinearAlgebra: I using Test using LogDensityProblems: LogDensityProblems -using Random: Xoshiro +using Random: Random, Xoshiro using StableRNGs: StableRNG using DifferentiationInterface: DifferentiationInterface using ForwardDiff: ForwardDiff using Mooncake: Mooncake +@model function rng_child() + x ~ Normal() + u = rand(__context__.rng) + 0.0 ~ Normal(x + u, 1) + return x +end +@model rng_parent() = child ~ to_submodel(rng_child()) +@model rng_deterministic() = x ~ Normal() + +@testset "LogDensityFunction: explicit RNG" begin + for adtype in (AutoForwardDiff(), AutoMooncake()), rng_type in (Xoshiro, StableRNG) + rng = rng_type(42) + ldf = LogDensityFunction(rng_deterministic(); rng, adtype) + expected_rng = copy(rng) + @test LogDensityProblems.logdensity(ldf, [0.2]) ≈ logpdf(Normal(), 0.2) + @test last(LogDensityProblems.logdensity_and_gradient(ldf, [0.2])) ≈ [-0.2] + @test rand(copy(rng)) == rand(expected_rng) + expected_rng = copy(rng) + @test only(rand(ldf)) == rand(expected_rng, Normal()) + @test rand(copy(rng)) == rand(expected_rng) + end + + for model in (rng_child(), rng_parent()) + values = get_vector_values(VarInfo(Xoshiro(1), model)) + ranges, x = DynamicPPL.get_rat_and_samplevec(values) + for (args, kwargs) in ( + ((UnlinkAll(),), (;)), + ((VarInfo(Xoshiro(1), model),), (;)), + ((values,), (;)), + ((values,), (; fix_transforms=true)), + ((ranges, x), (;)), + ) + rng = Xoshiro(42) + default_rng = copy(Random.default_rng()) + ldf = LogDensityFunction(model, getlogjoint_internal, args...; rng, kwargs...) + @test rand() == rand(default_rng) + expected_rng = copy(rng) + for default_seed in (1, 2) + Random.seed!(default_seed) + u = rand(expected_rng) + @test (@inferred LogDensityProblems.logdensity(ldf, [0.2])) ≈ + logpdf(Normal(), 0.2) + logpdf(Normal(0.2 + u, 1), 0.0) + end + for include_log_probs in (true, false) + u = rand(expected_rng) + output = ParamsWithStats([0.2], ldf; include_log_probs) + if include_log_probs + @test output.stats.logjoint ≈ + logpdf(Normal(), 0.2) + logpdf(Normal(0.2 + u, 1), 0.0) + end + end + @test rand(copy(rng)) == rand(expected_rng) + end + for adtype in (AutoForwardDiff(), AutoMooncake()) + rng = Xoshiro(42) + ldf = LogDensityFunction(model; rng, adtype) + expected_rng = copy(rng) + u = rand(expected_rng) + logp, grad = LogDensityProblems.logdensity_and_gradient(ldf, [0.2]) + @test logp ≈ logpdf(Normal(), 0.2) + logpdf(Normal(0.2 + u, 1), 0.0) + @test only(grad) ≈ -0.4 - u + @test rand(copy(rng)) == rand(expected_rng) + end + end +end + @model function issue_2844_nested_inner() p ~ Normal() return (; p) @@ -655,6 +721,25 @@ end end end +@testset "logdensity_internal defaults to the task-local RNG" begin + @model tiny() = x ~ Normal() + ldf = LogDensityFunction(tiny()) + args = ( + ldf.model, + DynamicPPL.getlogjoint_internal, + DynamicPPL.get_all_ranges_and_transforms(ldf), + ldf.transform_strategy, + ldf._accs, + ) + params = [0.3] + @test DynamicPPL.logdensity_internal(params, args...) ≈ + LogDensityProblems.logdensity(ldf, params) + prepared = AbstractPPL.prepare( + DynamicPPL.logdensity_internal, params; check_dims=false, context=args + ) + @test prepared(params) ≈ LogDensityProblems.logdensity(ldf, params) +end + @testset "LogDensityAt deprecation shim" begin @model tiny() = x ~ Normal() ldf = LogDensityFunction(tiny())