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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
23 changes: 21 additions & 2 deletions ext/DistributionsInferenceComposedDistributionsExt.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
30 changes: 26 additions & 4 deletions src/engine.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand All @@ -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;
Expand All @@ -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
Expand All @@ -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 "
Expand Down Expand Up @@ -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
Expand Down
2 changes: 1 addition & 1 deletion src/protocol.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
39 changes: 39 additions & 0 deletions test/composed_ext.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
22 changes: 22 additions & 0 deletions test/engine.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Loading