From 230a711f68b5c7b86c36e385c5df5c556d2e5ca7 Mon Sep 17 00:00:00 2001 From: Hong Ge Date: Wed, 16 Sep 2026 20:28:54 +0100 Subject: [PATCH 1/5] evaluation: pass context explicitly to model execution Assisted-by: Codex --- HISTORY.md | 4 +- Project.toml | 2 +- benchmarks/benchmarks.jl | 15 +- docs/src/accs/threadsafe.md | 4 +- docs/src/accs/values.md | 6 +- docs/src/api.md | 55 ++---- docs/src/evaluation.md | 22 +-- docs/src/migration.md | 6 +- docs/src/onboarding.md | 6 +- docs/src/tilde.md | 4 +- ext/DynamicPPLBridgeStanExt.jl | 13 +- ext/DynamicPPLInputProvenanceExt.jl | 3 +- src/DynamicPPL.jl | 17 +- src/accumulators.jl | 3 +- src/compiler.jl | 16 +- src/contexts.jl | 114 +----------- src/contexts/default.jl | 54 ------ src/contexts/init.jl | 64 +++---- src/debug_utils.jl | 18 +- src/logdensityfunction.jl | 4 +- src/model.jl | 256 +++++++++------------------ src/submodel.jl | 9 +- src/varinfo.jl | 8 +- test/compiler.jl | 3 +- test/conditionfix.jl | 2 +- test/context_implementations.jl | 41 ++--- test/contexts/init.jl | 2 +- test/debug_utils.jl | 5 +- test/ext/DynamicPPLBridgeStanExt.jl | 21 +++ test/ext/DynamicPPLForwardDiffExt.jl | 34 ++++ test/model.jl | 11 +- test/runtests.jl | 4 +- test/submodels.jl | 4 +- test/threadsafe.jl | 4 +- test/transformed_values.jl | 4 +- test/utils.jl | 2 +- test/varinfo.jl | 10 +- 37 files changed, 293 insertions(+), 557 deletions(-) delete mode 100644 src/contexts/default.jl diff --git a/HISTORY.md b/HISTORY.md index 57fba7676..4b02b1345 100644 --- a/HISTORY.md +++ b/HISTORY.md @@ -16,7 +16,7 @@ Re-evaluation and `LogDensityFunction` construction no longer copy fixed transfo `pointwise_loglikelihoods` and `pointwise_logdensities` now record observations for threadsafe models, such as `setthreadsafe(model, true)`; previously they were silently omitted. See [#1500](https://github.com/TuringLang/DynamicPPL.jl/pull/1500). -Added `evaluate!!(model, context, vi)` to evaluate with an explicit leaf context and collect outputs in `vi`, such as `evaluate!!(model, InitContext(rng, InitFromPrior(), UnlinkAll()), VarInfo())`. See [#1500](https://github.com/TuringLang/DynamicPPL.jl/pull/1500). +Added `evaluate!!(model, context, vi)` to evaluate with an explicit `Context` and collect outputs in `vi`, such as `evaluate!!(model, Context(rng, InitFromPrior(), UnlinkAll()), VarInfo())`. See [#1500](https://github.com/TuringLang/DynamicPPL.jl/pull/1500). Added the `template` keyword to `prefix(model, x::VarName)` to supply the enclosing container's shape and resolve `begin` and `end` in indexed prefixes. See [#1502](https://github.com/TuringLang/DynamicPPL.jl/pull/1502). @@ -30,6 +30,8 @@ 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). +`Context(rng, init_strategy, transform_strategy)` replaces `InitContext` and the context hierarchy. Pass it to `evaluate!!(model, context, outputs)`; custom initialisation and observation handling belong to strategies and accumulators. + `PrefixContext` and `extract_prefixes` are removed: use `prefix(model, vn; template)` to set prefixes and `DynamicPPL.getprefix(model)` to read the combined prefix (`nothing` when absent). The `prefix` field stores internal metadata for LHS variable addresses and nested submodel namespace storage templates. `condition` and `fix` reject binding addresses outside the model's prefix; use the prefixed address, such as `@varname(p.y)`, instead of `y`. See [#1502](https://github.com/TuringLang/DynamicPPL.jl/pull/1502). Partly removing bindings of a single multivariate LHS variable now throws `ArgumentError` during evaluation; declare separate LHS variables to remove their bindings independently. See [#1501](https://github.com/TuringLang/DynamicPPL.jl/pull/1501). diff --git a/Project.toml b/Project.toml index 07d07209d..2dc670d2d 100644 --- a/Project.toml +++ b/Project.toml @@ -62,7 +62,7 @@ BangBang = "0.4.1" Bijectors = "0.16" BridgeStan = "2" Chairmarks = "1.3.1" -Compat = "4" +Compat = "4.10" ComponentArrays = "0.15" ConstructionBase = "1.5.4" Distributions = "0.25" diff --git a/benchmarks/benchmarks.jl b/benchmarks/benchmarks.jl index f59f6caf2..c6d063bcc 100644 --- a/benchmarks/benchmarks.jl +++ b/benchmarks/benchmarks.jl @@ -116,6 +116,18 @@ end return (; x=x) end +@model _indexed_observation(obs, mu) = obs ~ Normal(mu, 1) + +"Indexed submodels with argument observations and one shared latent mean." +@model function indexed_submodels(obs) + mu ~ Normal() + x = similar(obs) + for i in eachindex(obs) + x[i] ~ to_submodel(_indexed_observation(obs[i], mu)) + end + return (; mu=mu) +end + "Variables whose support varies under linking, or otherwise nontrivial bijectors." @model function dynamic() eta ~ truncated(Normal(); lower=0.0, upper=0.1) @@ -167,7 +179,7 @@ function model_dimension(model, islinked) DynamicPPL.init!!( StableRNG(23), model, - VarInfo(), + VarInfo(DynamicPPL.VectorValueAccumulator()), DynamicPPL.InitFromPrior(), transform_strategy(islinked), ), @@ -360,6 +372,7 @@ function build_combinations(rng) end push!(models, ("Dynamic", dynamic())) push!(models, ("Submodel", parent(randn(rng)))) + push!(models, ("Indexed submodels 3k", indexed_submodels(randn(rng, 3_000)))) d = [1, 1, 1, 2, 2, 2] w = [1, 2, 3, 2, 1, 1] z = [1, 1, 2, 2, 1, 2] diff --git a/docs/src/accs/threadsafe.md b/docs/src/accs/threadsafe.md index 3c60b25f2..ed3b53242 100644 --- a/docs/src/accs/threadsafe.md +++ b/docs/src/accs/threadsafe.md @@ -48,8 +48,8 @@ tilde-statement. ```@example 1 x = 1.0 -context = DynamicPPL.InitContext(InitFromParams((; x=x)), UnlinkAll()) -_, tsvi = DynamicPPL._evaluate!!(contextualize(model, context), tsvi) +context = DynamicPPL.Context(InitFromParams((; x=x)), UnlinkAll()) +_, tsvi = DynamicPPL._evaluate!!(model, context, tsvi) length(tsvi.accs_by_task) ``` diff --git a/docs/src/accs/values.md b/docs/src/accs/values.md index 909fa5581..726c6b004 100644 --- a/docs/src/accs/values.md +++ b/docs/src/accs/values.md @@ -28,7 +28,7 @@ using Random: Xoshiro return x[1:3] ~ Dirichlet(ones(3)) end model = dirichlet() -context = InitContext(Xoshiro(1), InitFromPrior(), LinkAll()) +context = Context(Xoshiro(1), InitFromPrior(), LinkAll()) _, vi = evaluate!!(model, context, VarInfo(VectorValueAccumulator())) vector_values = get_vector_values(vi) keys(vector_values) @@ -51,7 +51,7 @@ A `RawValueAccumulator` records untransformed values. It does not retain LHS var block boundaries: indexed LHS variables are represented by their individual indices. ```@example 1 -context = InitContext(Xoshiro(1), InitFromPrior(), UnlinkAll()) +context = Context(Xoshiro(1), InitFromPrior(), UnlinkAll()) _, vi = evaluate!!(model, context, VarInfo(RawValueAccumulator(false))) raw_values = get_raw_values(vi) keys(raw_values) @@ -67,7 +67,7 @@ convert them explicitly outside evaluation; the context holds inputs and the `Va holds only outputs: ```@example 1 -context = InitContext(Xoshiro(1), InitFromParams(raw_values, nothing), LinkAll()) +context = Context(Xoshiro(1), InitFromParams(raw_values, nothing), LinkAll()) retval, outputs = evaluate!!(model, context, VarInfo(VectorValueAccumulator())) get_vector_values(outputs) ``` diff --git a/docs/src/api.md b/docs/src/api.md index 9bf2acd34..0f994a620 100644 --- a/docs/src/api.md +++ b/docs/src/api.md @@ -28,12 +28,6 @@ Model Model() ``` -The context of a model can be set using [`contextualize`](@ref): - -```@docs -contextualize -``` - Some models require threadsafe evaluation (see [the Turing docs](https://turinglang.org/docs/usage/threadsafe-evaluation/) for more information on when this is necessary). If this is the case, one must enable threadsafe evaluation for a model: @@ -74,7 +68,7 @@ get_sample_input_vector subsample ``` -Internally, this is accomplished using [`init!!`](@ref) on: +Internally, this is accomplished using [`init!!`](@ref) with [`VarInfo`](@ref). ```@docs to_vector_params @@ -471,7 +465,7 @@ unflatten!! internal_values_as_vector ``` -### Evaluation Contexts +### Evaluation contexts Internally, model evaluation is performed with [`AbstractPPL.evaluate!!`](@ref). @@ -479,59 +473,38 @@ Internally, model evaluation is performed with [`AbstractPPL.evaluate!!`](@ref). AbstractPPL.evaluate!! ``` -This method mutates the `varinfo` used for execution. -By default, it does not perform any actual sampling: it only evaluates the model using the values of the variables that are already in the `varinfo`. -If you wish to sample new values, see the section on [VarInfo initialisation](#VarInfo-initialisation) just below this. - -The behaviour of a model execution can be changed with evaluation contexts, which are a field of the model. - -All contexts are subtypes of `AbstractPPL.AbstractContext`. +Call `evaluate!!(model, context, varinfo)` to evaluate with an explicit context and collect +outputs in `varinfo`. Accumulators are reset before evaluation. -Contexts are split into two kinds: +The context is an evaluation input; it is not stored in the model. +Prefixes are stored separately from values. Conditioned and fixed values share one store, with each value carrying its role. Only latent sites reach the context; observations and tracked values go directly to accumulators. -**Leaf contexts**: These are the most important contexts as they ultimately decide how model evaluation proceeds. -For example, `DefaultContext` reuses values recorded by a `VectorValueAccumulator`, whereas `InitContext` obtains values either by sampling or from supplied parameters. -DynamicPPL has more leaf contexts which are used for internal purposes, but these are the two that are exported. +`Context` is the sole evaluation context. It supplies an RNG, an initialisation strategy, +and a transform strategy. The output `varinfo` never supplies latent inputs. ```@docs -DefaultContext -InitContext +DynamicPPL.Context ``` -Customise latent value selection with an [initialisation strategy](init.md) supplied to `InitContext`. +Customise latent value selection with an [initialisation strategy](init.md) supplied to `Context`. `tilde_assume!!` dispatches on that context to initialise, transform, and accumulate a latent value. Every observation calls `tilde_observe!!(prefix, prefix_template, right, left, vn, template, vi)`, which applies the prefix metadata and calls `accumulate_observe!!` without dispatching on a context. ```@docs tilde_assume!! tilde_observe!! +DynamicPPL.store_coloneq_value!! ``` -**Parent contexts**: These essentially act as 'modifiers' for leaf contexts. -Prefixes, conditioned values, and fixed values are stored on the model. - -To implement a parent context, you have to subtype `DynamicPPL.AbstractParentContext`, and implement the `childcontext` and `setchildcontext` methods. -If needed, you can also implement `tilde_assume!!` for your context. -This is optional; the default implementation is to simply delegate to the child context. - -```@docs -AbstractParentContext -childcontext -setchildcontext -``` - -Since contexts form a tree structure, these functions are automatically defined for manipulating context stacks. -They are mainly useful for modifying the fundamental behaviour (i.e. the leaf context), without affecting any of the modifiers (i.e. parent contexts). +Downstream evaluators that control execution directly can prepare arguments for `model.f`: ```@docs -leafcontext -setleafcontext +DynamicPPL.make_evaluate_args_and_kwargs ``` ### VarInfo initialisation -The function `init!!` is used to initialise, or overwrite, values in a VarInfo. -It is really a thin wrapper around using `evaluate!!` with an `InitContext`. +The function `init!!` constructs a `Context` and evaluates the model, resetting the output accumulators. ```@docs init!! diff --git a/docs/src/evaluation.md b/docs/src/evaluation.md index 384ef9440..1ac16ae00 100644 --- a/docs/src/evaluation.md +++ b/docs/src/evaluation.md @@ -54,7 +54,7 @@ The equivalent explicit-context call is: ```@example 1 using Random: Xoshiro -context = InitContext(Xoshiro(1), InitFromPrior(), UnlinkAll()) +context = Context(Xoshiro(1), InitFromPrior(), UnlinkAll()) retval, accs = evaluate!!(model, context, VarInfo()); ``` @@ -63,12 +63,12 @@ retval, accs = evaluate!!(model, context, VarInfo()); Evaluation separates the inputs that determine a model run from the outputs it records, as proposed in [#1469](https://github.com/TuringLang/DynamicPPL.jl/issues/1469). -| Object | Responsibility | -|:------------- |:--------------------------------------------------------------------- | -| `Model` | Model function, arguments, and conditioned or fixed data | -| `InitContext` | RNG, initialisation strategy, and requested transform strategy | -| `VarInfo` | Output accumulators, with no separate parameter or transform storage | -| `retval` | The model body's ordinary Julia return value, distinct from its trace | +| Object | Responsibility | +|:--------- |:--------------------------------------------------------------------- | +| `Model` | Model function, arguments, and conditioned or fixed data | +| `Context` | RNG, initialisation strategy, and requested transform strategy | +| `VarInfo` | Output accumulators, with no separate parameter or transform storage | +| `retval` | The model body's ordinary Julia return value, distinct from its trace | For a latent statement such as `x ~ Normal()`, the context's initialisation strategy supplies `x`. Its transform strategy determines the transformed value and Jacobian. @@ -81,9 +81,9 @@ use the same observation path. Fixed values are not scored, and tracked assignme such as `z := x + y` are recorded when requested. None of these operations uses the context to select a latent value. -The supplied leaf context replaces the model’s leaf context and is inherited by nested submodels. +The context belongs to the evaluation, not to `Model`, and is passed to nested submodels. Inside a model body, `__context__` refers to this context; use `rand(__context__.rng, ...)` -for explicit random draws controlled by the evaluation's RNG. `init!!` constructs a `InitContext` +for explicit random draws controlled by the evaluation's RNG. `init!!` constructs a `Context` and calls `evaluate!!`; custom value selection belongs in an initialisation strategy, not a custom context type. @@ -98,11 +98,11 @@ explicitly. For example, sample the model above, then evaluate it at the same pa ```@example 1 rng = Xoshiro(1) -context = InitContext(rng, InitFromPrior(), LinkAll()) +context = Context(rng, InitFromPrior(), LinkAll()) retval, recorded = evaluate!!(model, context, VarInfo(RawValueAccumulator(false))) params = get_raw_values(recorded) -context = InitContext(rng, InitFromParams(params, nothing), UnlinkAll()) +context = Context(rng, InitFromParams(params, nothing), UnlinkAll()) repeated, scores = evaluate!!(model, context, VarInfo()) @assert repeated == retval diff --git a/docs/src/migration.md b/docs/src/migration.md index e18a88d34..41117a6cb 100644 --- a/docs/src/migration.md +++ b/docs/src/migration.md @@ -14,10 +14,12 @@ to set a prefix and `DynamicPPL.getprefix(model)` to read the combined prefix (`nothing` when absent). The `prefix` field stores internal metadata for LHS variable addresses and nested submodel namespace storage templates. -To reuse previous values, extract them explicitly before evaluating: +Replace `InitContext` with `Context`. `DefaultContext` and context subtyping are removed: +custom value selection belongs in initialisation strategies. To reuse previous values, +extract them explicitly before evaluating: ```julia -context = InitContext(rng, InitFromParams(get_vector_values(previous), nothing), LinkAll()) +context = Context(rng, InitFromParams(get_vector_values(previous), nothing), LinkAll()) retval, outputs = evaluate!!(model, context, VarInfo()) ``` diff --git a/docs/src/onboarding.md b/docs/src/onboarding.md index fd83db501..4e4503aa3 100644 --- a/docs/src/onboarding.md +++ b/docs/src/onboarding.md @@ -42,13 +42,11 @@ Start with these docs: ### Prefer explicit evaluation state -Keep evaluation inputs in `InitContext` and choose output accumulators in `VarInfo`. +Keep evaluation inputs in `Context` and choose output accumulators in `VarInfo`. For example, to reuse recorded parameters while collecting a different set of outputs: ```julia -context = InitContext( - rng, InitFromParams(get_vector_values(previous), nothing), UnlinkAll() -) +context = Context(rng, InitFromParams(get_vector_values(previous), nothing), UnlinkAll()) retval, outputs = evaluate!!(model, context, VarInfo(accumulators...)) ``` diff --git a/docs/src/tilde.md b/docs/src/tilde.md index df237538e..847efab54 100644 --- a/docs/src/tilde.md +++ b/docs/src/tilde.md @@ -53,11 +53,11 @@ As described on the [Model evaluation page](./evaluation.md), there are three st 2. Transformation: figure out the untransformed (raw) value and the transformed value (where necessary); compute the relevant log-Jacobian. 3. Accumulation: pass all the relevant information to the accumulators, which individually decide what to do with it. -The method for `tilde_assume!!` (with `InitContext`) more or less implements this logic directly with three lines of code. +The method for `tilde_assume!!` (with `Context`) more or less implements this logic directly with three lines of code. The implementation in `src/contexts/init.jl` follows this structure: ```julia -function DynamicPPL.tilde_assume!!(ctx::InitContext, dist, vn, template, vi) +function DynamicPPL.tilde_assume!!(ctx::Context, dist, vn, template, vi) # 1. Initialisation init_tval = DynamicPPL.init(ctx.rng, vn, dist, ctx.strategy) diff --git a/ext/DynamicPPLBridgeStanExt.jl b/ext/DynamicPPLBridgeStanExt.jl index 27eca3d4a..b4e72e39c 100644 --- a/ext/DynamicPPLBridgeStanExt.jl +++ b/ext/DynamicPPLBridgeStanExt.jl @@ -524,7 +524,7 @@ function _stan_assume!!(distribution::StanDistribution, vn, template, vi, transf end function DynamicPPL.tilde_assume!!( - context::DynamicPPL.InitContext, + context::DynamicPPL.Context, distribution::StanDistribution, vn::DynamicPPL.VarName, template, @@ -535,17 +535,6 @@ function DynamicPPL.tilde_assume!!( return _stan_assume!!(distribution, vn, template, vi, transformed_value) end -function DynamicPPL.tilde_assume!!( - ::DynamicPPL.DefaultContext, - distribution::StanDistribution, - vn::DynamicPPL.VarName, - template, - vi::DynamicPPL.AbstractVarInfo, -) - transformed_value = DynamicPPL.get_transformed_value(vi, vn) - return _stan_assume!!(distribution, vn, template, vi, transformed_value) -end - function _stan_constrain(transform::StanTransform, u::AbstractVector{<:Real}) _check_length(u, transform.input_dimension, "u") return BridgeStan.param_constrain(transform.model, collect(Float64, u)) diff --git a/ext/DynamicPPLInputProvenanceExt.jl b/ext/DynamicPPLInputProvenanceExt.jl index 78f921c3f..72c088d1e 100644 --- a/ext/DynamicPPLInputProvenanceExt.jl +++ b/ext/DynamicPPLInputProvenanceExt.jl @@ -191,8 +191,7 @@ function check_input_provenance(rng, model, params) args, defaults, model.prefix, - values, - model.context; + values; args_on_lhs=DynamicPPL._args_on_lhs(model), ) vi = DynamicPPL.VarInfo((InputProvenanceAccumulator(),)) diff --git a/src/DynamicPPL.jl b/src/DynamicPPL.jl index 99e242c49..d50267fca 100644 --- a/src/DynamicPPL.jl +++ b/src/DynamicPPL.jl @@ -139,17 +139,8 @@ export AbstractVarInfo, get_range_and_transform, get_all_ranges_and_transforms, get_logdensity_callable, - # Leaf contexts - AbstractContext, - contextualize, - DefaultContext, - InitContext, - # Parent contexts - AbstractParentContext, - childcontext, - setchildcontext, - leafcontext, - setleafcontext, + # Contexts + Context, # Tilde pipeline tilde_assume!!, tilde_observe!!, @@ -226,7 +217,7 @@ export AbstractVarInfo, generated_quantities, typed_identity -@compat public getprefix +@compat public make_evaluate_args_and_kwargs, store_coloneq_value!!, getprefix # Reexport using Distributions: loglikelihood @@ -242,6 +233,7 @@ Abstract supertype for data structures that capture random variables when execut probabilistic model and accumulate log densities such as the log likelihood or the log joint probability of the model. +Implement `getaccs` and `setaccs!!` to provide an output container. See also: [`VarInfo`](@ref). """ abstract type AbstractVarInfo <: AbstractModelTrace end @@ -264,7 +256,6 @@ using .VarNamedTuples: include("transformed_values.jl") include("contexts.jl") -include("contexts/default.jl") include("contexts/init.jl") include("model.jl") include("distribution_wrappers.jl") diff --git a/src/accumulators.jl b/src/accumulators.jl index 6cbe5b07f..d7ab6209e 100644 --- a/src/accumulators.jl +++ b/src/accumulators.jl @@ -128,7 +128,8 @@ function combine end promote_for_threadsafe_eval(acc::AbstractAccumulator, ::Type{T}) where {T} Convert `acc` to a new accumulator that works with threadsafe evaluation. The type parameter -`T` is the element type of the parameters that will be used for model evaluation. +`T` is a floating-capable type derived from the parameters, or `Any` when their type is +unknown. Preserve the accumulator's numeric type when `T` is `Any`. See the docstring of `ThreadSafeVarInfo(vi, ::Type{T})` for more details. """ diff --git a/src/compiler.jl b/src/compiler.jl index d78e919bb..302d6dc11 100644 --- a/src/compiler.jl +++ b/src/compiler.jl @@ -648,7 +648,11 @@ function build_output(modeldef, linenumbernode, lhs_names) # Add the internal arguments to the user-specified arguments (positional + keywords). evaluatordef[:args] = vcat( - [:(__model__::$(DynamicPPL.Model)), :(__varinfo__::$(DynamicPPL.AbstractVarInfo))], + [ + :(__model__::$(DynamicPPL.Model)), + :(__context__::$(DynamicPPL.Context)), + :(__varinfo__::$(DynamicPPL.AbstractVarInfo)), + ], args, ) @@ -662,7 +666,6 @@ function build_output(modeldef, linenumbernode, lhs_names) # See the docstrings of `replace_returns` for more info. evaluatordef[:body] = MacroTools.@q begin $(linenumbernode) - __context__ = __model__.context $(replace_returns(add_return_to_last_statment(modeldef[:body]))) end @@ -706,12 +709,13 @@ function build_output(modeldef, linenumbernode, lhs_names) # Pass prepared keywords positionally so applicability checks their types too. definition[:kwargs] = [] definition[:args] = vcat( - definition[:args][1:2], + definition[:args][1:3], [MacroTools.combinearg(n, t, false, nothing) for (n, t, _, _) in kwargs_split], args, ) callargs = Any[ :__model__, + :__context__, :__varinfo__, [ is_splat ? :($(Base.pairs)($(NamedTuple)($n))) : n for @@ -766,11 +770,7 @@ function build_output(modeldef, linenumbernode, lhs_names) $(linenumbernode) $(normalize_kwargs...) return $(DynamicPPL.Model){false}( - $name, - $args_nt, - $kwargs_nt, - $(DynamicPPL.DefaultContext)(); - args_on_lhs=($(QuoteNode(Tuple(args_on_lhs)))), + $name, $args_nt, $kwargs_nt; args_on_lhs=($(QuoteNode(Tuple(args_on_lhs)))) ) end diff --git a/src/contexts.jl b/src/contexts.jl index d4a6ebc04..62104cfa7 100644 --- a/src/contexts.jl +++ b/src/contexts.jl @@ -1,112 +1,10 @@ """ - AbstractParentContext + tilde_assume!!(context::Context, dist::Distribution, vn::VarName, template, vi::AbstractVarInfo) -An abstract context that has a child context. +Obtain a latent value from the context and accumulate its outputs. -Subtypes of `AbstractParentContext` must implement the following interface: - -- `DynamicPPL.childcontext(context::AbstractParentContext)`: Return the child context. -- `DynamicPPL.setchildcontext(parent::AbstractParentContext, child::AbstractContext)`: Reconstruct - `parent` but now using `child` as its child context. -""" -abstract type AbstractParentContext <: AbstractContext end - -""" - childcontext(context::AbstractParentContext) - -Return the descendant context of `context`. -""" -function childcontext end - -""" - setchildcontext(parent::AbstractParentContext, child::AbstractContext) - -Reconstruct `parent` but now using `child` is its [`childcontext`](@ref), -effectively updating the child context. - -""" -function setchildcontext end - -""" - leafcontext(context::AbstractContext) - -Return the leaf of `context`, i.e. the first descendant context that is not an -`AbstractParentContext`. -""" -leafcontext(context::AbstractContext) = context -leafcontext(context::AbstractParentContext) = leafcontext(childcontext(context)) - -""" - setleafcontext(left::AbstractContext, right::AbstractContext) - -Return `left` but now with its leaf context replaced by `right`. - -Note that this also works even if `right` is not a leaf context, -in which case effectively append `right` to `left`, dropping the -original leaf context of `left`. - -# Examples -```jldoctest; setup=:(using Random) -julia> using DynamicPPL: leafcontext, setleafcontext, childcontext, setchildcontext, AbstractContext, InitContext - -julia> struct ParentContext{C} <: AbstractParentContext - context::C - end - -julia> DynamicPPL.childcontext(context::ParentContext) = context.context - -julia> DynamicPPL.setchildcontext(::ParentContext, child) = ParentContext(child) - -julia> Base.show(io::IO, c::ParentContext) = print(io, "ParentContext(", childcontext(c), ")") - -julia> ctx = ParentContext(ParentContext(DefaultContext())) -ParentContext(ParentContext(DefaultContext())) - -julia> # Replace the leaf context with another leaf. - leafcontext(setleafcontext(ctx, InitContext(MersenneTwister(23), InitFromPrior(), UnlinkAll()))) -InitContext{MersenneTwister, InitFromPrior, UnlinkAll}(MersenneTwister(23), InitFromPrior(), UnlinkAll()) - -julia> # Append another parent context. - setleafcontext(ctx, ParentContext(DefaultContext())) -ParentContext(ParentContext(ParentContext(DefaultContext()))) -``` -""" -function setleafcontext(left::AbstractParentContext, right::AbstractContext) - return setchildcontext(left, setleafcontext(childcontext(left), right)) -end -setleafcontext(::AbstractContext, right::AbstractContext) = right - -""" - DynamicPPL.tilde_assume!!( - context::AbstractContext, - right::Distribution, - vn::VarName, - template::Any, - vi::AbstractVarInfo - )::Tuple{Any,AbstractVarInfo} - -Handle assumed variables, i.e. anything which is not observed (see -[`tilde_observe!!`](@ref)). Accumulate the associated log probability, and return the -sampled value and updated `vi`. - -`vn` is the VarName on the left-hand side of the tilde statement. - -`template` is the value of the top-level symbol in `vn`. - -This function should return a tuple `(x, vi)`, where `x` is the sampled value (which must be -untransformed, i.e., `insupport(right, x)` must be true!) and `vi` is the updated VarInfo. +Return the model-space value and updated `vi`. The template describes the enclosing +variable's storage. Extend `init` for custom value selection, or accumulator methods +for custom output handling. """ -function tilde_assume!!( - context::AbstractParentContext, - right::Distribution, - vn::VarName, - template::Any, - vi::AbstractVarInfo, -) - return tilde_assume!!(childcontext(context), right, vn, template, vi) -end -function tilde_assume!!( - context::AbstractContext, ::Distribution, ::VarName, ::Any, ::AbstractVarInfo -) - return error("tilde_assume!! not implemented for context of type $(typeof(context))") -end +function tilde_assume!! end diff --git a/src/contexts/default.jl b/src/contexts/default.jl deleted file mode 100644 index 685052228..000000000 --- a/src/contexts/default.jl +++ /dev/null @@ -1,54 +0,0 @@ -""" - struct DefaultContext <: AbstractContext end - -`DefaultContext`, as the name suggests, is the default context used when instantiating a -model. - -```jldoctest -julia> @model f() = x ~ Normal(); - -julia> model = f(); model.context -DefaultContext() -``` - -As an evaluation context, the behaviour of `DefaultContext` is to require all variables to be -present in the `AbstractVarInfo` used for evaluation. Thus, semantically, evaluating a model -with `DefaultContext` means 'calculating the log-probability associated with the variables -in the `AbstractVarInfo`'. -""" -struct DefaultContext <: AbstractContext end - -""" - DynamicPPL.tilde_assume!!( - ::DefaultContext, - right::Distribution, - vn::VarName, - template::Any, - vi::AbstractVarInfo - ) - -Handle assumed variables. For `DefaultContext`, this function extracts the value associated -with `vn` from `vi`, If `vi` does not contain an appropriate value then this will error. -""" -function tilde_assume!!( - ::DefaultContext, right::Distribution, vn::VarName, template::Any, vi::AbstractVarInfo -) - # TODO(penelopeysm): Conceptually, this is the same as InitContext, except that: - # 1. init(...) is not called; instead we read the value from vi. - # 2. apply_transform_strategy(...) is not called; instead we infer from vi whether the - # value is supposed to be linked or not. - # This can definitely be unified in the future. - tval = get_transformed_value(vi, vn) - trf = if tval.transform isa DynamicLink - Bijectors.VectorBijectors.from_linked_vec(right) - elseif tval.transform isa Unlink - Bijectors.VectorBijectors.from_vec(right) - elseif tval.transform isa FixedTransform - tval.transform.transform - else - error("Expected transformed value to be a vectorised value") - end - x, inv_logjac = with_logabsdet_jacobian(trf, get_internal_value(tval)) - vi = accumulate_assume!!(vi, x, tval, -inv_logjac, vn, right, template) - return x, vi -end diff --git a/src/contexts/init.jl b/src/contexts/init.jl index 2f182ece8..cb4d67ce5 100644 --- a/src/contexts/init.jl +++ b/src/contexts/init.jl @@ -285,55 +285,49 @@ function get_param_eltype(strategy::InitFromVector) end """ - InitContext( - [rng::Random.AbstractRNG=Random.default_rng()], - strategy::AbstractInitStrategy, - transform_strategy::AbstractTransformStrategy, - ) + Context([rng::Random.AbstractRNG,] strategy::AbstractInitStrategy, transform_strategy::AbstractTransformStrategy) + +Supply the inputs for one model evaluation. + +The strategy obtains latent values, and the transform strategy determines their output +representation and log-Jacobian. Observations and fixed values come from the model. +Evaluation never reads parameter values or transforms from its output `VarInfo`. +Inside a model, use `rand(__context__.rng, ...)` for explicit draws from this RNG. + +# Examples -A leaf context that indicates that new values for random variables are currently being -obtained through sampling. Used e.g. when initialising a fresh VarInfo. +```jldoctest +julia> using Random: Xoshiro -The `strategy` argument specifies how new values are to be obtained (see -[`AbstractInitStrategy`](@ref) for details), while the `transform_strategy` argument specifies -whether values should be treated as being in linked or unlinked space. That also means that -`transform_strategy` determines whether the log-Jacobian of the link transform is included when -evaluating the model. +julia> @model example() = x ~ Normal(); -!!! note - If `leafcontext(model.context) isa InitContext`, then `evaluate!!(model, varinfo)` will - override all values in the VarInfo. +julia> ctx = Context(Xoshiro(1), InitFromParams((; x=2.0), nothing), UnlinkAll()); + +julia> result, vi = evaluate!!(example(), ctx, VarInfo(RawValueAccumulator(false))); + +julia> result == get_raw_values(vi)[@varname(x)] +true +``` """ -struct InitContext{ - R<:Random.AbstractRNG,S<:AbstractInitStrategy,L<:AbstractTransformStrategy -} <: AbstractContext +struct Context{R<:Random.AbstractRNG,S<:AbstractInitStrategy,L<:AbstractTransformStrategy} rng::R strategy::S transform_strategy::L +end - function InitContext( - rng::Random.AbstractRNG, - strategy::AbstractInitStrategy, - transform_strategy::AbstractTransformStrategy, - ) - return new{typeof(rng),typeof(strategy),typeof(transform_strategy)}( - rng, strategy, transform_strategy - ) - end - function InitContext( - strategy::AbstractInitStrategy, transform_strategy::AbstractTransformStrategy - ) - return InitContext(Random.default_rng(), strategy, transform_strategy) - end +function Context( + strategy::AbstractInitStrategy, transform_strategy::AbstractTransformStrategy +) + return Context(Random.default_rng(), strategy, transform_strategy) end +get_param_eltype(ctx::Context) = get_param_eltype(ctx.strategy) + function tilde_assume!!( - ctx::InitContext, dist::Distribution, vn::VarName, template::Any, vi::AbstractVarInfo + ctx::Context, dist::Distribution, vn::VarName, template::Any, vi::AbstractVarInfo ) init_tval = init(ctx.rng, vn, dist, ctx.strategy) x, tval, logjac = apply_transform_strategy(ctx.transform_strategy, init_tval, vn, dist) vi = accumulate_assume!!(vi, x, tval, logjac, vn, dist, template) - # We always return the untransformed value here, as that will determine - # what the lhs of the tilde-statement is set to. return x, vi end diff --git a/src/debug_utils.jl b/src/debug_utils.jl index e9a6f0ac1..292807ddd 100644 --- a/src/debug_utils.jl +++ b/src/debug_utils.jl @@ -277,7 +277,7 @@ function has_static_constraints(model::Model; num_evals::Int=5) end """ - gen_evaluator_call_with_types(model[, varinfo]; context=InitContext(InitFromParams(get_values(varinfo), nothing), UnlinkAll())) + gen_evaluator_call_with_types(model[, varinfo]; context=Context(InitFromParams(get_values(varinfo), nothing), UnlinkAll())) Generate the evaluator call and the types of the arguments. @@ -294,13 +294,9 @@ A 2-tuple with the following elements: function gen_evaluator_call_with_types( model::Model, varinfo::AbstractVarInfo=VarInfo(model); - context::AbstractContext=InitContext( - InitFromParams(get_values(varinfo), nothing), UnlinkAll() - ), + context::Context=Context(InitFromParams(get_values(varinfo), nothing), UnlinkAll()), ) - args, kwargs = DynamicPPL.make_evaluate_args_and_kwargs( - setleafcontext(model, context), varinfo - ) + args, kwargs = DynamicPPL.make_evaluate_args_and_kwargs(model, context, varinfo) f, args, kwargs = DynamicPPL._model_evaluator(model.f, args, kwargs) return if isempty(kwargs) (f, Base.typesof(args...)) @@ -310,7 +306,7 @@ function gen_evaluator_call_with_types( end """ - model_warntype(model[, varinfo, optimize=false]; context=InitContext(InitFromParams(get_values(varinfo), nothing), UnlinkAll())) + model_warntype(model[, varinfo, optimize=false]; context=Context(InitFromParams(get_values(varinfo), nothing), UnlinkAll())) Check the type stability of the model's evaluator, warning about any potential issues. @@ -321,7 +317,7 @@ This simply calls `@code_warntype` on the model's evaluator, filling in internal - `varinfo::AbstractVarInfo`: The varinfo to use when evaluating the model. Default: `VarInfo(model)`. # Keyword Arguments -- `context::AbstractContext`: The evaluation context. Defaults to the values supplied in `varinfo`, unlinked. +- `context::Context`: The evaluation context. Defaults to the values supplied in `varinfo`, unlinked. """ function model_warntype( model::Model, varinfo::AbstractVarInfo=VarInfo(model), optimize::Bool=false; kwargs... @@ -331,7 +327,7 @@ function model_warntype( end """ - model_typed(model[, varinfo, optimize=true]; context=InitContext(InitFromParams(get_values(varinfo), nothing), UnlinkAll())) + model_typed(model[, varinfo, optimize=true]; context=Context(InitFromParams(get_values(varinfo), nothing), UnlinkAll())) Return the type inference for the model's evaluator. @@ -342,7 +338,7 @@ This simply calls `@code_typed` on the model's evaluator, filling in internal ar - `varinfo::AbstractVarInfo`: The varinfo to use when evaluating the model. Default: `VarInfo(model)`. # Keyword Arguments -- `context::AbstractContext`: The evaluation context. Defaults to the values supplied in `varinfo`, unlinked. +- `context::Context`: The evaluation context. Defaults to the values supplied in `varinfo`, unlinked. """ function model_typed( model::Model, varinfo::AbstractVarInfo=VarInfo(model), optimize::Bool=true; kwargs... diff --git a/src/logdensityfunction.jl b/src/logdensityfunction.jl index 259826ebb..5113a9ecf 100644 --- a/src/logdensityfunction.jl +++ b/src/logdensityfunction.jl @@ -1,7 +1,7 @@ using DynamicPPL: AbstractVarInfo, AccumulatorTuple, - InitContext, + Context, InitFromVector, AbstractInitStrategy, LogJacobianAccumulator, @@ -126,7 +126,7 @@ For all other fields, please use the corresponding getter functions provided in # Extended help -`LogDensityFunction` supplies parameter inputs through an `InitContext` and collects +`LogDensityFunction` supplies parameter inputs through an `Context` and collects outputs in a `VarInfo`, which holds only evaluation outputs as accumulators, not inputs such as parameter values or transform strategies. diff --git a/src/model.jl b/src/model.jl index cdf126f65..fbe71781f 100644 --- a/src/model.jl +++ b/src/model.jl @@ -980,9 +980,11 @@ end function _reconstruct_model end """ - Model{Threaded}(f, args::NamedTuple, defaults::NamedTuple, context=DefaultContext(); args_on_lhs=()) + Model{Threaded}(f, args::NamedTuple, defaults::NamedTuple; args_on_lhs=()) -Store a model function, arguments, and context. Prefer [`@model`](@ref) for construction. +Store a model function, arguments, prefixes, and bindings. The evaluation context is passed +to [`evaluate!!`](@ref). Prefer [`@model`](@ref) for construction. +For direct construction, use [`condition`](@ref) or [`fix`](@ref) to supply bindings. The names of arguments with LHS variables are stored as immutable type metadata. Set `args_on_lhs` to the tuple of argument names that occur on the left-hand side of `~`, for example `Model{false}(f, (; y=1.0), (;); args_on_lhs=(:y,))`. @@ -1004,14 +1006,12 @@ struct Model{ Tdefaults, Prefix<:Union{VarName,Nothing,PrefixTemplate}, Values<:Union{VarNamedTuple,LocalModelValues,UnprefixedArgumentValues}, - C<:AbstractContext, Threaded, ArgsOnLHS, } <: AbstractProbabilisticProgram f::F args::NamedTuple{argnames,Targs} defaults::NamedTuple{defaultnames,Tdefaults} - context::C prefix::Prefix values::Values function Model{Threaded}( @@ -1019,24 +1019,23 @@ struct Model{ args::NamedTuple{A,Ta}, defaults::NamedTuple{D,Td}, prefix::P, - values::V, - context::C; + values::V; args_on_lhs::Union{Tuple{Vararg{Symbol}},Vector{Symbol}}=(), - ) where {F,A,Ta,D,Td,P,C,V,Threaded} + ) where {F,A,Ta,D,Td,P,V,Threaded} mapreduce( pair -> pair.second isa ModelValue, &, _model_values(values); init=true ) || throw(ArgumentError("Model values must carry a condition or fix role")) argument_names = Tuple(args_on_lhs) - return new{F,A,D,Ta,Td,P,V,C,Threaded,argument_names}( - f, args, defaults, context, prefix, values + return new{F,A,D,Ta,Td,P,V,Threaded,argument_names}( + f, args, defaults, prefix, values ) end # Internal reconstruction reuses already-validated bindings. function DynamicPPL._reconstruct_model( - model::Model{F,A,D,Ta,Td}, prefix::P, context::C, values::V, ::Val{Threaded} - ) where {F,A,D,Ta,Td,P,C,V,Threaded} - return new{F,A,D,Ta,Td,P,V,C,Threaded,_args_on_lhs(model)}( - model.f, model.args, model.defaults, context, prefix, values + model::Model{F,A,D,Ta,Td}, prefix::P, values::V, ::Val{Threaded} + ) where {F,A,D,Ta,Td,P,V,Threaded} + return new{F,A,D,Ta,Td,P,V,Threaded,_args_on_lhs(model)}( + model.f, model.args, model.defaults, prefix, values ) end end @@ -1050,20 +1049,19 @@ Storage templates for nested submodel namespaces remain internal to the model. getprefix(model::Model) = _getprefix(model.prefix) function _args_on_lhs( - ::Model{F,A,D,Ta,Td,P,V,C,Threaded,ArgsOnLHS} -) where {F,A,D,Ta,Td,P,V,C,Threaded,ArgsOnLHS} + ::Model{F,A,D,Ta,Td,P,V,Threaded,ArgsOnLHS} +) where {F,A,D,Ta,Td,P,V,Threaded,ArgsOnLHS} return ArgsOnLHS end Base.@constprop :aggressive function Model{Threaded}( f, args::NamedTuple, - defaults::NamedTuple, - context::AbstractContext=DefaultContext(); + defaults::NamedTuple; args_on_lhs::Union{Tuple{Vararg{Symbol}},Vector{Symbol}}=(), ) where {Threaded} values = _argument_defaults(merge(args, defaults), Val(Tuple(args_on_lhs))) - return Model{Threaded}(f, args, defaults, nothing, values, context; args_on_lhs) + return Model{Threaded}(f, args, defaults, nothing, values; args_on_lhs) end """ @@ -1088,15 +1086,10 @@ end Return whether `model` has been marked as needing threadsafe evaluation (using `setthreadsafe`). """ -requires_threadsafe( - ::Model{F,A,D,Ta,Td,P,V,C,Threaded} -) where {F,A,D,Ta,Td,P,V,C,Threaded} = Threaded -function _reconstruct_model( - model::Model; prefix=model.prefix, context=model.context, values=model.values -) - return _reconstruct_model( - model, prefix, context, values, Val(requires_threadsafe(model)) - ) +requires_threadsafe(::Model{F,A,D,Ta,Td,P,V,Threaded}) where {F,A,D,Ta,Td,P,V,Threaded} = + Threaded +function _reconstruct_model(model::Model; prefix=model.prefix, values=model.values) + return _reconstruct_model(model, prefix, values, Val(requires_threadsafe(model))) end function _materialize_argument_values(model::Model) model.values isa UnprefixedArgumentValues || return model @@ -1108,19 +1101,6 @@ function _materialize_argument_values(model::Model) return _reconstruct_model(model; values) end -""" - contextualize(model::Model, context::AbstractContext) - -Return a model with its context replaced by `context`. -""" -function contextualize(model::Model, context::AbstractContext) - return _reconstruct_model(model; context) -end -"""Return a model with its leaf context replaced by `context`.""" -function setleafcontext(model::Model, context::AbstractContext) - return contextualize(model, setleafcontext(model.context, context)) -end - """ setthreadsafe(model::Model, threadsafe::Bool) @@ -1144,9 +1124,7 @@ function setthreadsafe(model::Model, threadsafe::Bool) return if requires_threadsafe(model) == threadsafe model else - _reconstruct_model( - model, model.prefix, model.context, model.values, Val(threadsafe) - ) + _reconstruct_model(model, model.prefix, model.values, Val(threadsafe)) end end @@ -2307,6 +2285,10 @@ function prefix(model::Model, x) return prefix(model, VarName{Symbol(x)}()) end +optic_skip_length(::AbstractPPL.Iden) = 0 +optic_skip_length(optic::AbstractPPL.Index) = 1 + optic_skip_length(optic.child) +optic_skip_length(optic::AbstractPPL.Property) = 1 + optic_skip_length(optic.child) + function _prefix_varname_and_template(vn::VarName, template::Any, model::Model) return _prefix_varname_and_template(vn, template, getprefix(model), model.prefix) end @@ -2318,7 +2300,7 @@ end function tilde_assume!!( model::Model, - context::AbstractContext, + context::Context, right::Distribution, vn::VarName, template::Any, @@ -2410,11 +2392,12 @@ end [transform_strategy::AbstractTransformStrategy=UnlinkAll(),] ) -Evaluate `model` with the given initialisation and transform strategies, resetting and -filling the accumulators in `varinfo`. Parameter values are recorded only if a value -accumulator is present. +Construct a `Context` and evaluate `model`, resetting and collecting the requested outputs. -`transform_strategy` controls the output representation and defaults to `UnlinkAll()`. +The initialisation strategy supplies latent values. The transform strategy defaults to +`UnlinkAll()`, independently of the contents of `varinfo`. To reuse previous outputs, +explicitly pass `InitFromParams(get_vector_values(previous), nothing)` and the desired +transform strategy. Returns a tuple of the model's return value, plus the updated `varinfo` object. """ @@ -2425,9 +2408,8 @@ function init!!( init_strategy::AbstractInitStrategy, transform_strategy::AbstractTransformStrategy=UnlinkAll(), ) - ctx = InitContext(rng, init_strategy, transform_strategy) - model = DynamicPPL.setleafcontext(model, ctx) - return DynamicPPL.evaluate_nowarn!!(model, vi) + ctx = Context(rng, init_strategy, transform_strategy) + return AbstractPPL.evaluate!!(model, ctx, vi) end function init!!( model::Model, @@ -2439,88 +2421,41 @@ function init!!( end """ - evaluate!!(model::Model, varinfo) - -Evaluate the `model` with the given `varinfo`, wrapping it in a `ThreadSafeVarInfo` if the -model is marked as needing threadsafe evaluation. - -!!! warning - The semantics of this method are complicated. We **strongly** recommend that users do - *not* use this method unless absolutely necessary. In the future this method will be - deprecated and removed. As far as possible (and it should **always** be possible -- - please open an issue if you do not know how to adapt your code!) you should use the - five-argument `init!!([rng,] model, ::VarInfo, init_strategy, - transform_strategy)` method, which has more explicit semantics and allows you to have - more control over each part of the evaluation process. - -The exact semantics depend on the `model`'s context. Fundamentally, this method executes the -model evaluation function (i.e., the function used to define the model) using the given -`varinfo` as an argument. At each tilde-statement, `tilde_assume!!` or `tilde_observe!!` is -called, whose behaviour depends on the model's context. - -Broadly speaking, if the leaf context is an `InitContext`, then this function: - -- uses the initialisation strategy inside the `InitContext`; -- uses the transform strategy inside the `InitContext`; -- uses the accumulators inside `varinfo` (resetting them before evaluation); -- overwrites the values in `varinfo` with the new values obtained from the initialisation strategy. - -If the leaf context is a `DefaultContext`, then this function: - -- uses the values inside the `varinfo` as the initialisation strategy; -- derives a transform strategy from the `varinfo`'s stored variables (if a linked variable is - stored, then the transform strategy will treat that variable as linked; likewise for - unlinked) -- uses the accumulators inside `varinfo` (resetting them before evaluation); -- records the values of executed LHS variables in the reset value accumulator, omitting LHS variables - that are no longer executed. - -The long-term plan for this method is to: - -- Replace `DefaultContext` with `InitContext` by splitting up the functionality of `DefaultContext` - into its constituent components -- Remove the `VarInfo` argument, and instead use only an `AccumulatorTuple` -- Separate the initialisation and transform strategies into separate arguments, instead of storing - them inside the model's context. -""" -function AbstractPPL.evaluate!!(model::Model, varinfo::AbstractVarInfo) - @warn ( - "Calling `evaluate!!(model, varinfo)` directly is not recommended and will be" * - " deprecated in the future. Please switch to using `init!!([rng,] model," * - " ::VarInfo, init_strategy, transform_strategy)` instead, which" * - " has more explicit semantics and allows you to have more control over each" * - " part of the evaluation process. Please see the DynamicPPL documentation" * - " for more details: https://turinglang.org/DynamicPPL.jl/stable/evaluation" - ) maxlog = 5 - return DynamicPPL.evaluate_nowarn!!(model, varinfo) -end + evaluate!!(model::Model, context::Context, varinfo::AbstractVarInfo) -""" - evaluate_nowarn!!(model::Model, varinfo) +Reset the accumulators and evaluate `model` using `context`, returning `(retval, varinfo)`. + +The context belongs to this evaluation, not to the model. The same context is passed to +submodels and to [`tilde_assume!!`](@ref) for latent sites. Observations and tracked values +go directly to accumulators, independently of the context. Models marked with +[`setthreadsafe`](@ref) use a `ThreadSafeVarInfo` during evaluation. + +The [`Context`](@ref) supplies an RNG, initialisation strategy, and transform strategy. +The [`VarInfo`](@ref) contains only output accumulators. The convenience function +[`init!!`](@ref) constructs a `Context` and calls this method. Latent inputs are never +read from the output `varinfo`. + +# Examples + +```jldoctest +julia> using Random: Xoshiro + +julia> @model example(y) = (x ~ Normal(); y ~ Normal(x); return x + y); -This is the same as `evaluate!!(model, varinfo)` but without the deprecation warning. +julia> ctx = Context(Xoshiro(1), InitFromParams((; x=1.0)), UnlinkAll()); -!!! warning - This is meant for internal use in DynamicPPL.jl only! If you rely on this method in your - code, please note that it may break at any time. +julia> retval, vi = evaluate!!(example(2.0), ctx, VarInfo()); + +julia> retval +3.0 +``` """ -function evaluate_nowarn!!(model::Model, varinfo::AbstractVarInfo) - if leafcontext(model.context) isa DefaultContext - values, strategy = if hasacc(varinfo, Val(VECTORVAL_ACCNAME)) - copy(get_vector_values(varinfo)), - infer_transform_strategy_from_values(get_vector_values(varinfo)) - else - VarNamedTuple(), UnlinkAll() - end - ctx = InitContext(InitFromParams(values, nothing), strategy) - model = setleafcontext(model, ctx) - end +function AbstractPPL.evaluate!!(model::Model, context::Context, varinfo::AbstractVarInfo) return if requires_threadsafe(model) - # Use of float_type_with_fallback(eltype(x)) is necessary to deal with cases where x is - # a gradient type of some AD backend. - param_eltype = DynamicPPL.get_param_eltype(varinfo, model.context) + # Thread-local accumulators must accept AD values before evaluation starts. + param_eltype = DynamicPPL.get_param_eltype(context) wrapper = ThreadSafeVarInfo(varinfo, param_eltype) - result, wrapper_new = _evaluate!!(model, wrapper) + result, wrapper_new = _evaluate!!(model, context, wrapper) # TODO(penelopeysm): If seems that if you pass a TSVI to this method, it # will return the underlying VI, which is a bit counterintuitive (because # calling TSVI(::TSVI) returns the original TSVI, instead of wrapping it @@ -2534,42 +2469,42 @@ function evaluate_nowarn!!(model::Model, varinfo::AbstractVarInfo) end return result, setaccs!!(wrapper_new.varinfo, accs) else - _evaluate!!(model, resetaccs!!(varinfo)) + _evaluate!!(model, context, resetaccs!!(varinfo)) end end """ - _evaluate!!(model::Model, varinfo) + _evaluate!!(model::Model, context::Context, varinfo) -Evaluate the `model` with the given `varinfo`. +Evaluate the `model` with the given `context` and `varinfo`. This function does not wrap the varinfo in a `ThreadSafeVarInfo`. It also does not reset the log probability of the `varinfo` before running. """ -function _evaluate!!(model::Model, varinfo::AbstractVarInfo) - args, kwargs = make_evaluate_args_and_kwargs(model, varinfo) +function _evaluate!!(model::Model, context::Context, varinfo::AbstractVarInfo) + args, kwargs = make_evaluate_args_and_kwargs(model, context, varinfo) return model.f(args...; kwargs...) end """ - make_evaluate_args_and_kwargs(model, varinfo) + make_evaluate_args_and_kwargs(model, context, varinfo) + +Return the positional and keyword arguments for `model.f`, including the evaluation context. -Return the arguments and keyword arguments to be passed to the evaluator of the model, i.e. `model.f`e. +The positional arguments begin with `(model, context, varinfo)`, followed by the model +arguments converted for the parameter element type. Pass the result to +`model.f(args...; kwargs...)` when a downstream evaluator, such as a taped task, controls +execution directly. This prepares arguments without executing the model, resetting +accumulators, or wrapping `varinfo` for thread safety; use [`evaluate!!`](@ref) otherwise. """ @generated function make_evaluate_args_and_kwargs( - model::Model{_F,argnames,defaultnames}, varinfo::AbstractVarInfo + model::Model{_F,argnames,defaultnames}, context::Context, varinfo::AbstractVarInfo ) where {_F,argnames,defaultnames} unwrap_args = [ if is_splat_symbol(var) - :( - $convert_model_argument( - $get_param_eltype(varinfo, model.context), model.args.$var - )... - ) + :($convert_model_argument($get_param_eltype(context), model.args.$var)...) else - :($convert_model_argument( - $get_param_eltype(varinfo, model.context), model.args.$var - )) + :($convert_model_argument($get_param_eltype(context), model.args.$var)) end for var in argnames ] unwrap_kwargs = [ @@ -2577,37 +2512,12 @@ Return the arguments and keyword arguments to be passed to the evaluator of the var in defaultnames ] return quote - args = (model, varinfo, $(unwrap_args...)) + args = (model, context, varinfo, $(unwrap_args...)) kwargs = (; $(unwrap_kwargs...)) return args, kwargs end end -""" - get_param_eltype(varinfo::AbstractVarInfo, context::AbstractContext) - -Get the element type of the parameters being used to evaluate a model, using a `varinfo` -under the given `context`. For example, when evaluating a model with ForwardDiff AD, this -should return `ForwardDiff.Dual`. - -For `InitContext`, query its initialisation strategy. For other leaf contexts, infer -this type from recorded vectorised values, or return `Union{}` if no value accumulator -is present. Parent contexts delegate to their child context. - -See the docstring of `get_param_eltype(strategy::AbstractInitStrategy)` for the -strategy interface. -""" -function get_param_eltype(vi::AbstractVarInfo, ctx::AbstractParentContext) - return get_param_eltype(vi, DynamicPPL.childcontext(ctx)) -end -function get_param_eltype(vi::AbstractVarInfo, ::AbstractContext) - hasacc(vi, Val(VECTORVAL_ACCNAME)) || return Union{} - return get_param_eltype(InitFromParams(get_vector_values(vi), nothing)) -end -function get_param_eltype(::AbstractVarInfo, ctx::InitContext) - return get_param_eltype(ctx.strategy) -end - @generated function _argument_names(::NamedTuple{names}) where {names} return QuoteNode(map(unsplat_symbol, names)) end @@ -2815,11 +2725,3 @@ function returned(model::Model, parameters...) ), ) end - -function AbstractPPL.evaluate!!(model::Model, context::AbstractContext, vi::AbstractVarInfo) - return evaluate_nowarn!!(setleafcontext(model, context), vi) -end - -optic_skip_length(::AbstractPPL.Iden) = 0 -optic_skip_length(optic::AbstractPPL.Index) = 1 + optic_skip_length(optic.child) -optic_skip_length(optic::AbstractPPL.Property) = 1 + optic_skip_length(optic.child) diff --git a/src/submodel.jl b/src/submodel.jl index ef7584f6c..6129d1f82 100644 --- a/src/submodel.jl +++ b/src/submodel.jl @@ -202,7 +202,7 @@ end """ DynamicPPL.tilde_assume!!( parent_model::Model, - context::AbstractContext, + context::Context, submodel::DynamicPPL.Submodel, left_vn::VarName, template, @@ -213,7 +213,7 @@ Evaluate `submodel` under `parent_model`. """ function tilde_assume!!( parent_model::Model, - context::AbstractContext, + context::Context, submodel::Submodel{M,AutoPrefix}, left_vn::VarName, template, @@ -256,7 +256,7 @@ end # Specialize child evaluation on the selected submodel namespace bindings. @inline function _evaluate_submodel!!( parent_model::Model, - context::AbstractContext, + context::Context, submodel::Submodel{M,AutoPrefix}, left_vn::VarName, template, @@ -272,8 +272,7 @@ end end # Calling model.f directly avoids the inference recursion limit as nested prefixes # change the Model type; routing through _evaluate!! widens it to Any (Turing.jl#2844). - model = setleafcontext(model, context) - args, kwargs = make_evaluate_args_and_kwargs(model, vi) + args, kwargs = make_evaluate_args_and_kwargs(model, context, vi) return model.f(args...; kwargs...) end diff --git a/src/varinfo.jl b/src/varinfo.jl index da02d3feb..b3d5f502b 100644 --- a/src/varinfo.jl +++ b/src/varinfo.jl @@ -7,7 +7,7 @@ Collect model-evaluation outputs in accumulators. The default accumulators record log prior, log likelihood, and log Jacobian. Add a `RawValueAccumulator` or `VectorValueAccumulator` to record parameter values. -Inputs, including the transform strategy, belong to `InitContext`; this type stores +Inputs, including the transform strategy, belong to `Context`; this type stores no independent parameter values or transform state. """ struct VarInfo{Accs<:AccumulatorTuple} <: AbstractVarInfo @@ -38,7 +38,7 @@ end Evaluate `model` and collect vectorised parameter values and log densities. To select different outputs, pass `VarInfo(accumulators...)` to `evaluate!!` with -an explicit `InitContext`. +an explicit `Context`. """ function VarInfo( rng::Random.AbstractRNG, @@ -47,7 +47,7 @@ function VarInfo( transform_strategy::AbstractTransformStrategy=UnlinkAll(), ) vi = VarInfo(VectorValueAccumulator(), default_accumulators()...) - return last(evaluate!!(model, InitContext(rng, init_strategy, transform_strategy), vi)) + return last(evaluate!!(model, Context(rng, init_strategy, transform_strategy), vi)) end function VarInfo( model::Model, @@ -123,7 +123,7 @@ Leave other accumulators unchanged. function update_transform_status!!( vi::VarInfo, strategy::AbstractTransformStrategy, model::Model ) - ctx = InitContext(InitFromParams(get_vector_values(vi), nothing), strategy) + ctx = Context(InitFromParams(get_vector_values(vi), nothing), strategy) outputs = VarInfo(VectorValueAccumulator(), LogJacobianAccumulator()) _, outputs = evaluate!!(model, ctx, outputs) vi = _set_vector_values!!(vi, get_vector_values(outputs)) diff --git a/test/compiler.jl b/test/compiler.jl index 619198b18..faf6eef1d 100644 --- a/test/compiler.jl +++ b/test/compiler.jl @@ -228,8 +228,7 @@ end varinfo = VarInfo(model) @test getlogjoint(varinfo) == lp @test varinfo_ isa AbstractVarInfo - @test model_.f === model.f - @test model_.context isa InitContext + @test model_ === model # disable warnings @model function testmodel_missing4(x) diff --git a/test/conditionfix.jl b/test/conditionfix.jl index 8e742969e..8f897c246 100644 --- a/test/conditionfix.jl +++ b/test/conditionfix.jl @@ -1512,7 +1512,7 @@ end parent = fix(parent, @varname(child.x.a) => 3.0) result, _ = @inferred evaluate!!( parent, - InitContext( + Context( InitFromParams(VarNamedTuple(), nothing), DynamicPPL.infer_transform_strategy_from_values(VarNamedTuple()), ), diff --git a/test/context_implementations.jl b/test/context_implementations.jl index ba4dbe04e..f9e4b1b7b 100644 --- a/test/context_implementations.jl +++ b/test/context_implementations.jl @@ -10,7 +10,6 @@ using LinearAlgebra: I, norm using Random: Xoshiro using Test -struct ObservationContext <: AbstractContext end struct UnimplementedStrategy <: AbstractInitStrategy end struct ObserveHookDistribution <: ContinuousUnivariateDistribution @@ -59,11 +58,13 @@ end model = condition( observations(1.0, [2.0, 3.0], ObserveHookDistribution(seen)); z=4.0 ) - _, vi = DynamicPPL.evaluate_nowarn!!(model, VarInfo()) + _, vi = evaluate!!(model, Context(UnimplementedStrategy(), UnlinkAll()), VarInfo()) @test seen == [@varname(x), @varname(z), nothing, @varname(ys[1]), @varname(ys[2])] @test getloglikelihood(vi) ≈ sum(logpdf.(Normal(), 0.0:4.0)) empty!(seen) - _, vi = DynamicPPL.evaluate_nowarn!!(outer(model), VarInfo()) + _, vi = evaluate!!( + outer(model), Context(UnimplementedStrategy(), UnlinkAll()), VarInfo() + ) @test length(seen) == 5 @test getloglikelihood(vi) ≈ sum(logpdf.(Normal(), 0.0:4.0)) end @@ -81,24 +82,6 @@ end @test Set(keys(get_vector_values(vi))) == expected end - @testset "leaf contexts without parameter outputs" begin - @model function observed_only(x) - x ~ Normal() - 0.0 ~ Normal() - return x - end - for context in (ObservationContext(), DefaultContext()), threaded in (false, true) - model = contextualize(setthreadsafe(observed_only(1.0), threaded), context) - value, vi = DynamicPPL.evaluate_nowarn!!(model, VarInfo()) - @test value == 1.0 - @test getloglikelihood(vi) == logpdf(Normal(), 1.0) + logpdf(Normal(), 0.0) - end - @model latent() = x ~ Normal() - @test_throws "No value was provided" DynamicPPL.evaluate_nowarn!!( - latent(), VarInfo() - ) - end - @testset "observations do not dispatch on context" begin for T in (Float32, Float64, BigFloat) dist = Normal(zero(T), one(T)) @@ -118,7 +101,7 @@ end @testset "no context hooks needed without latent LHS variables" begin accs = VarInfo((DynamicPPL.default_accumulators()..., RawValueAccumulator(true))) result, vi = evaluate!!( - fix(child(); x=1.0), InitContext(UnimplementedStrategy(), UnlinkAll()), accs + fix(child(); x=1.0), Context(UnimplementedStrategy(), UnlinkAll()), accs ) @test result == 3.0 @test get_raw_values(vi)[@varname(z)] == 3.0 @@ -133,7 +116,7 @@ end for x in (1.0, 3.0) strategy = InitFromParams(VarNamedTuple((@varname(b.a.x) => x,))) recording = RecordingStrategy(strategy) - ctx = InitContext(Xoshiro(1), recording, UnlinkAll()) + ctx = Context(Xoshiro(1), recording, UnlinkAll()) accs = VarInfo(( DynamicPPL.default_accumulators()..., RawValueAccumulator(true) )) @@ -150,7 +133,7 @@ end end @testset "Context supplies inputs independently of outputs" begin - empty_context = InitContext( + empty_context = Context( Xoshiro(1), InitFromParams(VarNamedTuple(), nothing), UnlinkAll() ) @test_throws ErrorException evaluate!!( @@ -159,7 +142,7 @@ end for T in (Float32, Float64, BigFloat), threaded in (false, true) model = setthreadsafe(child(T(2)), threaded) input = VarInfo(Xoshiro(1), model, InitFromParams((; x=one(T)))) - context = InitContext( + context = Context( Xoshiro(1), InitFromParams(get_vector_values(input), nothing), UnlinkAll() ) for output in @@ -185,7 +168,7 @@ end # Inputs determine transforms even when the reused output recorded linked values. @model positive() = x ~ Exponential() old = VarInfo(positive(), InitFromParams((; x=2.0)), LinkAll()) - context = InitContext(Xoshiro(1), InitFromParams((; x=3.0), nothing), UnlinkAll()) + context = Context(Xoshiro(1), InitFromParams((; x=3.0), nothing), UnlinkAll()) result, output = evaluate!!(positive(), context, old) @test result == 3.0 @test iszero(getlogjac(output)) @@ -198,10 +181,10 @@ end return y ~ Normal() end outputs = VarInfo(VectorValueAccumulator(), RawValueAccumulator(false)) - context = InitContext(Xoshiro(1), InitFromPrior(), UnlinkAll()) + context = Context(Xoshiro(1), InitFromPrior(), UnlinkAll()) _, outputs = evaluate!!(optional_lhs(true), context, outputs) inputs = get_vector_values(outputs) - context = InitContext(Xoshiro(1), InitFromParams(inputs, nothing), UnlinkAll()) + context = Context(Xoshiro(1), InitFromParams(inputs, nothing), UnlinkAll()) _, outputs = evaluate!!(optional_lhs(false), context, outputs) @test !haskey(get_vector_values(outputs), @varname(x)) @test !haskey(get_raw_values(outputs), @varname(x)) @@ -214,7 +197,7 @@ end end model = dependent_support() input = VarInfo(Xoshiro(1), model, InitFromParams((; x=2.0, y=0.5)), LinkAll()) - context = InitContext( + context = Context( InitFromParams(get_values(input), nothing), DynamicPPL.infer_transform_strategy_from_values(get_values(input)), ) diff --git a/test/contexts/init.jl b/test/contexts/init.jl index 9a3d62583..c6958c890 100644 --- a/test/contexts/init.jl +++ b/test/contexts/init.jl @@ -216,7 +216,7 @@ using Test TransformedValue([missing], Unlink()), ) strategy = InitFromParams((; x=value), nothing) - context = InitContext(Xoshiro(1), strategy, UnlinkAll()) + context = Context(Xoshiro(1), strategy, UnlinkAll()) @test_throws ArgumentError evaluate!!( missing_parameter(), context, VarInfo(()) ) diff --git a/test/debug_utils.jl b/test/debug_utils.jl index b71e5e1a5..1966edb62 100644 --- a/test/debug_utils.jl +++ b/test/debug_utils.jl @@ -32,8 +32,7 @@ end (; y, checking=true), __model__.defaults, __model__.prefix, - __model__.values, - __model__.context; + __model__.values; args_on_lhs=DynamicPPL._args_on_lhs(__model__), ) @test check_model(child) @@ -220,7 +219,7 @@ end @test codeinfo isa Core.CodeInfo @test retype <: Tuple - context = InitContext(Xoshiro(1), InitFromParams((; y=2.0)), UnlinkAll()) + context = Context(Xoshiro(1), InitFromParams((; y=2.0)), UnlinkAll()) _, retype = DynamicPPL.DebugUtils.model_typed(model, VarInfo(); context) @test retype <: Tuple{Float64,VarInfo} diff --git a/test/ext/DynamicPPLBridgeStanExt.jl b/test/ext/DynamicPPLBridgeStanExt.jl index 28cb73eed..e5ccf27bf 100644 --- a/test/ext/DynamicPPLBridgeStanExt.jl +++ b/test/ext/DynamicPPLBridgeStanExt.jl @@ -217,6 +217,27 @@ parameters { @test logjoint(model, (; theta=invalid)) == -Inf @test logprior(model, (; theta=invalid)) == -Inf + context = Context( + InitFromParams( + VarNamedTuple(; + theta=TransformedValue(u, FixedTransform(distribution.transform)) + ), + nothing, + ), + DynamicPPL.infer_transform_strategy_from_values( + VarNamedTuple(; + theta=TransformedValue(u, FixedTransform(distribution.transform)) + ), + ), + ) + for output in + (VarInfo(), VarInfo(VectorValueAccumulator(), DynamicPPL.default_accumulators()...)) + result, output = evaluate!!(model, context, output) + @test result ≈ constrained + @test getlogjoint(output) ≈ expected_unlinked + @test getlogjac(output) ≈ -logjac + end + conditioned_model = condition(model, (; theta=constrained)) @test logjoint(conditioned_model, NamedTuple()) ≈ expected_unlinked diff --git a/test/ext/DynamicPPLForwardDiffExt.jl b/test/ext/DynamicPPLForwardDiffExt.jl index 38d5aa417..3dbe7fbde 100644 --- a/test/ext/DynamicPPLForwardDiffExt.jl +++ b/test/ext/DynamicPPLForwardDiffExt.jl @@ -62,4 +62,38 @@ end ) isa Any end +@model function observed_input(y) + x ~ Normal() + y ~ Normal(x) + 0.0 ~ Normal(x) + return x +end +@model nested_input(y) = a ~ to_submodel(observed_input(y)) + +@testset "Context parameter types come from inputs" begin + for T in (Float32, Float64, BigFloat) + y = T(2) + x = T(0.5) + for (model, vn) in + ((observed_input(y), @varname(x)), (nested_input(y), @varname(a.x))) + for threaded in (false, true) + m = setthreadsafe(model, threaded) + function density(value) + values = VarNamedTuple((vn => TransformedValue([value], Unlink()),)) + _, output = evaluate!!( + m, + Context( + InitFromParams(values, nothing), + DynamicPPL.infer_transform_strategy_from_values(values), + ), + VarInfo(), + ) + return getlogjoint(output) + end + @test ForwardDiff.derivative(density, x) ≈ y - 3x + end + end + end +end + end diff --git a/test/model.jl b/test/model.jl index 1af3e5f95..cd6c88423 100644 --- a/test/model.jl +++ b/test/model.jl @@ -79,9 +79,7 @@ const GDEMO_DEFAULT = DynamicPPL.TestUtils.demo_assume_observe_literal() for m in ( DynamicPPL.Model{false}(f, (; x=missing, y=1.0), (;)), DynamicPPL.Model{false}(f, (; x=missing); y=1.0), - DynamicPPL.Model{false}( - f, (; x=missing, y=1.0), (;), nothing, VarNamedTuple(), DefaultContext() - ), + DynamicPPL.Model{false}(f, (; x=missing, y=1.0), (;), nothing, VarNamedTuple()), ) @test isempty(conditioned(m)) @test isempty(DynamicPPL._args_on_lhs(m)) @@ -105,8 +103,7 @@ const GDEMO_DEFAULT = DynamicPPL.TestUtils.demo_assume_observe_literal() model.args, model.defaults, model.prefix, - model.values, - model.context; + model.values; args_on_lhs=DynamicPPL._args_on_lhs(model), ) @test direct.prefix === model.prefix @@ -432,7 +429,7 @@ const GDEMO_DEFAULT = DynamicPPL.TestUtils.demo_assume_observe_literal() @inferred( evaluate!!( model, - InitContext( + Context( InitFromParams(get_values(varinfo), nothing), UnlinkAll() ), VarInfo(), @@ -446,7 +443,7 @@ const GDEMO_DEFAULT = DynamicPPL.TestUtils.demo_assume_observe_literal() @inferred( evaluate!!( model, - InitContext( + Context( InitFromParams(get_values(varinfo_linked), nothing), LinkAll(), ), diff --git a/test/runtests.jl b/test/runtests.jl index 74947f594..1ec57c25c 100644 --- a/test/runtests.jl +++ b/test/runtests.jl @@ -36,9 +36,9 @@ Random.seed!(100) include("pointwise_logdensities.jl") include("lkj.jl") + include("prefix.jl") include("contexts/init.jl") include("conditionfix.jl") - include("prefix.jl") include("context_implementations.jl") include("threadsafe.jl") include("debug_utils.jl") @@ -77,7 +77,7 @@ Random.seed!(100) # why...) -- if we don't import them here then the doctest output will include # the prefixed module name using Distributions: Normal - using DynamicPPL: DefaultContext, Condition, Fix + using DynamicPPL: Context Documenter.doctest(DynamicPPL; manual=false, doctestfilters=doctestfilters) end diff --git a/test/submodels.jl b/test/submodels.jl index d6cca53d1..311f51850 100644 --- a/test/submodels.jl +++ b/test/submodels.jl @@ -839,9 +839,7 @@ end vi = VarInfo(model) @test @inferred( evaluate!!( - model, - InitContext(InitFromParams(get_values(vi), nothing), UnlinkAll()), - vi, + model, Context(InitFromParams(get_values(vi), nothing), UnlinkAll()), vi ) ) isa Tuple end diff --git a/test/threadsafe.jl b/test/threadsafe.jl index d4833581b..83101fe0a 100644 --- a/test/threadsafe.jl +++ b/test/threadsafe.jl @@ -261,7 +261,7 @@ end threadsafe_model = setthreadsafe(model, true) expected = VarInfo(Xoshiro(1), model) vi = VarInfo(Xoshiro(1), threadsafe_model) - ctx = InitContext(Xoshiro(1), InitFromPrior(), UnlinkAll()) + ctx = Context(Xoshiro(1), InitFromPrior(), UnlinkAll()) outputs() = VarInfo(VectorValueAccumulator(), DynamicPPL.default_accumulators()...) for result in ( vi, @@ -403,7 +403,7 @@ end # But init!! should return the original VarInfo @test vi isa DynamicPPL.VarInfo # Same with evaluate!! - ctx = InitContext(Xoshiro(1), InitFromParams((; x=2.0)), UnlinkAll()) + ctx = Context(Xoshiro(1), InitFromParams((; x=2.0)), UnlinkAll()) result, vi = evaluate!!(model, ctx, vi) @test result == 2.0 @test vi_ isa DynamicPPL.ThreadSafeVarInfo diff --git a/test/transformed_values.jl b/test/transformed_values.jl index 8316c01be..5f98fb0a3 100644 --- a/test/transformed_values.jl +++ b/test/transformed_values.jl @@ -174,7 +174,9 @@ end ) strategy = DynamicPPL.infer_transform_strategy_from_values(get_vector_values(vi)) @test DynamicPPL.target_transform(strategy, @varname(x)) === ft - retval, vi = DynamicPPL.evaluate_nowarn!!(single(), vi) + retval, vi = evaluate!!( + single(), Context(InitFromParams(get_vector_values(vi), nothing), strategy), vi + ) @test retval == 3.0 @test get_transform(get_vector_values(vi)[@varname(x)]) === ft end diff --git a/test/utils.jl b/test/utils.jl index f86712f45..70730bfca 100644 --- a/test/utils.jl +++ b/test/utils.jl @@ -61,7 +61,7 @@ end model = test() for transform in (UnlinkAll(), LinkAll()) input = VarInfo(Xoshiro(1), model, InitFromPrior(), transform) - context = InitContext(InitFromParams(get_values(input), nothing), transform) + context = Context(InitFromParams(get_values(input), nothing), transform) value, output = evaluate!!(model, context, VarInfo()) @test getlogjoint(output) ≈ logpdf(dist, value) @test getlogjac(output) ≈ getlogjac(input) diff --git a/test/varinfo.jl b/test/varinfo.jl index 286e7ffa5..0363a0b25 100644 --- a/test/varinfo.jl +++ b/test/varinfo.jl @@ -185,7 +185,7 @@ end vi = last( evaluate!!( m, - InitContext( + Context( InitFromParams(get_values(vi), nothing), DynamicPPL.infer_transform_strategy_from_values(get_values(vi)), ), @@ -230,7 +230,7 @@ end vi = last( evaluate!!( m, - InitContext( + Context( InitFromParams(get_values(vi), nothing), DynamicPPL.infer_transform_strategy_from_values(get_values(vi)), ), @@ -253,7 +253,7 @@ end # Test evaluating without any accumulators. vi = last( evaluate!!( - m, InitContext(InitFromParams(values, nothing), UnlinkAll()), VarInfo(()) + m, Context(InitFromParams(values, nothing), UnlinkAll()), VarInfo(()) ), ) @test_throws "Missing accumulator :LogPrior." getlogprior(vi) @@ -278,7 +278,7 @@ end # And evaluate the model once so that they are populated. _, vi_orig = evaluate!!( model, - InitContext( + Context( InitFromParams(get_values(vi_orig), nothing), DynamicPPL.infer_transform_strategy_from_values(get_values(vi_orig)), ), @@ -329,7 +329,7 @@ end # Thus after re-evaluation, the accs should be exactly the same as before. _, vi = evaluate!!( model, - InitContext( + Context( InitFromParams(get_values(vi_orig), nothing), DynamicPPL.infer_transform_strategy_from_values(get_values(vi_orig)), ), From a193d58ad7ebfec0556b53b007c4b5f50ad69aa1 Mon Sep 17 00:00:00 2001 From: Hong Ge Date: Wed, 16 Sep 2026 21:14:14 +0100 Subject: [PATCH 2/5] tests: reuse BridgeStan context inputs Assisted-by: Codex --- test/ext/DynamicPPLBridgeStanExt.jl | 16 +++++----------- 1 file changed, 5 insertions(+), 11 deletions(-) diff --git a/test/ext/DynamicPPLBridgeStanExt.jl b/test/ext/DynamicPPLBridgeStanExt.jl index e5ccf27bf..9ab84eafe 100644 --- a/test/ext/DynamicPPLBridgeStanExt.jl +++ b/test/ext/DynamicPPLBridgeStanExt.jl @@ -217,18 +217,12 @@ parameters { @test logjoint(model, (; theta=invalid)) == -Inf @test logprior(model, (; theta=invalid)) == -Inf + values = VarNamedTuple(; + theta=TransformedValue(u, FixedTransform(distribution.transform)) + ) context = Context( - InitFromParams( - VarNamedTuple(; - theta=TransformedValue(u, FixedTransform(distribution.transform)) - ), - nothing, - ), - DynamicPPL.infer_transform_strategy_from_values( - VarNamedTuple(; - theta=TransformedValue(u, FixedTransform(distribution.transform)) - ), - ), + InitFromParams(values, nothing), + DynamicPPL.infer_transform_strategy_from_values(values), ) for output in (VarInfo(), VarInfo(VectorValueAccumulator(), DynamicPPL.default_accumulators()...)) From 0c659ad4228485b57cf0674c246b9f06d4e7ecb1 Mon Sep 17 00:00:00 2001 From: Hong Ge Date: Thu, 1 Oct 2026 13:40:26 +0100 Subject: [PATCH 3/5] contexts: align glossary and document evaluation API migration Assisted-by: Codex --- HISTORY.md | 12 ++++++++++-- benchmarks/benchmarks.jl | 2 +- docs/src/api.md | 2 +- docs/src/migration.md | 21 ++++++++++++++++++--- src/model.jl | 2 +- 5 files changed, 31 insertions(+), 8 deletions(-) diff --git a/HISTORY.md b/HISTORY.md index 4b02b1345..2b015c47d 100644 --- a/HISTORY.md +++ b/HISTORY.md @@ -20,7 +20,7 @@ Added `evaluate!!(model, context, vi)` to evaluate with an explicit `Context` an Added the `template` keyword to `prefix(model, x::VarName)` to supply the enclosing container's shape and resolve `begin` and `end` in indexed prefixes. See [#1502](https://github.com/TuringLang/DynamicPPL.jl/pull/1502). -Added the `context` keyword to `DynamicPPL.DebugUtils.model_typed`, `model_warntype`, and `gen_evaluator_call_with_types` to select the evaluation context. See [#1501](https://github.com/TuringLang/DynamicPPL.jl/pull/1501). +Added the `context` keyword to `DynamicPPL.DebugUtils.model_typed`, `model_warntype`, and `gen_evaluator_call_with_types` to select the evaluation context. See [#1503](https://github.com/TuringLang/DynamicPPL.jl/pull/1503). `subsample` and `independent_problem` now accept argument-supplied observations; previously observations had to be supplied through `condition`. See [#1501](https://github.com/TuringLang/DynamicPPL.jl/pull/1501). @@ -30,7 +30,15 @@ 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). -`Context(rng, init_strategy, transform_strategy)` replaces `InitContext` and the context hierarchy. Pass it to `evaluate!!(model, context, outputs)`; custom initialisation and observation handling belong to strategies and accumulators. +`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). + +`AbstractContext` / `AbstractParentContext` subtyping → initialisation strategies for custom value selection and accumulators for custom output handling. These names are no longer exported. See [#1503](https://github.com/TuringLang/DynamicPPL.jl/pull/1503). + +`childcontext`, `setchildcontext`, `leafcontext`, and `setleafcontext` are removed: use a single `Context(rng, init_strategy, transform_strategy)` passed directly to evaluation instead of traversing or rebuilding a context hierarchy. See [#1503](https://github.com/TuringLang/DynamicPPL.jl/pull/1503). + +`make_evaluate_args_and_kwargs(m, vi)` → `DynamicPPL.make_evaluate_args_and_kwargs(m, ctx, vi)`, now `public`; the prepared positional arguments begin with `(m, ctx, vi)`. See [#1503](https://github.com/TuringLang/DynamicPPL.jl/pull/1503). `PrefixContext` and `extract_prefixes` are removed: use `prefix(model, vn; template)` to set prefixes and `DynamicPPL.getprefix(model)` to read the combined prefix (`nothing` when absent). The `prefix` field stores internal metadata for LHS variable addresses and nested submodel namespace storage templates. `condition` and `fix` reject binding addresses outside the model's prefix; use the prefixed address, such as `@varname(p.y)`, instead of `y`. See [#1502](https://github.com/TuringLang/DynamicPPL.jl/pull/1502). diff --git a/benchmarks/benchmarks.jl b/benchmarks/benchmarks.jl index c6d063bcc..b0da0fa9a 100644 --- a/benchmarks/benchmarks.jl +++ b/benchmarks/benchmarks.jl @@ -118,7 +118,7 @@ end @model _indexed_observation(obs, mu) = obs ~ Normal(mu, 1) -"Indexed submodels with argument observations and one shared latent mean." +"Indexed submodels with argument-supplied observations and one shared latent mean." @model function indexed_submodels(obs) mu ~ Normal() x = similar(obs) diff --git a/docs/src/api.md b/docs/src/api.md index 0f994a620..228a4c019 100644 --- a/docs/src/api.md +++ b/docs/src/api.md @@ -477,7 +477,7 @@ Call `evaluate!!(model, context, varinfo)` to evaluate with an explicit context outputs in `varinfo`. Accumulators are reset before evaluation. The context is an evaluation input; it is not stored in the model. -Prefixes are stored separately from values. Conditioned and fixed values share one store, with each value carrying its role. Only latent sites reach the context; observations and tracked values go directly to accumulators. +Prefixes are stored separately from values. Conditioned and fixed values share one store, with each value carrying its role. Only latent LHS variables reach the context; observations and tracked values go directly to accumulators. `Context` is the sole evaluation context. It supplies an RNG, an initialisation strategy, and a transform strategy. The output `varinfo` never supplies latent inputs. diff --git a/docs/src/migration.md b/docs/src/migration.md index 41117a6cb..7126e4982 100644 --- a/docs/src/migration.md +++ b/docs/src/migration.md @@ -14,9 +14,24 @@ to set a prefix and `DynamicPPL.getprefix(model)` to read the combined prefix (`nothing` when absent). The `prefix` field stores internal metadata for LHS variable addresses and nested submodel namespace storage templates. -Replace `InitContext` with `Context`. `DefaultContext` and context subtyping are removed: -custom value selection belongs in initialisation strategies. To reuse previous values, -extract them explicitly before evaluating: +Replace `DefaultContext()` / `InitContext(...)` with +`Context(rng, init_strategy, transform_strategy)`, specifying parameter inputs and output +transforms explicitly. `Model.context` and `contextualize` are removed: +`contextualize(m, ctx); evaluate!!(m, vi)` becomes `evaluate!!(m, ctx, vi)`. +The two-argument `evaluate!!(m, vi)` is removed. + +`AbstractContext` and `AbstractParentContext` are no longer exported, and context +subtyping is replaced by initialisation strategies for custom value selection and +accumulators for custom output handling. `childcontext`, `setchildcontext`, `leafcontext`, +and `setleafcontext` are removed; pass a single `Context` directly to evaluation instead +of traversing or rebuilding a context hierarchy. + +For downstream evaluators, replace `make_evaluate_args_and_kwargs(m, vi)` with +`DynamicPPL.make_evaluate_args_and_kwargs(m, ctx, vi)`, now `public`. Its prepared +positional arguments begin with `(m, ctx, vi)`. It does not execute the model, reset +accumulators, or wrap them for thread safety; use `evaluate!!` for those steps. + +To reuse previous values, extract them explicitly before evaluating: ```julia context = Context(rng, InitFromParams(get_vector_values(previous), nothing), LinkAll()) diff --git a/src/model.jl b/src/model.jl index fbe71781f..59dcf9075 100644 --- a/src/model.jl +++ b/src/model.jl @@ -2426,7 +2426,7 @@ end Reset the accumulators and evaluate `model` using `context`, returning `(retval, varinfo)`. The context belongs to this evaluation, not to the model. The same context is passed to -submodels and to [`tilde_assume!!`](@ref) for latent sites. Observations and tracked values +submodels and to [`tilde_assume!!`](@ref) for latent LHS variables. Observations and tracked values go directly to accumulators, independently of the context. Models marked with [`setthreadsafe`](@ref) use a `ThreadSafeVarInfo` during evaluation. From c7cd675d953aa5b341d503dabfc0f9c29ec49ffa Mon Sep 17 00:00:00 2001 From: Hong Ge Date: Thu, 1 Oct 2026 16:07:07 +0100 Subject: [PATCH 4/5] benchmarks: time a parent binding into an indexed submodel Assisted-by: Claude Code --- benchmarks/benchmarks.jl | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) diff --git a/benchmarks/benchmarks.jl b/benchmarks/benchmarks.jl index b0da0fa9a..871481639 100644 --- a/benchmarks/benchmarks.jl +++ b/benchmarks/benchmarks.jl @@ -372,7 +372,13 @@ function build_combinations(rng) end push!(models, ("Dynamic", dynamic())) push!(models, ("Submodel", parent(randn(rng)))) - push!(models, ("Indexed submodels 3k", indexed_submodels(randn(rng, 3_000)))) + indexed = indexed_submodels(randn(rng, 3_000)) + push!(models, ("Indexed submodels 3k", indexed)) + # A parent binding at a child's prefixed address. + push!( + models, + ("Indexed submodels conditioned", condition(indexed, @varname(x[1].obs) => 0.0)), + ) d = [1, 1, 1, 2, 2, 2] w = [1, 2, 3, 2, 1, 1] z = [1, 1, 2, 2, 1, 2] From d58dff183c4b3704b50750c736bf9c17f68650e9 Mon Sep 17 00:00:00 2001 From: Hong Ge Date: Thu, 1 Oct 2026 17:12:39 +0100 Subject: [PATCH 5/5] contexts: test prepared evaluator arguments with bindings Assisted-by: Codex --- test/model.jl | 16 ++++++++++++++++ 1 file changed, 16 insertions(+) diff --git a/test/model.jl b/test/model.jl index cd6c88423..54d216f9e 100644 --- a/test/model.jl +++ b/test/model.jl @@ -61,6 +61,22 @@ end const GDEMO_DEFAULT = DynamicPPL.TestUtils.demo_assume_observe_literal() @testset "model.jl" begin + @testset "prepared evaluator arguments apply effective bindings" begin + @model function prepared_bindings(observation, replaced; held=0.0) + observation ~ Normal() + replaced ~ Normal() + held ~ Normal() + return (; observation, replaced, held) + end + model = fix(condition(prepared_bindings(1.0, 2.0); replaced=3.0); held=4.0) + context = DynamicPPL.Context(Xoshiro(1), InitFromPrior(), UnlinkAll()) + args, kwargs = DynamicPPL.make_evaluate_args_and_kwargs(model, context, VarInfo()) + result, vi = model.f(args...; kwargs...) + @test result == (; observation=1.0, replaced=3.0, held=4.0) + @test getloglikelihood(vi) ≈ logpdf(Normal(), 1.0) + logpdf(Normal(), 3.0) + @test getlogprior(vi) == 0.0 + end + @testset "immutable metadata for arguments with LHS variables" begin @model argument_lhs(x) = x ~ Normal() @test isbitstype(typeof(DynamicPPL.Model{false}(identity, (;), (;))))