From 753f58a3016c1650e4d985eca1580779ed513eb6 Mon Sep 17 00:00:00 2001 From: seabbs-bot Date: Wed, 12 Aug 2026 12:55:31 +0200 Subject: [PATCH] perf(engine): convert a NamedTuple-record batch once at FitLogDensity construction Batch record scoring allocated a fresh Missing-admitting vector per record per evaluation, via ComposedDistributions' _named_value_vector inside logpdf(d, ::NamedTuple) -- directly in the ForwardDiff gradient loop for the shape rand(d, n) produces. Add a scored_data field to FitLogDensity, prepared once at distribution_to_logdensity construction via a new _prepare_scored_data/_prepare_default_loglik_data hook (a no-op unless overridden). The ComposedDistributions extension overrides it for Sequential/Parallel composers scoring an AbstractVector{<:NamedTuple} batch under the default loglik, converting each record once and holding the converted form for the life of the problem. observations(prob) still returns exactly what the caller passed in; the converted batch is held separately and only reaches loglik when loglik is the default reducer, so a caller-supplied loglik always sees the untouched data. Closes #115. --- ...utionsInferenceComposedDistributionsExt.jl | 23 ++++++++++- src/engine.jl | 30 ++++++++++++-- src/protocol.jl | 2 +- test/composed_ext.jl | 39 +++++++++++++++++++ test/engine.jl | 22 +++++++++++ 5 files changed, 109 insertions(+), 7 deletions(-) diff --git a/ext/DistributionsInferenceComposedDistributionsExt.jl b/ext/DistributionsInferenceComposedDistributionsExt.jl index b136396..edae9ac 100644 --- a/ext/DistributionsInferenceComposedDistributionsExt.jl +++ b/ext/DistributionsInferenceComposedDistributionsExt.jl @@ -12,9 +12,11 @@ # belongs upstream (CD#189), not here. module DistributionsInferenceComposedDistributionsExt -using ComposedDistributions: ComposedDistributions, AbstractComposedDistribution +using ComposedDistributions: ComposedDistributions, AbstractComposedDistribution, + Sequential, Parallel import DistributionsInference: parameter_rows, estimated_rows, flat_dimension, - reconstruct, extra_logprior, extra_prior_state + reconstruct, extra_logprior, extra_prior_state, + _prepare_default_loglik_data # A centred pooled row's prior is a `CentredPoolPrior` marker, not a fixed # distribution: it is scored against the reconstructed hyperparameters, so it @@ -81,4 +83,21 @@ function extra_logprior(d::AbstractComposedDistribution, ::Any, return ComposedDistributions.pool_centred_logprior(state, nt) end +# A `NamedTuple`-record batch (e.g. `[rand(d) for _ in 1:n]`) is matched to +# `d`'s value-name layout by `ComposedDistributions._named_value_vector` on +# every `logpdf(d, ::NamedTuple)` call, heap-allocating a fresh +# `Missing`-admitting vector per record per evaluation — directly in the +# gradient loop under the default `loglik`. Since `reconstruct` only ever +# changes `d`'s numeric fields, not its tree structure, the by-name match is +# invariant across the whole fit: do it once here, at +# `distribution_to_logdensity` construction, and score the converted vectors +# directly thereafter. The `Missing`-admitting element type is preserved +# exactly, so the same censored `logpdf` specialisation +# (`AbstractVector{>:Missing}`) is still selected (#115). +function _prepare_default_loglik_data( + d::Union{Sequential, Parallel}, data::AbstractVector{<:NamedTuple}) + return [ComposedDistributions._named_value_vector(d, record) + for record in data] +end + end # module DistributionsInferenceComposedDistributionsExt diff --git a/src/engine.jl b/src/engine.jl index 1a786fb..e46ff51 100644 --- a/src/engine.jl +++ b/src/engine.jl @@ -5,6 +5,20 @@ _default_loglik(obj, data) = sum(record -> Distributions.logpdf(obj, record), data) +# The batch a `FitLogDensity` actually scores `loglik` against, prepared once +# at construction rather than re-derived on every `logdensity` call (#115). +# Generic no-op: an extension overrides this for its own object/record shape +# where a per-evaluation conversion would otherwise sit in the gradient loop +# (e.g. `ComposedDistributions`' `NamedTuple`-record batches). Only the +# *default* `loglik` gets a prepared batch: a caller-supplied `loglik` may +# read `data` in ways a converted form would silently break, so it always +# sees the untouched `data` it was handed at construction. +_prepare_scored_data(obj, data, loglik) = data +function _prepare_scored_data(obj, data, ::typeof(_default_loglik)) + _prepare_default_loglik_data(obj, data) +end +_prepare_default_loglik_data(obj, data) = data + @doc " A PPL-neutral log-density over a fit-protocol object's estimated parameters. @@ -24,6 +38,12 @@ directly, so it is sampleable by any LogDensityProblems consumer. - `data`: the observed records scored by `loglik`. - `loglik`: a reducer `(obj, data) -> Real` (default sums `logpdf(obj, record)`). +- `scored_data`: the batch `loglik` is actually called against, prepared once + from `data` at construction (`_prepare_scored_data`, #115). Equal to `data` + unless `loglik` is the default reducer and an extension has registered a + cheaper per-evaluation form for `obj`'s and `data`'s shape (e.g. + `ComposedDistributions` converting a `NamedTuple`-record batch to its + `Missing`-admitting vector form once, rather than on every evaluation). - `flat_priors`: the estimated rows' priors, in [`parameter_rows`](@ref) order, collected once at construction. An entry is `nothing` for an estimated row scored instead through [`extra_logprior`](@ref) (an object-dependent prior; @@ -40,10 +60,11 @@ directly, so it is sampleable by any LogDensityProblems consumer. - [`distribution_to_logdensity`](@ref): the assembler. - [`logdensity`](@ref): evaluate on a flat vector. " -struct FitLogDensity{D, T, L, FP, ES, CF} +struct FitLogDensity{D, T, L, SD, FP, ES, CF} obj::D data::T loglik::L + scored_data::SD flat_priors::FP extra_state::ES concrete_fields::CF @@ -54,8 +75,9 @@ function FitLogDensity(obj, data, loglik) flat_priors = [row.prior for row in rows] extra_state = extra_prior_state(obj) concrete_fields = _concrete_field_candidates(typeof(obj), rows) - return FitLogDensity( - obj, data, loglik, flat_priors, extra_state, concrete_fields) + scored_data = _prepare_scored_data(obj, data, loglik) + return FitLogDensity(obj, data, loglik, scored_data, flat_priors, + extra_state, concrete_fields) end @doc " @@ -357,7 +379,7 @@ function logdensity(prob::FitLogDensity, x::AbstractVector) _check_generic_fields(typeof(prob.obj), prob.concrete_fields, x) obj = reconstruct(prob.obj, x) lp += extra_logprior(prob.obj, obj, x, prob.extra_state) - return lp + prob.loglik(obj, prob.data) + return lp + prob.loglik(obj, prob.scored_data) end # A `nothing` prior (a fixed row, or an estimated one scored through diff --git a/src/protocol.jl b/src/protocol.jl index 9884382..d66fcd4 100644 --- a/src/protocol.jl +++ b/src/protocol.jl @@ -134,7 +134,7 @@ at its value in `obj`. `x` is [`flat_dimension`](@ref)`(obj)` long — empty whe This is the companion hook every fittable object implements with its own method, alongside [`parameter_rows`](@ref). The engine's -[`logdensity`](@ref) calls it once per evaluation to score `prob.data` +[`logdensity`](@ref) calls it once per evaluation to score the observations against the object collapsed at `x`. An estimated field's type must stay generic (e.g. `shape::S`, not diff --git a/test/composed_ext.jl b/test/composed_ext.jl index d44b4c2..4993527 100644 --- a/test/composed_ext.jl +++ b/test/composed_ext.jl @@ -93,6 +93,45 @@ end end end +@testitem "NamedTuple-record batch: cached scoring matches uncached, missing preserved (#115)" setup=[ + ComposedFixture] begin + using ComposedDistributions: _named_value_vector + + # A record with a `missing` field: the censored `AbstractVector{>:Missing}` + # `logpdf` specialisation must still be selected after the batch is + # converted once at construction, not the plain `AbstractVector` one. + data = [(onset_admit = 1.5, admit_death = missing), + (onset_admit = 2.0, admit_death = 3.0)] + prob = DistributionsInference.distribution_to_logdensity(plain_tree, data) + + # `observations` still returns exactly what was passed in — the + # conversion is held internally, not exposed through the public accessor. + @test DistributionsInference.observations(prob) === data + + # The internally-cached scoring batch is the by-name-matched, + # `Missing`-admitting vector form, matching CD's own per-record + # conversion exactly. + @test prob.scored_data isa Vector{Vector{Union{Missing, Float64}}} + for i in eachindex(data) + @test isequal(prob.scored_data[i], _named_value_vector(plain_tree, data[i])) + end + + # The method actually selected for the cached vector is the censored + # specialisation, not a plain-`AbstractVector` fallback. + m = which(Distributions.logpdf, (typeof(plain_tree), typeof(prob.scored_data[1]))) + @test occursin("Missing", string(m.sig)) + + # End-to-end: the cached-batch `logdensity` matches a hand-rolled + # uncached NamedTuple walk exactly, for a missing-admitting batch. + x = [2.5] + reconstructed = DistributionsInference.reconstruct(plain_tree, x) + rows = DistributionsInference.estimated_rows(plain_tree) + lp_prior = sum(logpdf(rows[i].prior, x[i]) for i in eachindex(x)) + lp_uncached = lp_prior + + sum(r -> Distributions.logpdf(reconstructed, r), data) + @test DistributionsInference.logdensity(prob, x) ≈ lp_uncached rtol=1e-12 +end + @testitem "logdensity matches the CD reference for a plain tree" setup=[ComposedFixture] begin data = [[0.5, 2.0], [1.0, 3.0]] prob = DistributionsInference.distribution_to_logdensity(plain_tree, data) diff --git a/test/engine.jl b/test/engine.jl index 67a036d..456dc9f 100644 --- a/test/engine.jl +++ b/test/engine.jl @@ -223,3 +223,25 @@ end @test DistributionsInference.flat_priors(prob) == [leaf.shape_prior] @test DistributionsInference.flat_priors(prob) === prob.flat_priors end + +@testitem "FitLogDensity: scored_data defaults to data, custom loglik sees the untouched data (#115)" setup=[ + ToyFixture] begin + leaf = ToyGammaLeaf(2.0, 1.0, LogNormal(log(2.0), 0.2)) + data = [1.5, 2.0, 3.2] + + # No batch-conversion hook registered for `ToyGammaLeaf`: the internal + # scoring batch is exactly `data`, not a copy. + default_prob = DistributionsInference.distribution_to_logdensity(leaf, data) + @test default_prob.scored_data === data + + # A caller-supplied `loglik` always receives the untouched `data`, even + # though a `_prepare_scored_data` hook could in principle exist for this + # object/data shape: only the default reducer's batch may be swapped. + seen = Ref{Any}(nothing) + capture_loglik(obj, d) = (seen[] = d; sum(y -> logpdf(obj, y), d)) + custom_prob = DistributionsInference.distribution_to_logdensity( + leaf, data; loglik = capture_loglik) + @test custom_prob.scored_data === data + DistributionsInference.logdensity(custom_prob, [2.5]) + @test seen[] === data +end