diff --git a/Project.toml b/Project.toml index b79473df8..f66d6d85d 100644 --- a/Project.toml +++ b/Project.toml @@ -75,6 +75,7 @@ MGVI = "fdae7790-d271-4276-880d-f72bbddf129c" NestedSamplers = "41ceaf6f-1696-4a54-9b49-2e7a9ec3782e" Optim = "429524aa-4258-5aef-a3af-852621145aeb" OptimizationBase = "bca83a33-5cc9-4baa-983d-23429ab6bcbb" +OptimizationLBFGSB = "22f7324a-a79d-40f2-bebe-3af60c77bd15" Plots = "91a5bcdd-55d7-5caf-9e0b-520d859cae80" PropertyFunctions = "09e99361-2bb8-48a2-a80f-de58f0739eb4" PythonCall = "6099a3de-0909-46bc-b1f4-468b9a2dfc0d" @@ -91,6 +92,7 @@ BATMGVIExt = "MGVI" BATNestedSamplersExt = "NestedSamplers" BATOptimExt = "Optim" BATOptimizationBaseExt = ["OptimizationBase"] +BATOptimizationLBFGSBExt = "OptimizationLBFGSB" BATPlotsExt = "Plots" BATPropertyFunctionsExt = "PropertyFunctions" BATSliceSamplingExt = "SliceSampling" @@ -147,6 +149,7 @@ NestedSamplers = "0.8, 0.9" OneTwoMany = "0.1.2" Optim = "1.12, 2" OptimizationBase = "3.3, 4, 5" +OptimizationLBFGSB = "1.5" PDMats = "0.9, 0.10, 0.11" ParallelProcessingTools = "0.4" Parameters = "0.12, 0.13" diff --git a/docs/make.jl b/docs/make.jl index 6b87ae207..1ce3cda8a 100644 --- a/docs/make.jl +++ b/docs/make.jl @@ -42,6 +42,7 @@ makedocs( "Home" => "index.md", "Installation" => "installation.md", "List of algorithms" => "list_of_algorithms.md", + "Molewhacker importance sampling" => "molewhacker.md", "Tutorial" => "tutorial.md", "API Documentation" => "stable_api.md", "Plotting" => "plotting.md", diff --git a/docs/src/experimental_api.md b/docs/src/experimental_api.md index 67bc422fc..664422ca5 100644 --- a/docs/src/experimental_api.md +++ b/docs/src/experimental_api.md @@ -37,6 +37,8 @@ EllipticalSliceMCMCSampling GridSampler HierarchicalDistribution PriorImportanceSampler +MolewhackerRefit +MolewhackerSampling ReactiveNestedSampling SliceMCMCSampling SobolSampler diff --git a/docs/src/internal_api.md b/docs/src/internal_api.md index 80773f621..f675f4e32 100644 --- a/docs/src/internal_api.md +++ b/docs/src/internal_api.md @@ -60,6 +60,7 @@ BAT.LogDVal BAT.MCMCSampleGenerator BAT.MCMCStepInfo BAT.MeasureLike +BAT.MultiThreadedExec BAT.NoWhitening BAT.OnlineMvCov BAT.OnlineMvMean diff --git a/docs/src/list_of_algorithms.md b/docs/src/list_of_algorithms.md index 0c034f0f0..c172865ea 100644 --- a/docs/src/list_of_algorithms.md +++ b/docs/src/list_of_algorithms.md @@ -112,6 +112,22 @@ bat_sample(target, PriorImportanceSampler(nsamples=10^5)) ``` +## Molewhacker importance sampler (experimental) + +BAT sampling algorithm type: [`MolewhackerSampling`](@ref) + +```julia +import ForwardDiff, OptimizationLBFGSB +context = BATContext(ad = ForwardDiff) +bat_sample(target, MolewhackerSampling(nsamples = 10^4), context) +``` + +Fits a defensive Gaussian mixture with local Fisher geometry. Returns fresh +importance samples from the fitted proposal. No MGVI dependency is required. +See [Molewhacker importance sampling](molewhacker.md) for supported models, +the sampling law, tuning, and diagnostics. + + ## Integration algorithms BAT function: [`bat_integrate`](@ref) diff --git a/docs/src/molewhacker.md b/docs/src/molewhacker.md new file mode 100644 index 000000000..289129ad2 --- /dev/null +++ b/docs/src/molewhacker.md @@ -0,0 +1,266 @@ +# Molewhacker importance sampling + +[`MolewhackerSampling`](@ref) builds a Gaussian mixture with local Fisher geometry +and draws fresh importance samples. The sampler is standalone within BAT. +Its API is experimental. + +```julia +using BAT, Distributions, StableRNGs +using MeasureBase: Likelihood +import ForwardDiff, OptimizationLBFGSB + +posterior = PosteriorMeasure( + Likelihood(z -> Normal(z[1], 0.5), 2.0), + MvNormal([0.0], [1.0;;]), +) +context = BATContext(rng = StableRNG(71), ad = ForwardDiff) +algorithm = MolewhackerSampling(nsamples = 4000, batchsize = 512, maxiter = 5) +result = evalmeasure(posterior, algorithm, context) +samples = BAT.samplesof(result) +diagnostics = result.evalinfo.result +``` + +This posterior has mean `1.6` and variance `0.2`. `bat_sample` also accepts the +algorithm. Use `evalmeasure` to retain the fitted proposal and diagnostics. + +## Proposal construction + +The proposal follows the [Newtrinos Molewhacker algorithm](https://github.com/Newtrinos-org/Newtrinos.jl/blob/bat-v5-migration/src/analysis/molewhacker.jl), +with Laplace seeds and fresh rounds added by default: + +1. Transform the prior to standard-normal coordinates. By default, run parallel + L-BFGS-B searches from ten Sobol starts to find initial centers. With Laplace + seeds, each search stops after 50 iterations and a Newton step finishes it. +2. At each center, form a Gaussian with precision `I + J' F J`. Here `J` is the + forward model's parameter Jacobian, and `F` is its observation distribution's + Fisher information. The identity adds prior information once. Laplace seeds + add a second Gaussian with the observed information as precision. +3. Let `q_uniform` be the equally weighted mixture of all local Gaussians. + Assign component masses proportional to `target(center) / q_uniform(center)`. +4. Draw an initial discovery pool. Rank its points by their current + `logtarget - logproposal` values. Add a Fisher Gaussian at each selected point. + These added centers need not be modes. +5. Recompute all component masses using the center-ratio rule. From each new + component, draw `floor(component_mass * previous_pool_size)` discovery points. + Repeat until a hard budget or an explicit pool-ESS or efficiency threshold applies. +6. Run three more rounds by default, each after adding fresh draws from the + current proposal to the pool. This step is an addition to the source algorithm. +7. Refit the proposal to those fresh draws by importance-weighted EM, keeping the + adaptive mixture as a defensive share. This step is also an addition. + +The adaptive pool guides proposal construction only. Its points have different +sampling laws, so reweighting the pool by the latest proposal does not produce +valid final importance weights. Its reported `pilot_ess` is a fitting heuristic. + +## Final sampling law + +By default, the final proposal is the fitted mixture. An explicit positive +`exploration_mass` adds a prior component after fitting: +`q_final = exploration_mass * prior + (1 - exploration_mass) * q_fitted`. +With bounded likelihood and positive prior mass, this bounds importance ratios. It does not guarantee +mode discovery. The prior component enters after fitting, preserving the source +algorithm's discovery proposals. + +Freeze both the proposal and production count before drawing the final samples. +Conditional on all earlier work, these samples are IID from `q_final`. +Only these fresh samples enter the returned empirical measure. Their log +importance ratios are `logtarget - logproposal`. + +Weights use one common exponential scale for numerical stability. +`diagnostics.logweight_scale` retains this scale. The target mass stays unchanged. +The normalized proposal appears in `result.approx`. Self-normalized estimates +retain their usual finite-sample bias. + +## Controls and limits + +- `nseeds` sets the number of initial centers. The default `init = nothing` uses + Sobol starts in normal coordinates. `init = ExplicitInit(...)` supplies centers + in original coordinates, which BAT copies and transforms. +- `init_mode` sets the initial optimizer. The default requires + `import OptimizationLBFGSB`. It stops after 50 iterations with Laplace seeds, + and after 1,000 without them. Set `init_mode = nothing` to keep supplied centers, + or `nseeds = 0` to skip initialization. Prior-only discovery can fail badly on + concentrated targets in higher dimensions. +- Mode searches receive equal shares of the available fitting-call budget, with + remainder calls assigned in seed order. Each search owns its optimizer copy + and RNG. A search that exhausts its share supplies no Gaussian. Unused calls + remain available for adaptation. +- `ncandidates` sets the number of centers selected per round. It defaults to 14, + independent of the thread count, so results do not depend on the machine. + Geometry calculations use the selected `executor`. +- `executor = BAT.MultiThreadedExec(ntasks = 14)` limits concurrent BAT work to + fourteen tasks, including target evaluation, geometry, and mixture scoring. + The default `MultiThreadedExec()` uses `Threads.nthreads()` tasks. This setting + does not change `ncandidates`, the Julia thread pool, or BLAS threads. + Allocation-heavy likelihoods can run faster with fewer active tasks. +- `batchsize` sets the initial discovery-pool size and the independent sizing-pilot + size. Later discovery batches follow component masses, so they can be empty. + Target values at existing pool points are reused. + Gaussians are cached by pool index, preserving every selection's mass update. + Center density sums add only the selected occurrences each round. + Mixture scoring uses the selected executor and bounded, reusable workspaces. + A bounded cache retains component densities across rounds. Components beyond + the cache budget are scored separately, preserving the cached work. +- `maxiter` is a strict limit on adaptation rounds. `maxiter = 0` skips discovery, + adaptation, and fresh rounds. `maxcomponents` limits proposed occurrences, + including repeated selections, refit Gaussians, and any added prior component. +- `fresh_rounds` defaults to three. After adaptation stops, for any reason, each + fresh round adds `batchsize` fresh draws from the current proposal to the pool, + then selects and adds candidates as before. Discovery batches come only from + new components, so the pool holds few draws from the current proposal. It then + misses the ratio spikes that production draws hit, and one weight can dominate + the output. Each fresh round costs `batchsize` target calls. Fresh rounds stop + early at `maxevals`, `maxcomponents`, or when no candidate has finite weight. + On the public DeepCore model, three fresh rounds cut the largest normalized + squared weight from 0.15–0.58 to 0.05–0.11 over four seeds. On well-fitted + targets they cost calls and gain nothing. Set `fresh_rounds = 0` for the source + algorithm's rounds only. +- The fresh-round draws then refit the proposal, in the style of MitISEM + (Hoogerheide, Opschoor and van Dijk). Each draw is weighted by the proposal that + drew it, and importance-weighted EM fits one to six Gaussians to the target. The + number maximizes the weighted log-likelihood of held-out draws, the cross-entropy part + of KL(p || fit), and the refit on all draws starts from the held-out winner. Each fitted + covariance is inflated by 1.1 to reduce weight concentration. This does not + guarantee bounded weights: the fitted covariance can still underestimate target spread. + On the test targets this cost about 6% ESS and cut the mean Pareto shape + from 0.15–0.32 to −0.02–0.25. The final proposal gives these 80% of the mass and + keeps the adaptive mixture at 20% for defence. Center ratios see the target only at + component centers, so they cannot see proposal mass placed where the target is + small. The fit needs no extra target calls and is skipped when the fresh draws have + too few effective samples or no component budget remains. The fit uses at most + the remaining component budget. On six + 12-dimensional test targets it raised production ESS by 13–210%, and on the public + DeepCore model by 13–36% over two seeds. Production draws come after the fit, from + the frozen result, so they stay IID from one proposal. Pass a `MolewhackerRefit` as + `refit` to tune the fit, or `nothing` to keep the adaptive mixture. +- `exploration_mass` defaults to zero. A positive value mixes the prior into the + final proposal. This option leaves discovery unchanged. +- `laplace_seeds` defaults to `true`. It adds a second Gaussian at each seed, with + the observed information `-∇² logtarget` as precision, variance inflated by + `laplace_inflation` (default 1.2). + Fisher information misses curvature where the forward model is stationary in a + parameter, for example a mixing angle near maximal mixing. The Hessian comes + from central differences of AD gradients, once per distinct seed. Where it is + positive definite, one Newton step polishes the seed if the target increases + and a target call remains after reserving center, discovery, and output calls. + This step lets the default mode search stop early. Elsewhere, such as at kinks, + the seed keeps its Fisher Gaussian alone. The two Gaussians share a center, so + center-ratio fitting gives them equal mass. The target must support AD + gradients, as for the default mode search. Laplace Gaussians count toward + `maxcomponents`, so it must allow twice `nseeds`. Gradient calls for these + Hessians do not count toward `maxevals`. Set `laplace_seeds = false` for the + source algorithm's Fisher-only seeds. +- `maxevals` caps target calls, including mode searches, initial centers, + discovery, sizing, and production. Geometry calls are separate and counted + in `ngeometries`. The default adds no call limit beyond the other stopping + rules. Explicit caps reserve production and any sizing pilot, and must leave + room for enabled initialization and discovery. With automatic output, each new + discovery draw also reserves one fresh output draw. +- `nsamples = nothing` is the default. Without an ESS goal, the output count + matches the final discovery pool, or `batchsize` when adaptation is disabled. + An explicit integer fixes the output count. +- `target_pool_ess` stops adaptation when the recycled-pool ESS exceeds its value. + `target_efficiency` stops adaptation when that ESS divided by the pool size + exceeds its value. Both default to `Inf`, leaving adaptation to the hard budgets. + When both are set, either threshold can stop adaptation. +- A finite `target_ess` sizes production only. An independent fresh pilot estimates the output count. + An explicit `nsamples` caps that count. With `nsamples = nothing`, only the + remaining `maxevals` budget caps it. The achieved ESS is not guaranteed. + Production never stops based on its current weights. + +Choose adaptation thresholds separately from the production goal: + +| Adaptation rule | Setting | +| --- | --- | +| Hard budgets only (default) | Leave both thresholds at `Inf` | +| Smaller fixed refinement budget | `maxiter = 16, ncandidates = 14` | +| Source pool-ESS heuristic | `target_pool_ess = 5000` | +| Source efficiency heuristic | `target_efficiency = 0.2` | +| Projected production ESS | `target_efficiency = target_ess / nsamples` | + +A smaller `maxiter` trades refinement for a smaller mixture, independently of +the production sample count. With ten initial components, fourteen candidates, +and sixteen completed rounds, the mixture has 234 proposed occurrences. Laplace +seeds can add up to ten more, and the default fresh rounds up to 42. Reselection +can store fewer Gaussians. Failed geometries or earlier stops can reduce both +counts. The refit can add up to six Gaussians within `maxcomponents`. +Optional prior mixing can add one component. +This is a user-selected budget, not an automatic convergence test. Compare fresh +weighted estimates of the observables you need before reducing the budget. + +The projected rule requires an explicit output cap and a goal below that cap. +It extrapolates from recycled-pool efficiency and can stop before finding tails or modes. +These thresholds do not validate the proposal or guarantee the production ESS. +Earlier versions coupled `target_ess` to pool-ESS stopping. Set +`target_pool_ess = target_ess` explicitly to retain that behavior. + +Target draws follow the context RNG's serial order before parallel evaluation. +For deterministic optimizers without wall-time limits, changing the executor +preserves that order. The model must support the context's AD selector. + +The sampler supports dense CPU geometry for Normal, MvNormal, Poisson, +Exponential, and product observation models. Singular local geometry rejects +that candidate. If initialization supplies no usable proposal, sampling uses +the prior. Unsupported models and unrelated errors propagate. An arbitrary +log-density closure does not expose the required forward model. + +Product models share one parameter Jacobian. Diagonal and isotropic Normal +covariances use compact parameter charts. Whitened Jacobian rows form one Fisher +pullback. The default ForwardDiff path avoids a redundant primal model call. +Local Gaussians retain their precision factor, avoiding explicit inversion. + +Dense parameter geometry needs quadratic storage and cubic factorization work. +Each round can propose `ncandidates` components, as in the source algorithm. +Reselecting a discovery point reuses its stored Gaussian. Its multiplicity remains +in the fitting density, mixture masses, and per-occurrence draw allocation. +Combining identical mixture categories can change draws for a fixed RNG seed +while preserving the proposal distribution. +With ten seeds and 14 candidates, 1,000 rounds can propose 14,010 components while +storing fewer Gaussians. `maxcomponents` counts proposed occurrences, preserving +the refinement budget even when components are reused. +Raising `nsamples` alone does not enlarge the discovery pool. +Mixture evaluations still grow with component count. Include initialization, geometry, +adaptation, and production when comparing total cost. These heuristics do not +establish global coverage or a general convergence guarantee. + +## Diagnostics + +`evalinfo.result` records iteration, target-call, geometry, component, and output +counts. `ncomponents` counts stored Gaussians. `ncomponent_proposals` includes +repeated selections and any added prior component. Geometry counts exclude cache +hits. `nseed_evals` counts initialization +target calls. `nseed_exhausted` counts mode searches that reach their assigned +budget. `nhessians` counts seed Hessians computed for `laplace_seeds`. +`niterations` counts adaptation rounds and `nfresh` counts fresh rounds. `history` +records pool growth, both component counts, and whether the round was fresh. +`stop_reason` describes adaptation, not the fresh rounds that follow. + +`stop_reason` distinguishes `:maxiter`, `:maxcomponents`, `:maxevals`, +`:pilot_ess`, `:pool_efficiency`, `:no_finite_candidate`, and `:geometry_failure`. +`pilot_ess` describes the recycled discovery pool. Its efficiency is `pilot_ess / npilot`. +`pilot_efficiency` comes from +the independent sizing pilot, when enabled. `ess` and `efficiency` describe the +fresh production weights. + +ESS averages weight dispersion and can hide one dominant weight. `pareto_k` is the +generalized Pareto shape of the largest production weights, as in PSIS. Values +above 0.7 mark unreliable estimates, and values below 0.5 are good. It is `NaN` +for fewer than 21 finite weights, and `Inf` when the largest weights exceed the +rest beyond the floating-point range. A small `pareto_k` does not rule out one +dominant weight, so also check `max_weight`, the largest normalized weight. Its +inverse is the L∞ effective sample size of Martino, Elvira and Louzada (2017). +Delta-method ESS values for single observables fail in the same case: a dominant +draw sits at the weighted mean and hides its own variance. Check relevant +observables and repeat runs when missed modes matter. Zero-target draws receive zero weight. A production +batch with no finite positive target mass fails. + +`smooth_weights = true` replaces the largest production weights by expected +order statistics of that fit, capped at the largest raw weight (Pareto-smoothed +importance sampling). This lowers estimator variance and adds a small bias, +including for the target mass. `pareto_k` still describes the raw weights. +Smoothing cannot repair `pareto_k` above 0.7. + +`weight_diagnostic` accepts an optional function of final log importance ratios. +For example, callers using MGVI can pass `MGVI.pareto_diagnostic`. Its return +value appears in `diagnostics.diagnostic`. It cannot replace the sampler's raw +weights. Its own sample-size requirements still apply. diff --git a/ext/BATOptimizationLBFGSBExt.jl b/ext/BATOptimizationLBFGSBExt.jl new file mode 100644 index 000000000..7c0710dd6 --- /dev/null +++ b/ext/BATOptimizationLBFGSBExt.jl @@ -0,0 +1,18 @@ +# This file is a part of BAT.jl, licensed under the MIT License (MIT). + +module BATOptimizationLBFGSBExt + +using BAT +import OptimizationLBFGSB + +BAT.pkgext(::Val{:OptimizationLBFGSB}) = BAT.PackageExtension{:OptimizationLBFGSB}() +BAT.ext_default(::BAT.PackageExtension{:OptimizationLBFGSB}, ::Val{:LBFGSB_ALG}) = OptimizationLBFGSB.LBFGSB() + +# The Fortran backend requires Float64 input. The sampler converts fitted +# centers back to the context precision before forming proposal geometry. +function BAT._mw_mode(center::AbstractVector{Float32}, logtarget, + mode::BAT.OptimizationAlg{<:OptimizationLBFGSB.LBFGSB}, remaining, context) + return BAT._mw_mode(Float64.(center), logtarget, mode, remaining, context) +end + +end diff --git a/src/samplers/importance/molewhacker.jl b/src/samplers/importance/molewhacker.jl new file mode 100644 index 000000000..80815bc08 --- /dev/null +++ b/src/samplers/importance/molewhacker.jl @@ -0,0 +1,806 @@ +# This file is a part of BAT.jl, licensed under the MIT License (MIT). + +""" + MolewhackerRefit(; kwargs...) + +Importance-weighted EM refit of the [`MolewhackerSampling`](@ref) proposal to its +fresh-round draws, in the style of MitISEM (Hoogerheide, Opschoor and van Dijk). +The number of fitted Gaussians maximizes the weighted log-likelihood of held-out draws. + +Fields: + +$(TYPEDFIELDS) +""" +@with_kw struct MolewhackerRefit + "Largest number of fitted Gaussians." + maxcomponents::Int = 6 + "Mass share that the adaptive mixture keeps for defence." + defence::Float64 = 0.2 + "Covariance shrinkage toward the diagonal." + shrinkage::Float64 = 0.1 + "Covariance inflation of the fitted Gaussians to reduce weight concentration." + inflation::Float64 = 1.1 + "Effective draws needed per dimension and fitted Gaussian." + min_ess_per_dim::Float64 = 2.0 + "Fitted mass below which a Gaussian drops out." + min_mass::Float64 = 1e-3 + "Maximum EM iterations per fit." + maxiter::Int = 200 + "Stop EM when the weighted log-likelihood per draw rises by less than this." + tol::Float64 = 1e-6 +end +export MolewhackerRefit + +""" + MolewhackerSampling(; kwargs...) + +Adaptive Gaussian-mixture importance sampling with local Fisher geometry. +Finds initial modes, adds Gaussians at high target-to-proposal ratios, and +updates all component masses from center ratios. Requires a differentiable +forward-model likelihood and a standard-normal prior after `pretransform`. + +The adaptive pool guides proposal construction only. Final samples are fresh +IID draws from the frozen mixture, with optional prior mixing. +Weights are `exp(logtarget - logproposal - logweight_scale)`; the common +scale is retained in `evalinfo.result`. + +Fields: + +$(TYPEDFIELDS) +""" +@with_kw struct MolewhackerSampling{TR<:TransformIntent,IA,IM,E<:BATExecutor,D,RF<:Union{Nothing,MolewhackerRefit}} <: AbstractSamplingAlgorithm + pretransform::TR = NormalBased() + "Seed source in original coordinates, or `nothing` for Sobol starts in normal coordinates." + init::IA = nothing + "Number of initial seeds. Zero skips mode initialization." + nseeds::Int = 10 + "Add a Laplace Gaussian (observed information) beside each seed's Fisher Gaussian, and Newton-polish the seed." + laplace_seeds::Bool = true + "Laplace variance inflation. On the public DeepCore model, 1.2 beat both 1.0 and 1.5." + laplace_inflation::Float64 = 1.2 + "Seed optimizer. The default L-BFGS-B backend requires `import OptimizationLBFGSB`. It stops after 50 iterations with Laplace seeds, whose Newton step finishes the search." + init_mode::IM = nseeds == 0 ? nothing : OptimizationAlg(optalg = ext_default(pkgext(Val(:OptimizationLBFGSB)), Val(:LBFGSB_ALG)), maxiters = laplace_seeds ? 50 : 1_000) + "Production count, or its cap with finite `target_ess`. Nothing follows the pool size or ESS goal." + nsamples::Union{Nothing,Int} = nothing + "Optional production ESS goal. A fresh pilot sizes production; achieved ESS is not guaranteed." + target_ess::Float64 = Inf + "Stop adaptation when recycled-pool ESS exceeds this heuristic threshold." + target_pool_ess::Float64 = Inf + "Stop adaptation when recycled-pool ESS per point exceeds this heuristic threshold." + target_efficiency::Float64 = Inf + "Initial discovery pool and production-sizing pilot count." + batchsize::Int = 1000 + maxiter::Int = 100 + "Maximum component proposals, including repeated selections, refit Gaussians, and any added prior component." + maxcomponents::Int = typemax(Int) + "Maximum target calls, excluding geometry. No additional call limit by default." + maxevals::Int = typemax(Int) + "Number of candidate centers added per round, independent of the thread count." + ncandidates::Int = 14 + "Optional prior coefficient added after fitting. Zero leaves the fitted proposal unchanged." + exploration_mass::Float64 = 0.0 + executor::E = default_executor() + "Optional function of final log importance ratios." + weight_diagnostic::D = nothing + "Pareto-smooth the largest production weights (PSIS). Lowers variance and adds a small bias." + smooth_weights::Bool = false + "Rounds after adaptation stops that first add `batchsize` fresh proposal draws to the pool." + fresh_rounds::Int = 3 + "Refit of the proposal to the fresh-round draws, or `nothing` to keep the adaptive mixture." + refit::RF = MolewhackerRefit() +end +export MolewhackerSampling + +function _mw_check(alg, context) + @argcheck (isnothing(alg.nsamples) || alg.nsamples > 0) && alg.batchsize > 0 + @argcheck alg.maxiter >= 0 && alg.nseeds >= 0 && alg.ncandidates > 0 && alg.fresh_rounds >= 0 + # Laplace seeds can double the initial occurrences. + @argcheck alg.maxcomponents >= max(1, (1 + alg.laplace_seeds) * alg.nseeds + Int(alg.exploration_mass > 0)) + @argcheck alg.target_ess > 0 && 0 <= alg.exploration_mass < 1 + @argcheck alg.target_pool_ess > 0 && alg.target_efficiency > 0 + @argcheck alg.target_efficiency <= 1 || alg.target_efficiency == Inf + pilot_count = isfinite(alg.target_ess) ? alg.batchsize : 0 + production_count = something(alg.nsamples, alg.batchsize) + @argcheck alg.maxevals >= production_count && alg.maxevals - production_count >= pilot_count + reserve = production_count + pilot_count + discovery_count = alg.maxiter > 0 ? alg.batchsize : 0 + center_count = alg.nseeds > 0 ? alg.nseeds : Int(alg.maxiter > 0) + @argcheck alg.maxevals - reserve >= discovery_count + center_count "Leave target calls for initial centers and discovery, or disable initialization and adaptation." + @argcheck isnothing(alg.init_mode) || alg.maxevals - reserve - discovery_count - center_count >= alg.nseeds "Leave target calls for mode searches, or disable the seed optimizer." + @argcheck get_compute_unit(context) isa CPUnit + @argcheck iszero(alg.exploration_mass) || 0 < get_precision(context)(alg.exploration_mass) < 1 + @argcheck alg.laplace_inflation > 0 + refit = alg.refit + @argcheck isnothing(refit) || (refit.maxcomponents > 0 && 0 <= refit.defence < 1 && 0 <= refit.shrinkage <= 1 && + refit.inflation > 0 && refit.min_ess_per_dim > 0 && 0 <= refit.min_mass < 1 && refit.maxiter > 0 && refit.tol >= 0) + return reserve +end + +function _mw_gaussian(center::AbstractVector{T}, precision) where T + P = PDMat{T}(precision) + c = copy(center) + return MvNormalCanon(c, P * c, P) +end + +# Idle tasks split Jacobian columns when fewer geometries than tasks are pending. +# Blocks keep at least three columns: narrower blocks added allocation but no speed. +_mw_jacobian_blocks(executor::MultiThreadedExec, n, dim) = clamp(fld(executor.ntasks, max(n, 1)), 1, cld(dim, 3)) +_mw_jacobian_blocks(::BATExecutor, n, dim) = 1 + +function _mw_draw(q, logtarget, n, executor, context) + v = VectorOfSimilarVectors(rand(get_rng(context), q, n)) + first_logp = logtarget(first(v)) + logp = Vector{typeof(float(first_logp))}(undef, n) + logp[1] = first_logp + exec_map!(logtarget, executor, view(logp, 2:n), view(v, 2:n)) + all(x -> isfinite(x) || x == -Inf, logp) || throw(ArgumentError("MolewhackerSampling encountered an invalid target log density.")) + logr = _mw_batched_logpdf(q, flatview(v), executor) + all(isfinite, logr) || throw(ArgumentError("MolewhackerSampling encountered a non-finite generating log density.")) + return (; v, logp, logr) +end + +function _mw_batched_logpdf(d, x::AbstractMatrix, executor = default_executor()) + T = promote_type(Distributions.partype(d), eltype(x)) + # Mixture logpdf! stores component densities in the mixture-weight type. + if d isa MixtureModel && T != eltype(probs(d)) + return logpdf.(Ref(d), eachcol(x)) + end + r = Vector{T}(undef, size(x, 2)) + logpdf!(r, d, x) + if d isa MixtureModel + # The batch kernel can produce NaN when every component returns -Inf. + for i in eachindex(r) + isnan(r[i]) && (r[i] = logpdf(d, view(x, :, i))) + end + end + return r +end + +function _mw_batched_logpdf(d::MixtureModel{Multivariate,Continuous,<:MvNormalCanon}, x::AbstractMatrix, + executor = default_executor()) + T = promote_type(eltype(mean(first(d.components))), eltype(x), eltype(probs(d))) + n = size(x, 2) + n == 0 && return T[] + ntasks = executor isa MultiThreadedExec ? executor.ntasks : Threads.nthreads() + nchunks = executor isa SequentialExec ? 1 : + min(ntasks, n, max(1, n * length(d.components) ÷ 65536)) + logweights = log.(probs(d)) + constants = Distributions.mvnormal_c0.(d.components) + nchunks == 1 && return _mw_gaussian_logpdf(d, x, logweights, constants) + # Keep matrix batches wide when the pool is small relative to the worker count. + by_component = n < 256nchunks + count = by_component ? length(d.components) : n + width = cld(count, nchunks) + ranges = [i:min(i + width - 1, count) for i in 1:width:count] + blocks = Vector{Vector{T}}(undef, length(ranges)) + score = r -> by_component ? _mw_component_logpdf(d, x, constants, r) : + _mw_gaussian_logpdf(d, view(x, :, r), logweights, constants) + exec_map!(score, executor, blocks, ranges) + if by_component + result = first(blocks) + for block in Iterators.drop(blocks, 1) + result .= _logaddexp.(result, block) + end + return result + end + return reduce(vcat, blocks) +end + +function _mw_component_logpdf(d, x, constants, indices) + weights = probs(d)[indices] + mass = sum(weights) + T = promote_type(eltype(mean(first(d.components))), eltype(x), eltype(weights)) + iszero(mass) && return fill(T(-Inf), size(x, 2)) + part = MixtureModel(d.components[indices], weights ./ mass) + result = _mw_gaussian_logpdf(part, x, log.(probs(part)), constants[indices]) + result .+= log(mass) + return result +end + +function _mw_gaussian_logpdf(d, x, logweights, constants) + T = promote_type(eltype(mean(first(d.components))), eltype(x), eltype(probs(d))) + n, k = size(x, 2), length(d.components) + # Bound the density workspace to about 2 MiB per task in Float64, and reuse it. + width = min(n, 512, max(1, 262144 ÷ k)) + terms = Matrix{T}(undef, width, k) + shifted = Matrix{T}(undef, size(x, 1), width) + maxima = Vector{T}(undef, width) + result = zeros(T, n) + for first in 1:width:n + indices = first:min(first + width - 1, n) + count = length(indices) + delta = view(shifted, :, 1:count) + fill!(maxima, T(-Inf)) + for i in eachindex(d.components) + iszero(probs(d)[i]) && continue + component = d.components[i] + delta .= view(x, :, indices) .- mean(component) + lmul!(cholesky(component.J).U, delta) + for j in 1:count + value = constants[i] - sum(abs2, view(delta, :, j)) / 2 + logweights[i] + terms[j, i] = value + maxima[j] = max(maxima[j], value) + end + end + r = view(result, indices) + for i in eachindex(d.components) + iszero(probs(d)[i]) && continue + for j in 1:count + r[j] += exp(terms[j, i] - maxima[j]) + end + end + for j in 1:count + r[j] = log(r[j]) + maxima[j] + isnan(r[j]) && (r[j] = logpdf(d, view(x, :, indices[j]))) + end + end + return result +end + +# Bound the cache entries. Components beyond this budget are scored without caching. +const _MW_SCORE_CACHE_LIMIT = 2^24 + +function _mw_logpdf_column(component, x) + delta = x .- mean(component) + lmul!(cholesky(component.J).U, delta) + return Distributions.mvnormal_c0(component) .- vec(sum(abs2, delta, dims = 1)) ./ 2 +end + +# Rounds change every mass but no existing component density, so only new pool +# points and new components need whitening. +function _mw_extend_scores(cache, components, points, executor) + N, K = size(points, 2), length(components) + size(cache) == (N, K) && return cache + n, k = min(size(cache, 1), N), min(size(cache, 2), K) + grown = similar(cache, N, K) + grown[1:n, 1:k] = view(cache, 1:n, 1:k) + columns = Vector{Vector{eltype(cache)}}(undef, K) + rows = view(points, :, n+1:N) + exec_map!(j -> _mw_logpdf_column(components[j], j <= k ? rows : points), executor, columns, collect(1:K)) + for j in 1:K + j <= k ? (grown[n+1:N, j] = columns[j]) : (grown[:, j] = columns[j]) + end + return grown +end + +function _mw_cached_logpdf_rows(cache, logmass, r) + T = eltype(cache) + maxima = fill(T(-Inf), length(r)) + @inbounds for k in eachindex(logmass) + lm, column = logmass[k], view(cache, r, k) + @simd for j in eachindex(maxima) + maxima[j] = max(maxima[j], column[j] + lm) + end + end + sums = zeros(T, length(r)) + @inbounds for k in eachindex(logmass) + lm, column = logmass[k], view(cache, r, k) + isfinite(lm) || continue + @simd for j in eachindex(sums) + sums[j] += ifelse(maxima[j] == -Inf, zero(T), exp(column[j] + lm - maxima[j])) + end + end + return log.(sums) .+ maxima +end + +function _mw_pool_logpdf(q, components, points, cache, executor, limit = _MW_SCORE_CACHE_LIMIT) + ncached = min(length(components), limit ÷ size(points, 2)) + if ncached == 0 + return _mw_batched_logpdf(q, points, executor), similar(cache, 0, 0) + end + cache = _mw_extend_scores(cache, view(components, 1:ncached), points, executor) + n = size(cache, 1) + ntasks = executor isa MultiThreadedExec ? executor.ntasks : 1 + width = cld(n, max(1, min(ntasks, n ÷ 256))) + ranges = [i:min(i + width - 1, n) for i in 1:width:n] + blocks = Vector{Vector{eltype(cache)}}(undef, length(ranges)) + logmass = log.(view(probs(q), 1:ncached)) + exec_map!(r -> _mw_cached_logpdf_rows(cache, logmass, r), executor, blocks, ranges) + result = reduce(vcat, blocks) + if ncached < length(components) + weights = probs(q)[ncached+1:end] + mass = sum(weights) + if mass > 0 + remainder = MixtureModel(components[ncached+1:end], weights ./ mass) + terms = _mw_batched_logpdf(remainder, points, executor) + result .= _logaddexp.(result, terms .+ log(mass)) + end + end + return result, cache +end + +struct MolewhackerBudgetReached <: Exception end + +function _mw_with_budget(f, logtarget, remaining) + ncalls = Ref(0) + counted = x -> begin + ChainRulesCore.ignore_derivatives() do + ncalls[] < remaining || throw(MolewhackerBudgetReached()) + ncalls[] += 1 + end + logtarget(x) + end + try + return f(counted), ncalls[], false + catch err + err isa MolewhackerBudgetReached || rethrow() + return nothing, ncalls[], true + end +end + +function _mw_mode(center, logtarget, mode, remaining, context) + isnothing(mode) && return center, 0, false + result, ncalls, exhausted = _mw_with_budget(logtarget, remaining) do counted + # Optimizer results do not infer. The assertion keeps later seed work concrete. + convert(typeof(center), maximize_density(counted, center, mode, context).result)::typeof(center) + end + return exhausted ? center : result, ncalls, exhausted +end + +# Mixture refinement (MolewhackerRefit): fit Gaussians to the target by importance-weighted EM on +# the fresh-round draws, each weighted by its own proposal, and keep the adaptive mixture as a +# defensive share. Center ratios cannot see proposal mass placed where the target is small. These +# draws can. The held-out weighted log-likelihood is the cross-entropy part of KL(p || fit). +function _mw_weighted_moments(x, w, shrinkage) + mu = x * w ./ sum(w) + c = x .- mu + S = (c .* w') * c' ./ sum(w) + # Shrink toward the diagonal: a few effective draws must fit many covariances. + return mu, Symmetric((1 - shrinkage) .* S .+ shrinkage .* Diagonal(diag(S))) +end + +# M step: weighted moments per component. Components with negligible mass drop out. +function _mw_em_masses(x, w, R, refit) + T = eltype(x) + mass = vec(sum(w .* R, dims = 1)) + keep = findall(>(T(refit.min_mass)), mass) + return [_mw_weighted_moments(x, w .* view(R, :, k), T(refit.shrinkage)) for k in keep], mass[keep] ./ sum(mass[keep]) +end + +# Gaussian log density from the Cholesky factor. MvNormal with a Symmetric covariance does not infer. +function _mw_normal_logpdf(mu, S, x) + U = cholesky(S).U + z = U' \ (x .- mu) + return .-vec(sum(abs2, z, dims = 1)) ./ 2 .- (logdet(U) + size(x, 1) * log(2 * eltype(x)(pi)) / 2) +end + +# E step: responsibilities of each component for each draw, and each draw's mixture log density. +function _mw_em_estep(x, fits, pis) + L = reduce(hcat, [log(p) .+ _mw_normal_logpdf(mu, S, x) for ((mu, S), p) in zip(fits, pis)]) + m = maximum(L, dims = 2) + R = exp.(L .- m) + s = sum(R, dims = 2) + return R ./ s, vec(m .+ log.(s)) +end + +# Weighted k-means++ starts, with hard assignment to the nearest start. +function _mw_em_start(x, w, K, rng) + T = eltype(x) + centers = [x[:, argmax(w)]] + while length(centers) < K + d2 = [minimum(c -> sum(abs2, view(x, :, i) .- c), centers) for i in axes(x, 2)] .* w + push!(centers, x[:, something(findfirst(>=(rand(rng, T) * sum(d2)), cumsum(d2)), size(x, 2))]) + end + R = zeros(T, size(x, 2), K) + for i in axes(x, 2) + R[i, argmin(k -> sum(abs2, view(x, :, i) .- centers[k]), 1:K)] = 1 + end + return R +end + +function _mw_em(x, w, R, refit) + T = eltype(x) + fits, pis = Tuple{Vector{T},Symmetric{T,Matrix{T}}}[], T[] + loglik = T(-Inf) + for _ in 1:refit.maxiter + fits, pis = _mw_em_masses(x, w, R, refit) + all(f -> isposdef(last(f)), fits) || return nothing + R, l = _mw_em_estep(x, fits, pis) + previous, loglik = loglik, dot(w, l) + loglik - previous < refit.tol && break + end + return fits, pis +end + +function _mw_fit_mixture(q, x, logw, refit, context, maxcomponents = refit.maxcomponents) + T = eltype(x) + w = exp.(logw .- maximum(logw)) + fit, held = x[:, 1:2:end], x[:, 2:2:end] + wfit, wheld = w[1:2:end] ./ sum(w[1:2:end]), w[2:2:end] ./ sum(w[2:2:end]) + best, score = nothing, T(-Inf) + for K in 1:min(refit.maxcomponents, maxcomponents) + 1 / sum(abs2, wfit) > refit.min_ess_per_dim * K * size(x, 1) || break + em = _mw_em(fit, wfit, _mw_em_start(fit, wfit, K, get_rng(context)), refit) + isnothing(em) && continue + s = dot(wheld, last(_mw_em_estep(held, em...))) + s > score && ((best, score) = (em, s)) + end + isnothing(best) && return q + # Refit on all draws from the held-out winner. A new start can fall into a worse optimum. + w ./= sum(w) + em = _mw_em(x, w, first(_mw_em_estep(x, best...)), refit) + isnothing(em) && return q + fits, pis = em + fitted = [_mw_gaussian(mu, Matrix(Symmetric(inv(cholesky(S)))) ./ T(refit.inflation)) for (mu, S) in fits] + defence = T(refit.defence) + masses = vcat(defence .* probs(q), (1 - defence) .* pis) + return MixtureModel(vcat(q.components, fitted), masses ./ sum(masses)) +end + +function _mw_efficiency(logw) + c = maximum(logw) + isfinite(c) || throw(ArgumentError("MolewhackerSampling drew no finite positive target mass.")) + w = exp.(logw .- c) + ess = sum(w)^2 / sum(abs2, w) + return (; weight = w, ess, efficiency = ess / length(w), logweight_scale = c) +end + +# Generalized Pareto fit to the largest importance ratios, relative to the ratio at the tail +# cutoff (Zhang & Stephens 2009). The shape gets the weakly informative adjustment of PSIS. +function _mw_pareto_fit(logw) + finite = findall(isfinite, logw) + order = finite[sortperm(logw[finite])] + n = length(order) + ntail = min(ceil(Int, n / 5), ceil(Int, 3sqrt(n))) + n > ntail >= 5 || return nothing + tail = order[n-ntail+1:n] + # Shift by the largest log ratio so exceedances stay in [0, 1] for any weight spread. + logmax = float(logw[last(tail)]) + base = exp(logw[order[n-ntail]] - logmax) + x = exp.(logw[tail] .- logmax) .- base + last(x) > 0 || return nothing + m = 30 + floor(Int, sqrt(ntail)) + xstar = x[max(1, floor(Int, ntail / 4 + 1 / 2))] + # Tail spread beyond the Float64 range: the largest ratios dominate completely. + xstar > 0 || return (; k = oftype(logmax, Inf), sigma = oftype(logmax, NaN), tail, logmax, base) + θ = [1 / last(x) + (1 - sqrt(m / (j - 1 / 2))) / (3xstar) for j in 1:m] + ks = [mean(log1p.(-t .* x)) for t in θ] + loglik = ntail .* (log.(-θ ./ ks) .- ks .- 1) + weights = exp.(loglik .- maximum(loglik)) + θhat = sum(θ .* weights) / sum(weights) + k = mean(log1p.(-θhat .* x)) + return (; k = (ntail * k + 5) / (ntail + 10), sigma = -k / θhat, tail, logmax, base) +end + +# ESS misses a single dominant weight; the tail shape does not. +_mw_pareto_k(logw) = (fit = _mw_pareto_fit(logw); isnothing(fit) ? NaN : fit.k) + +# PSIS: replace the largest ratios by expected GPD order statistics, capped at the largest raw ratio. +function _mw_smooth(logw) + fit = _mw_pareto_fit(logw) + (isnothing(fit) || !isfinite(fit.k)) && return logw + p =((1:length(fit.tail)) .- 1 / 2) ./ length(fit.tail) + q = abs(fit.k) < sqrt(eps()) ? -fit.sigma .* log1p.(-p) : fit.sigma .* expm1.(-fit.k .* log1p.(-p)) ./ fit.k + smoothed = copy(logw) + smoothed[fit.tail] .= fit.logmax .+ min.(log.(fit.base .+ q), 0) + return smoothed +end + +# BAT's weightedmeasure represents a likelihood scale as a log-function chain. +# Recover only these known constant shifts, never an arbitrary log-density model. +_mw_model(likelihood::Union{Likelihood,_SimpleLikelihood}) = _get_model(likelihood) +_mw_model(likelihood::DensityInterface.LogFuncDensity) = _mw_logmodel(likelihood._log_f) +_mw_logmodel(f::Base.Fix1{typeof(logdensityof)}) = _mw_model(f.x) +function _mw_logmodel(f::FunctionChain) + fs = fchainfs(f) + tail = last(fs) + if tail isa Base.Fix2{typeof(+),<:Real} && isfinite(tail.x) + return _mw_logmodel(ffchain(Base.front(fs)...)) + end + throw(ArgumentError("MolewhackerSampling requires a supported forward-model likelihood.")) +end +_mw_logmodel(f) = throw(ArgumentError("MolewhackerSampling requires a supported forward-model likelihood.")) + + +function _mw_center_mixture(components, center_logp, executor, previous = similar(center_logp, 0), + multiplicities = ones(Int, length(components)), added = length(previous)+1:length(components)) + T = eltype(first(components)) + n, nold = sum(multiplicities), length(previous) + uniform = MixtureModel(components, T.(multiplicities) ./ T(n)) + centers = reduce(hcat, mean.(components)) + # Existing center/component pairs never change. Add only the new densities + # to their unnormalized sums, retaining every component's multiplicity. + new_uniform = MixtureModel(components[added], fill(inv(T(length(added))), length(added))) + old_terms = _mw_batched_logpdf(new_uniform, view(centers, :, 1:nold), executor) .+ log(T(length(added))) + new_terms = _mw_batched_logpdf(uniform, view(centers, :, nold+1:length(components)), executor) .+ log(T(n)) + center_logsum = vcat(_logaddexp.(previous, old_terms), new_terms) + logw = center_logp .- (center_logsum .- log(T(n))) + offset = maximum(logw) + isfinite(offset) || return nothing, center_logsum + weights = T.(exp.(logw .- offset)) .* multiplicities + return MixtureModel(components, weights ./ sum(weights)), center_logsum +end + +function _mw_initial_proposal(transformed_m, f, model, logtarget, gprior, alg, fit_budget, ad, context) + T = get_precision(context) + components = typeof(gprior)[] + center_logp = T[] + nevals, ngeometries, nfailed, nexhausted, nhessians = 0, 0, 0, 0, 0 + if alg.nseeds > 0 + seeds = if isnothing(alg.init) + bat_sample(StandardMvNormal{T}(length(gprior)), SobolSampler(nsamples = alg.nseeds), context).result.v + else + bat_initval(transformed_m, alg.nseeds, apply_trafo_to_init(f, alg.init), context).result + end + starts = [T.(collect(seed)) for seed in seeds] + results = Vector{Tuple{Vector{T},Int,Bool}}(undef, alg.nseeds) + if isnothing(alg.init_mode) + results .= [(c, 0, false) for c in starts] + else + # Preserve center-density and discovery calls after the parallel searches. + mode_budget = fit_budget - alg.nseeds - (alg.maxiter > 0 ? alg.batchsize : 0) + share, extra = divrem(mode_budget, alg.nseeds) + contexts = [set_rng(context, Philox4x((rand(get_rng(context), UInt64), UInt64(i)))::Philox4x{UInt64,10}) for i in 1:alg.nseeds] + search = i -> _mw_mode(starts[i], logtarget, deepcopy(alg.init_mode), share + (i <= extra), contexts[i]) + exec_map!(search, alg.executor, results, collect(1:alg.nseeds)) + end + nevals = sum(r -> r[2], results) + nexhausted = count(r -> r[3], results) + centers = [r[1] for r in results if !r[3]] + precisions = Vector{Union{Nothing,typeof(gprior.J)}}(undef, length(centers)) + nblocks = _mw_jacobian_blocks(alg.executor, length(centers), length(gprior)) + exec_map!(c -> _mw_local_precision(model, c, ad, nblocks), alg.executor, precisions, centers) + ngeometries = length(centers) + keep = findall(!isnothing, precisions) + nfailed = ngeometries - length(keep) + seeds, fisher = centers[keep], precisions[keep] + observed = Vector{Union{Nothing,Matrix{T}}}(nothing, length(seeds)) + if alg.laplace_seeds && !isempty(seeds) + remaining = fit_budget - nevals - length(seeds) - (alg.maxiter > 0 ? alg.batchsize : 0) + seeds, observed, nhessians, ncalls = _mw_laplace_seeds(logtarget, seeds, fisher, ad, alg.executor, remaining) + nevals += ncalls + end + components = typeof(gprior)[_mw_gaussian(seeds[i], fisher[i]) for i in eachindex(seeds)] + center_logp = logtarget.(seeds) + nevals += length(seeds) + # A Laplace Gaussian shares its seed's center, so center-ratio fitting gives the pair equal mass. + for i in eachindex(seeds) + if !isnothing(observed[i]) + push!(components, _mw_gaussian(seeds[i], observed[i] ./ T(alg.laplace_inflation))) + push!(center_logp, center_logp[i]) + end + end + end + q, center_logsum = isempty(components) ? (nothing, similar(center_logp, 0)) : + _mw_center_mixture(components, center_logp, alg.executor) + if isnothing(q) + q = MixtureModel([gprior], [one(T)]) + center_logp = alg.maxiter > 0 ? [logtarget(mean(gprior))] : Float64[] + center_logsum = alg.maxiter > 0 ? [logpdf(gprior, mean(gprior))] : Float64[] + nevals += length(center_logp) + end + return q, center_logp, center_logsum, nevals, ngeometries, nfailed, nexhausted, nhessians +end + +# Laplace seeds: the observed information −∇²log p at each seed, from central differences of AD +# gradients run on the executor. Fisher information misses curvature where the forward model is +# stationary. Where the observed information is positive definite, one Newton step polishes the +# center if it raises the log target. Seeds within squared Mahalanobis 1 of an earlier seed share +# its Hessian: typical draws of a d-dimensional Gaussian lie near d, so this holds in any dimension. +function _mw_laplace_seeds(logtarget, centers, fisher, ad, executor, maxevals = typemax(Int)) + T = eltype(first(centers)) + d = length(first(centers)) + owner = collect(eachindex(centers)) + for j in eachindex(centers) + i = findfirst(i -> owner[i] == i && dot(centers[j] - centers[i], fisher[i] * (centers[j] - centers[i])) < 1, 1:j-1) + isnothing(i) || (owner[j] = i) + end + distinct = findall(j -> owner[j] == j, eachindex(centers)) + h = cbrt(eps(T)) + valgrad = valgrad_func(logtarget, ad) + # Per distinct seed: the gradient at the center, then at ±h along each coordinate. + offsets = [(0, zero(T)); [(k, σ * h) for k in 1:d for σ in (1, -1)]] + jobs = [(s, k, δ) for s in distinct for (k, δ) in offsets] + evaluated = Vector{Tuple{T,Vector{T}}}(undef, length(jobs)) + exec_map!(job -> _mw_shifted_valgrad(valgrad, centers[job[1]], job[2], job[3]), executor, evaluated, jobs) + observed = Vector{Union{Nothing,Matrix{T}}}(nothing, length(centers)) + polished = copy(centers) + ncalls = 0 + for (n, s) in enumerate(distinct) + block = evaluated[(n-1)*length(offsets)+1:n*length(offsets)] + H = reduce(hcat, [(last(block[2k]) .- last(block[2k+1])) ./ (2h) for k in 1:d]) + P = -Matrix(Symmetric((H + H') ./ 2)) + all(isfinite, P) && isposdef(Symmetric(P)) || continue + for j in findall(==(s), owner) + observed[j] = P + end + ncalls < maxevals || continue + candidate = centers[s] .+ P \ last(first(block)) + ncalls += 1 + logtarget(candidate) > first(first(block)) && (polished[s] = candidate) + end + return polished, observed, length(distinct), ncalls +end + +function _mw_shifted_valgrad(valgrad, center, k, δ) + x = copy(center) + k > 0 && (x[k] += δ) + value, gradient = valgrad(x) + return (value, collect(gradient)) +end + +function evalmeasure_impl(em::EvaluatedMeasure, alg::MolewhackerSampling, context::BATContext) + reserve = _mw_check(alg, context) + transformed_m, f = transform_and_unshape(alg.pretransform, em, context) + m = unevaluated(transformed_m) + is_std_mvnormal(getprior(m)) || throw(ArgumentError("MolewhackerSampling requires a standard-normal prior after pretransform.")) + model = ffcomp(_mw_model(getlikelihood(unevaluated(em))), inverse(f)) + logtarget = checked_logdensityof(m) + T = get_precision(context) + dim = some_dof(transformed_m) + gprior = _mw_gaussian(zeros(T, dim), Matrix{T}(I, dim, dim)) + ad = alg.maxiter > 0 || alg.nseeds > 0 ? get_valid_adselector(context, alg) : get_adselector(context) + fit_budget = alg.maxevals - reserve + q, center_logp, center_logsum, nevals, ngeometries, nfailed, nseed_exhausted, nhessians = + _mw_initial_proposal(transformed_m, f, model, logtarget, gprior, alg, fit_budget, ad, context) + nseed_evals = nevals + components = copy(q.components) + multiplicities = ones(Int, length(components)) + ncomponent_proposals = length(components) + has_prior = first(q.components) === gprior + component_limit = alg.maxcomponents - Int(!has_prior && alg.exploration_mass > 0) + # Automatic output reserves one fresh draw for each discovery point. + draw_cost = isnothing(alg.nsamples) ? 2 : 1 + niterations, nfresh, npilot = 0, 0, 0 + stop_reason = nseed_exhausted > 0 && nseed_exhausted == alg.nseeds ? :maxevals : :maxiter + history = StructArray((; iteration = Int[], fresh = Bool[], npilot = Int[], drawn = Int[], ncomponents = Int[], ncomponent_proposals = Int[])) + pilot_ess = nothing + # Pool and fresh draws grow in place across rounds. + fresh_points, fresh_logw = ElasticArray{T}(undef, dim, 0), T[] + + if alg.maxiter > 0 && fit_budget - nevals >= alg.batchsize + batch = _mw_draw(q, logtarget, alg.batchsize, alg.executor, context) + nevals += alg.batchsize + points, logp = ElasticArray{T}(flatview(batch.v)), batch.logp + npilot = length(logp) + # Zero marks unvisited points; -1 marks failed geometry. Positive entries + # identify stored Gaussians, so reselection needs no equality search. + component_index = zeros(Int, npilot) + score_cache = Matrix{T}(undef, 0, 0) + iteration, adapting = 0, true + while true + iteration += 1 + adapting &= iteration <= alg.maxiter + if !adapting + # The pool holds few draws from the current proposal, so it misses the ratio + # spikes that production would hit. Fresh draws expose them to selection, + # whatever stopped adaptation. + nfresh < alg.fresh_rounds && fit_budget - nevals >= alg.batchsize * draw_cost && + ncomponent_proposals < component_limit || break + nfresh += 1 + batch = _mw_draw(q, logtarget, alg.batchsize, alg.executor, context) + append!(fresh_points, flatview(batch.v)) + append!(fresh_logw, batch.logp .- batch.logr) + nevals += alg.batchsize + fit_budget -= (draw_cost - 1) * alg.batchsize + append!(points, flatview(batch.v)) + append!(logp, batch.logp) + append!(component_index, zeros(Int, alg.batchsize)) + npilot = length(logp) + end + logq, score_cache = _mw_pool_logpdf(q, components, points, score_cache, alg.executor) + scores = logp .- logq + reason = if !isfinite(maximum(scores)) + :no_finite_candidate + elseif !adapting + nothing + # Recycled scores guide fitting only; they never become output weights. + elseif (pilot_ess = _mw_efficiency(scores).ess) > alg.target_pool_ess + :pilot_ess + elseif pilot_ess / npilot > alg.target_efficiency + :pool_efficiency + elseif fit_budget - nevals < draw_cost + :maxevals + elseif ncomponent_proposals >= component_limit + :maxcomponents + end + if !isnothing(reason) + adapting || break + stop_reason, adapting = reason, false + continue + end + adapting && (niterations = iteration) + nselected = min(alg.ncandidates, npilot, component_limit - ncomponent_proposals) + indices = partialsortperm(scores, 1:nselected, rev = true) + uncached = filter(i -> iszero(component_index[i]), indices) + centers = collect.(eachcol(view(points, :, uncached))) + precisions = Vector{Union{Nothing,typeof(gprior.J)}}(undef, length(uncached)) + nblocks = _mw_jacobian_blocks(alg.executor, length(uncached), dim) + exec_map!(c -> _mw_local_precision(model, c, ad, nblocks), alg.executor, precisions, centers) + ngeometries += length(uncached) + nfailed += count(isnothing, precisions) + for (index, center, precision) in zip(uncached, centers, precisions) + if isnothing(precision) + component_index[index] = -1 + else + push!(components, _mw_gaussian(center, precision)) + push!(multiplicities, 0) + push!(center_logp, logp[index]) + component_index[index] = length(components) + end + end + added = filter(>(0), component_index[indices]) + if isempty(added) + adapting || break + stop_reason, adapting = :geometry_failure, false + continue + end + multiplicities[added] .+= 1 + ncomponent_proposals += length(added) + q, center_logsum = _mw_center_mixture(components, center_logp, alg.executor, + center_logsum, multiplicities, added) + # Preserve the request for each selected occurrence before flooring. + requested = floor.(Int, (probs(q)[added] ./ multiplicities[added]) .* npilot) + drawn = 0 + for (index, count) in zip(added, requested) + n = min(count, (fit_budget - nevals) ÷ draw_cost) + n == 0 && continue + batch = _mw_draw(components[index], logtarget, n, alg.executor, context) + nevals += n + fit_budget -= (draw_cost - 1) * n + drawn += n + append!(points, flatview(batch.v)) + append!(logp, batch.logp) + append!(component_index, zeros(Int, n)) + end + npilot = length(logp) + push!(history, (; iteration, fresh = !adapting, npilot, drawn, ncomponents = length(q.components), ncomponent_proposals)) + end + scores = logp .- first(_mw_pool_logpdf(q, components, points, score_cache, alg.executor)) + pilot_ess = isfinite(maximum(scores)) ? _mw_efficiency(scores).ess : zero(T) + elseif alg.maxiter > 0 + stop_reason = :maxevals + end + if !isnothing(alg.refit) && !isempty(fresh_logw) && ncomponent_proposals < component_limit + nprevious = length(q.components) + q = _mw_fit_mixture(q, fresh_points, fresh_logw, alg.refit, context, component_limit - ncomponent_proposals) + ncomponent_proposals += length(q.components) - nprevious + end + + # Optional prior mixing does not affect the source discovery strategy. + epsilon = T(alg.exploration_mass) + if epsilon > 0 + ncomponent_proposals += Int(!has_prior) + weights = (one(T) - epsilon) .* probs(q) + q = if has_prior + weights[1] += epsilon + MixtureModel(q.components, weights) + else + MixtureModel(vcat([gprior], q.components), vcat(epsilon, weights)) + end + end + nproduction = something(alg.nsamples, max(alg.batchsize, npilot)) + pilot_efficiency = nothing + if isfinite(alg.target_ess) + pilot = _mw_draw(q, logtarget, alg.batchsize, alg.executor, context) + nevals += alg.batchsize + pilot_efficiency = _mw_efficiency(pilot.logp .- pilot.logr).efficiency + limit = something(alg.nsamples, alg.maxevals - nevals) + requested = alg.target_ess / pilot_efficiency + nproduction = requested >= limit ? limit : ceil(Int, requested) + end + production = _mw_draw(q, logtarget, nproduction, alg.executor, context) + nevals += nproduction + logw = production.logp .- production.logr + output = _mw_efficiency(alg.smooth_weights ? _mw_smooth(logw) : logw) + diagnostic = isnothing(alg.weight_diagnostic) ? nothing : alg.weight_diagnostic(logw) + smpls_z = DensitySampleVector(v = production.v, logd = production.logp, weight = output.weight) + smpls = inverse(f).(smpls_z) + dsm = DensitySampleMeasure(smpls, dof = dim, ess = output.ess) + q_z = batmeasure(q) + q_original = pushfwd(inverse(f), q_z) + approx = alg.pretransform isa DoNotTransform ? BispacedMeasure(q_original) : BispacedMeasure(q_original, q_z, hash(f)) + result = (; stop_reason, niterations, nfresh, nevals, ngeometries, nfailed, nproduction, + nseed_evals, nseed_exhausted, nhessians, npilot, pilot_ess, ncomponents = length(q.components), ncomponent_proposals, + ess = output.ess, efficiency = output.efficiency, logweight_scale = output.logweight_scale, + pareto_k = _mw_pareto_k(logw), max_weight = maximum(output.weight) / sum(output.weight), + pilot_efficiency, diagnostic, history) + return EvaluatedMeasure(em; + transform_intent = alg.pretransform, + f_transform = _viewrep_f(f, alg.pretransform), + empirical = _viewrep_empirical(dsm, smpls_z, f, alg.pretransform, dim, output.ess), + approx, samplegen = nothing, dof = dim, + transformed = _viewrep_measure(transformed_m, alg.pretransform), + evalinfo = MeasureEvalInfo(alg, result) + ) +end diff --git a/src/samplers/importance/molewhacker_geometry.jl b/src/samplers/importance/molewhacker_geometry.jl new file mode 100644 index 000000000..c1e4a104c --- /dev/null +++ b/src/samplers/importance/molewhacker_geometry.jl @@ -0,0 +1,160 @@ +# This file is a part of BAT.jl, licensed under the MIT License (MIT). + +# Pull back distribution Fisher information without constructing its parameter-space +# matrix. In particular, the covariance block of an m-dimensional normal would +# otherwise require O(m^4) storage. +struct MolewhackerGeometryError <: Exception end + +function _mw_regular_positive(x) + isfinite(x) && x > 0 || throw(MolewhackerGeometryError()) + return x +end + +_mw_parameters(d::Union{Normal,Poisson,Exponential}) = collect(params(d)) +_mw_parameters(d::MvNormal) = vcat(mean(d), vec(Matrix(cov(d)))) +_mw_parameters(d::MvNormal{<:Real,<:PDiagMat}) = vcat(mean(d), d.Σ.diag) +_mw_parameters(d::MvNormal{<:Real,<:ScalMat}) = vcat(mean(d), d.Σ.value) + +_mw_factors(d::Distributions.Product) = d.v +_mw_factors(d::Distributions.ProductDistribution) = vec(d.dists) +_mw_factors(d::NamedTupleDist) = values(d) + +function _mw_parameters(d::Union{Distributions.Product,Distributions.ProductDistribution,NamedTupleDist}) + return reduce(vcat, map(_mw_parameters, _mw_factors(d))) +end + +_mw_nparams(::Normal) = 2 +_mw_nparams(::Union{Poisson,Exponential}) = 1 +_mw_nparams(d::MvNormal) = length(d) + length(d)^2 +_mw_nparams(d::MvNormal{<:Real,<:PDiagMat}) = 2 * length(d) +_mw_nparams(d::MvNormal{<:Real,<:ScalMat}) = length(d) + 1 +function _mw_nparams(d::Union{Distributions.Product,Distributions.ProductDistribution,NamedTupleDist}) + return sum(_mw_nparams, _mw_factors(d)) +end + +function _mw_parameters(d) + throw(ArgumentError("MolewhackerSampling has no Fisher geometry for $(typeof(d)). Supported families are Normal, MvNormal, Poisson, Exponential, and their products.")) +end + +function _mw_whitened_jacobian(d::Normal, J) + σ = _mw_regular_positive(std(d)) + A = J ./ σ + view(A, 2:2, :) .*= sqrt(eltype(A)(2)) + return A +end + +function _mw_whitened_jacobian(d::Union{Poisson,Exponential}, J) + m = _mw_regular_positive(mean(d)) + scale = d isa Poisson ? sqrt(m) : m + return J ./ scale +end + +function _mw_whitened_jacobian(d::MvNormal{<:Real,<:PDiagMat}, J) + m = length(d) + variances = _mw_regular_positive.(d.Σ.diag) + scales = sqrt.(variances) + Jμ, Jv = view(J, 1:m, :), view(J, m+1:2m, :) + T = promote_type(eltype(J), eltype(variances)) + # The compact covariance parameters are variances, with information 1/(2v²). + return vcat(Jμ ./ scales, Jv ./ variances ./ sqrt(T(2))) +end + +function _mw_whitened_jacobian(d::MvNormal{<:Real,<:ScalMat}, J) + m = length(d) + variance = _mw_regular_positive(d.Σ.value) + Jμ, Jv = view(J, 1:m, :), view(J, m+1:m+1, :) + T = promote_type(eltype(J), typeof(variance)) + scale = sqrt(T(m) / T(2)) + return vcat(Jμ ./ sqrt(variance), Jv ./ variance .* scale) +end + +function _mw_whitened_jacobian(d::MvNormal, J) + m = length(d) + Σ = Matrix(cov(d)) + all(isfinite, Σ) || throw(MolewhackerGeometryError()) + L = cholesky(Symmetric(Σ)).L + A = L \ view(J, 1:m, :) + covariance_J = view(J, m+1:size(J, 1), :) + all(iszero, covariance_J) && return A + B = map(eachcol(covariance_J)) do v + vec(L \ reshape(v, m, m) / L') + end + C = reduce(hcat, B) + C ./= sqrt(eltype(C)(2)) + return vcat(A, C) +end + +function _mw_whitened_jacobian(d::Union{Distributions.Product,Distributions.ProductDistribution,NamedTupleDist}, J) + factors = _mw_factors(d) + ends = cumsum(map(_mw_nparams, factors)) + blocks = map(eachindex(factors)) do i + firstrow = i == firstindex(factors) ? 1 : ends[i - 1] + 1 + _mw_whitened_jacobian(factors[i], view(J, firstrow:ends[i], :)) + end + return reduce(vcat, blocks) +end + +function _mw_pullback(d, J) + # Stack independent information factors before forming the parameter-space Gram matrix. + A = _mw_whitened_jacobian(d, J) + return A' * A +end + +function _mw_jacobian(f, x, ad, nblocks = 1) + # Avoid an unused primal evaluation for default ForwardDiff. Other selectors, + # including configured chunk sizes or tags, retain their generic AD path. + if forward_adtype(ad) == ADSelector(ForwardDiff) + return nblocks > 1 ? _mw_blocked_jacobian(f, x, nblocks) : ForwardDiff.jacobian(f, x) + end + return last(with_jacobian(f, x, AbstractMatrix, ad)) +end + +# Differentiate only the columns in `r`, holding the other coordinates fixed. +function _mw_jacobian_columns(f, x, r) + return ForwardDiff.jacobian(t -> f(vcat(view(x, 1:first(r)-1), t, view(x, last(r)+1:length(x)))), x[r]) +end + +# Idle tasks share one Jacobian through contiguous column blocks. +function _mw_blocked_jacobian(f, x, nblocks) + n = length(x) + ranges = [fld((b - 1) * n, nblocks)+1:fld(b * n, nblocks) for b in 1:nblocks] + tasks = [Threads.@spawn _mw_jacobian_columns(f, x, r) for r in ranges[2:end]] + first_block = try + _mw_jacobian_columns(f, x, first(ranges)) + finally + # Join every block before returning or propagating an exception. + foreach(task -> try wait(task) catch end, tasks) + end + return _mw_assemble_columns(first_block, tasks, ranges, n) +end + +# Function barrier: ForwardDiff picks its chunk at run time, so the block type is known only here. +function _mw_assemble_columns(first_block, tasks, ranges, n) + J = similar(first_block, size(first_block, 1), n) + J[:, first(ranges)] = first_block + for (r, task) in zip(ranges[2:end], tasks) + # Rethrow the block's own exception so geometry failures stay classifiable. + istaskfailed(task) && throw(task.exception) + J[:, r] = fetch(task)::typeof(first_block) + end + return J +end + +function _mw_local_precision(f, x::AbstractVector{T}, ad, nblocks = 1) where T + try + all(isfinite, x) || throw(MolewhackerGeometryError()) + d = f(x) + # Share one model Jacobian across product leaves, then add the prior. + J = _mw_jacobian(_mw_parameters ∘ f, x, ad, nblocks) + # The Jacobian type depends on the run-time chunk, so assert the Gram matrix type. + G = Matrix{T}(_mw_pullback(d, J))::Matrix{T} + all(isfinite, G) || throw(MolewhackerGeometryError()) + P = Matrix(Symmetric(G)) + I + return PDMat(P, cholesky(Symmetric(P))) + catch err + if err isa Union{MolewhackerGeometryError,PosDefException,SingularException} + return nothing + end + rethrow() + end +end diff --git a/src/samplers/samplers.jl b/src/samplers/samplers.jl index ef9e2be4a..cc19c7f7b 100644 --- a/src/samplers/samplers.jl +++ b/src/samplers/samplers.jl @@ -5,3 +5,5 @@ include("pathfinder.jl") include("bat_sample.jl") include("mcmc/mcmc.jl") include("importance/importance_sampler.jl") +include("importance/molewhacker_geometry.jl") +include("importance/molewhacker.jl") diff --git a/src/utils/executors.jl b/src/utils/executors.jl index 103cae19a..6d03d6974 100644 --- a/src/utils/executors.jl +++ b/src/utils/executors.jl @@ -18,12 +18,30 @@ function exec_map!(f::Base.Callable, executor::SequentialExec, Y::AbstractVector end -struct MultiThreadedExec <: BATExecutor end +""" + MultiThreadedExec(; ntasks = Threads.nthreads()) + +Execute work in at most `ntasks` Julia tasks. The default uses the number of +Julia worker threads. Task count does not set the number of sampler candidates. +""" +struct MultiThreadedExec <: BATExecutor + ntasks::Int + + function MultiThreadedExec(; ntasks::Integer = Threads.nthreads()) + @argcheck ntasks > 0 + return new(ntasks) + end +end -function exec_map!(f::Base.Callable, executor::MultiThreadedExec, Y::AbstractVector, X::AbstractVector) +function exec_map!(f::F, executor::MultiThreadedExec, Y::AbstractVector, X::AbstractVector) where {F<:Base.Callable} @argcheck length(X) == length(Y) throw(ArgumentError("Input and output arrays must have equal lengths.")) - @threads for i in 0:(length(Y) - 1) - Y[firstindex(Y) + i] = f(X[firstindex(X) + i]) + n = length(Y) + ntasks = min(executor.ntasks, n) + ntasks <= 1 && return exec_map!(f, SequentialExec(), Y, X) + @sync for task in 1:ntasks + Threads.@spawn for i in fld((task - 1) * n, ntasks):(fld(task * n, ntasks) - 1) + Y[firstindex(Y) + i] = f(X[firstindex(X) + i]) + end end return Y end diff --git a/test/runtests.jl b/test/runtests.jl index b2d2f6e8a..3ef376b2c 100644 --- a/test/runtests.jl +++ b/test/runtests.jl @@ -48,6 +48,7 @@ const test_groups = [ "samplers/mcmc/test_fisher_tuner.jl", "samplers/mcmc/test_multitrafo_tuner.jl", "samplers/importance/test_importance_sampler.jl", + "samplers/importance/test_molewhacker.jl", ], "hmc" => [ "samplers/mcmc/test_hmc_nuts.jl", diff --git a/test/samplers/importance/test_molewhacker.jl b/test/samplers/importance/test_molewhacker.jl new file mode 100644 index 000000000..daa68c34c --- /dev/null +++ b/test/samplers/importance/test_molewhacker.jl @@ -0,0 +1,302 @@ +# This file is a part of BAT.jl, licensed under the MIT License (MIT). + +using BAT, Test, Distributions, LinearAlgebra, Statistics, StableRNGs +using DensityInterface: logdensityof +using MeasureBase: Likelihood, weightedmeasure +using ValueShapes: NamedTupleDist +import ForwardDiff, Optim, OptimizationLBFGSB + +@testset "Molewhacker importance sampling" begin + context(seed = 71; precision = Float64) = BATContext(rng = StableRNG(seed), ad = ForwardDiff, precision = precision) + prior = MvNormal([0.0], [1.0;;]) + target = PosteriorMeasure(Likelihood(z -> Normal(z[1], 0.5), 2.0), prior) + + @testset "Gaussian estimator and adaptation" begin + alg = MolewhackerSampling(nsamples = 4000, batchsize = 512, maxiter = 5, nseeds = 0, exploration_mass = 0.1) + em = evalmeasure(target, alg, context()) + smpls = BAT.samplesof(em) + z = BAT.samplesof(em.empirical.transformed) + q = Distribution(em.approx.transformed) + info = em.evalinfo.result + @test mean(smpls)[1] ≈ 1.6 atol = 0.035 + @test var(smpls)[1] ≈ 0.2 atol = 0.025 + @test info.efficiency > 0.7 + ratios = z.logd .- logpdf.(Ref(q), z.v) + @test z.weight ≈ exp.(ratios .- info.logweight_scale) + mass_estimate = mean(z.weight) * exp(info.logweight_scale) + @test mass_estimate ≈ pdf(Normal(0, sqrt(1.25)), 2) rtol = 0.05 + end + + @testset "Reselection preserves fitting multiplicity and budgets" begin + alg = MolewhackerSampling(nsamples = 64, batchsize = 1, ncandidates = 1, + nseeds = 0, maxiter = 8, maxcomponents = 4) + em = evalmeasure(target, alg, context()) + q, info = Distribution(em.approx.transformed), em.evalinfo.result + @test (info.niterations, info.ncomponents, info.ncomponent_proposals, info.stop_reason) == + (3, 2, 4, :maxcomponents) + @test info.npilot == 1 + # One pool point supplies three selected occurrences of the same Gaussian. + c = mean(last(q.components))[1] + a, b = Normal(), Normal(c, sqrt(0.2)) + fitting(x) = (pdf(a, x) + 3pdf(b, x)) / 4 + masses = [pdf(Normal(1.6, sqrt(0.2)), 0) / fitting(0), + 3pdf(Normal(1.6, sqrt(0.2)), c) / fitting(c)] + masses ./= sum(masses) + points = [[-1.0], [0.0], [1.6], [3.0]] + expected = [masses[1] * pdf(a, x[1]) + masses[2] * pdf(b, x[1]) for x in points] + @test pdf.(Ref(q), points) ≈ expected + end + + @testset "Fisher covariance and independent factors" begin + c = atanh(0.5) + model = z -> MvNormal([2z[1], 2z[1]], [4.0 2tanh(z[1]); 2tanh(z[1]) 1.0]) + p = PosteriorMeasure(Likelihood(model, [0.0, 0.0]), prior) + seeds = [[c]] + alg = MolewhackerSampling(nsamples = 32, maxiter = 0, nseeds = 1, init_mode = nothing, init = ExplicitInit(seeds), laplace_seeds = false) + em = evalmeasure(p, alg, context()) + q = Distribution(em.approx.transformed) + @test cov(last(q.components))[1, 1] ≈ 4 / 25 + + # Compact Normal covariance charts retain the variance-derivative information. + for (model, variance) in ((z -> MvNormal([2z[1], z[1]], Diagonal(exp.([2z[1], z[1]]))), 2 / 17), + (z -> MvNormal([2z[1], z[1]], exp(2z[1])I(2)), 1 / 10)) + p_compact = PosteriorMeasure(Likelihood(model, [0.0, 0.0]), prior) + em_compact = evalmeasure(p_compact, MolewhackerSampling(nsamples = 32, maxiter = 0, + nseeds = 1, init_mode = nothing, init = ExplicitInit([[0.0]]), laplace_seeds = false), context()) + @test cov(last(Distribution(em_compact.approx.transformed).components))[1, 1] ≈ variance + end + + # Likelihood information at zero is 1/4 + 2 + 4. Add the prior once. + product_model = z -> NamedTupleDist(a = Normal(z[1], 2.0), rest = NamedTupleDist(b = Poisson(2exp(z[1])), c = Exponential(3exp(2z[1])))) + p_product = PosteriorMeasure(Likelihood(product_model, (a = 0.0, rest = (b = 2, c = 3.0))), prior) + em_product = evalmeasure(p_product, MolewhackerSampling(nsamples = 32, maxiter = 0, + nseeds = 1, init_mode = nothing, init = ExplicitInit([[0.0]]), laplace_seeds = false), context()) + @test cov(last(Distribution(em_product.approx.transformed).components))[1, 1] ≈ 4 / 29 + + # A stationary forward model: Fisher sees only the prior, the target has curvature 1 + 16. + p_stationary = PosteriorMeasure(Likelihood(z -> Normal(z[1]^2, 0.5), -2.0), prior) + em_laplace = evalmeasure(p_stationary, MolewhackerSampling(nsamples = 32, maxiter = 0, nseeds = 1, + init_mode = nothing, init = ExplicitInit([[0.0]]), laplace_seeds = true, maxevals = 33), context()) + q_laplace = Distribution(em_laplace.approx.transformed) + @test sort([invcov(c)[1, 1] for c in q_laplace.components]) ≈ [1, 17 / 1.2] rtol = 1e-6 + @test probs(q_laplace) ≈ [0.5, 0.5] + @test em_laplace.evalinfo.result.nhessians == 1 + @test em_laplace.evalinfo.result.nevals == 33 + # The Newton polish lets the default mode search stop early. + @test (MolewhackerSampling().init_mode.maxiters, MolewhackerSampling(laplace_seeds = false).init_mode.maxiters) == (50, 1000) + + # Idle tasks split the Jacobian columns. A nonsymmetric map exposes their order. + A = [1.0 2.0 0.0 -1.0 0.5 0.0; 0.0 1.0 3.0 0.0 -2.0 1.0] + p_linear = PosteriorMeasure(Likelihood(z -> MvNormal(A * z, Diagonal([0.5, 2.0])), zeros(2)), MvNormal(zeros(6), I(6))) + em_linear = evalmeasure(p_linear, MolewhackerSampling(nsamples = 32, maxiter = 0, nseeds = 1, init_mode = nothing, + init = ExplicitInit([zeros(6)]), executor = BAT.MultiThreadedExec(ntasks = 2), laplace_seeds = false), context()) + @test invcov(last(Distribution(em_linear.approx.transformed).components)) ≈ I + A' * Diagonal([2.0, 0.5]) * A + + a, b = [-0.5], [1.5] + repeated = [a, a, b] + saved = deepcopy(repeated) + repeated_result = evalmeasure(target, MolewhackerSampling(nsamples = 32, maxiter = 0, + nseeds = 3, init_mode = nothing, init = ExplicitInit(repeated), laplace_seeds = false), context()) + q_repeated = Distribution(repeated_result.approx.transformed) + points = [[-0.5], [0.0], [1.5], [3.0]] + ga, gb = Normal(-0.5, sqrt(0.2)), Normal(1.5, sqrt(0.2)) + uniform(x) = (2pdf(ga, x) + pdf(gb, x)) / 3 + masses = [2pdf(Normal(1.6, sqrt(0.2)), -0.5) / uniform(-0.5), + pdf(Normal(1.6, sqrt(0.2)), 1.5) / uniform(1.5)] + masses ./= sum(masses) + expected = [masses[1] * pdf(ga, x[1]) + masses[2] * pdf(gb, x[1]) for x in points] + @test pdf.(Ref(q_repeated), points) ≈ expected + @test repeated == saved + end + + @testset "Separated nonlinear modes" begin + p = PosteriorMeasure(Likelihood(z -> Normal(z[1]^2, 0.3), 4.0), prior) + alg = MolewhackerSampling(nsamples = 6000, batchsize = 512, maxiter = 5, + nseeds = 2, init_mode = nothing, init = ExplicitInit([[-2.0], [2.0]])) + em = evalmeasure(p, alg, context(72)) + s = BAT.samplesof(em) + right_mass = sum(s.weight .* (first.(s.v) .> 0)) / sum(s.weight) + @test right_mass ≈ 0.5 atol = 0.04 + end + + @testset "Structured prior and coordinate law" begin + p = PosteriorMeasure(Likelihood(x -> Normal(log(x.rate), 0.7), 0.4), NamedTupleDist(rate = LogNormal())) + em = evalmeasure(p, MolewhackerSampling(nsamples = 16, maxiter = 0, + nseeds = 1, init_mode = nothing, init = ExplicitInit([(rate = 1.0,)]), laplace_seeds = false), context(73)) + @test cov(last(Distribution(em.approx.transformed).components))[1, 1] ≈ 0.49 / 1.49 + # Validate both proposal and empirical coordinate pairs, including the Jacobian. + @test BAT.validate_evalmeasure(em; context = context(73)) === em + end + + @testset "RNG, precision, and target scale" begin + alg = MolewhackerSampling(nsamples = 256, batchsize = 128, maxiter = 2, executor = BAT.SequentialExec()) + threaded = MolewhackerSampling(nsamples = 256, batchsize = 128, maxiter = 2, executor = BAT.MultiThreadedExec(ntasks = 2)) + a = evalmeasure(target, alg, context(74, precision = Float32)) + b = evalmeasure(target, threaded, context(74, precision = Float32)) + @test BAT.samplesof(a) == BAT.samplesof(b) + # Test weight-scale invariance away from tied scores in the exact + # Gaussian proposal produced by mode initialization. Fresh rounds draw from + # the whole mixture, whose Float32 masses shift at rounding level with scale. + scale_alg = MolewhackerSampling(nsamples = 256, batchsize = 128, maxiter = 2, nseeds = 0, fresh_rounds = 0) + a = evalmeasure(target, scale_alg, context(74, precision = Float32)) + shifted = evalmeasure(weightedmeasure(1e8, BAT.batmeasure(target)), scale_alg, context(74, precision = Float32)) + @test BAT.samplesof(a).v ≈ BAT.samplesof(shifted).v rtol = 1e-6 + @test BAT.samplesof(a).weight ≈ BAT.samplesof(shifted).weight rtol = 1e-6 + @test BAT.samplesof(a).logd ≈ logdensityof.(Ref(BAT.unevaluated(a)), BAT.samplesof(a).v) + end + + @testset "Mixture refit from fresh draws" begin + # From the prior, one fresh round and the weighted EM fit recover this Gaussian + # posterior: efficiency 0.95-0.98 over six seeds, against 0.63-0.86 without them. + refit(k; kw...) = evalmeasure(target, MolewhackerSampling(; nsamples = 2000, batchsize = 500, maxiter = 1, + nseeds = 0, fresh_rounds = k, kw...), context(72)).evalinfo.result + with, without = refit(1), refit(0) + @test with.efficiency > 0.93 > without.efficiency + # Fourteen fresh candidates and one to six fitted Gaussians. + @test with.ncomponents - without.ncomponents - 14 in 1:6 + @test refit(1; refit = nothing).ncomponents - without.ncomponents == 14 + @test refit(1; refit = MolewhackerRefit(maxcomponents = 1)).ncomponents - without.ncomponents == 15 + limited = refit(1; ncandidates = 1, maxcomponents = 3) + @test limited.ncomponents <= limited.ncomponent_proposals <= 3 + fitted = refit(1; ncandidates = 1, maxcomponents = 4) + @test fitted.ncomponents == fitted.ncomponent_proposals == 4 + end + + @testset "Cached pool scoring" begin + # A second round adds points and components. The cache matches full scoring, + # and a zero limit takes the uncached path. + rng = StableRNG(3) + comps = [BAT._mw_gaussian(randn(rng, 4), let A = randn(rng, 4, 4); A'A / 4 + I end) for _ in 1:12] + x1, x = randn(rng, 4, 50), randn(rng, 4, 80) + x[:, 1:50] = x1 + w = rand(rng, 12) + w[3] = 0 + q1, q2 = MixtureModel(comps[1:7], w[1:7] ./ sum(w[1:7])), MixtureModel(comps, w ./ sum(w)) + ex = BAT.MultiThreadedExec(ntasks = 2) + l1, cache = BAT._mw_pool_logpdf(q1, comps[1:7], x1, zeros(0, 0), ex) + l2, cache = BAT._mw_pool_logpdf(q2, comps, x, cache, ex) + @test l1 ≈ BAT._mw_batched_logpdf(q1, x1) && l2 ≈ BAT._mw_batched_logpdf(q2, x) && size(cache) == (80, 12) + partial, cache = BAT._mw_pool_logpdf(q2, comps, x, cache, ex, 80 * 6) + @test partial ≈ l2 && length(cache) <= 80 * 6 + @test first(BAT._mw_pool_logpdf(q2, comps, x, cache, ex, 0)) ≈ l2 + # The uncached part can have finite density where all cached terms underflow. + separated = [BAT._mw_gaussian([c], [1.0;;]) for c in (1e200, 0.0)] + q = MixtureModel(separated, [0.5, 0.5]) + @test only(first(BAT._mw_pool_logpdf(q, separated, zeros(1, 1), zeros(0, 0), ex, 1))) ≈ logpdf(q, [0.0]) + end + + @testset "Bounded target concurrency" begin + workers, guard = Set{Task}(), ReentrantLock() + function model(z) + lock(guard) do + push!(workers, current_task()) + end + return Normal(z[1], 0.5) + end + p = PosteriorMeasure(Likelihood(model, 2.0), prior) + for (ntasks, expected) in ((1, 1), (2, 3)) + empty!(workers) + alg = MolewhackerSampling(nsamples = 17, maxiter = 0, nseeds = 0, + executor = BAT.MultiThreadedExec(ntasks = ntasks)) + evalmeasure(p, alg, context()) + # The first point uses the caller. A one-task limit keeps all work there. + @test length(workers) == expected + end + end + + @testset "Budgets, pilot sizing, and mode initialization" begin + bounded = evalmeasure(target, MolewhackerSampling(batchsize = 64, + maxiter = 10, maxevals = 320, nseeds = 0), context()) + @test bounded.evalinfo.result.nevals <= 320 + @test length(BAT.samplesof(bounded)) == bounded.evalinfo.result.npilot + + flat = PosteriorMeasure(Likelihood(z -> Normal(0.0, 1.0), 0.0), prior) + sized = evalmeasure(flat, MolewhackerSampling(target_ess = 12_000, + batchsize = 64, maxiter = 0, maxevals = 20_064, nseeds = 0), context()) + @test 12_000 <= length(BAT.samplesof(sized)) <= 12_001 + @test sized.evalinfo.result.nevals == 64 + length(BAT.samplesof(sized)) + # The prior proposal matches this flat target, so all weights are equal. + @test sized.evalinfo.result.max_weight ≈ 1 / length(BAT.samplesof(sized)) + budgets = evalmeasure(flat, MolewhackerSampling(nsamples = 256, target_ess = 32, + batchsize = 64, maxiter = 1, nseeds = 0), context()) + @test budgets.evalinfo.result.niterations == 1 + for (rule, reason) in (((; target_pool_ess = 32), :pilot_ess), + ((; target_efficiency = 128 / 256), :pool_efficiency)) + stopped = evalmeasure(flat, MolewhackerSampling(; nsamples = 256, target_ess = 128, + batchsize = 64, maxiter = 10, nseeds = 0, fresh_rounds = 0, rule...), context()) + info = stopped.evalinfo.result + @test (info.niterations, info.ncomponents, info.stop_reason) == (0, 1, reason) + @test info.ess ≈ 128 + # Fresh rounds still follow a threshold stop. + refreshed = evalmeasure(flat, MolewhackerSampling(; nsamples = 256, batchsize = 64, maxiter = 10, + nseeds = 0, fresh_rounds = 1, rule...), context()).evalinfo.result + @test (refreshed.niterations, refreshed.nfresh, refreshed.stop_reason) == (0, 1, reason) + end + + concentrated = PosteriorMeasure(Likelihood(z -> MvNormal(z, 0.25I(18)), fill(2.0, 18)), + MvNormal(zeros(18), I(18))) + seeded = evalmeasure(concentrated, MolewhackerSampling(nsamples = 1000, maxiter = 0, maxcomponents = 20), context()) + @test seeded.evalinfo.result.efficiency > 0.8 + @test maximum(abs, mean(BAT.samplesof(seeded)) .- 1.6) < 0.06 + limited = evalmeasure(target, MolewhackerSampling(nsamples = 64, maxiter = 0, nseeds = 1, + init = ExplicitInit([[0.0]]), init_mode = OptimAlg(optalg = Optim.LBFGS()), maxevals = 66), context()) + @test limited.evalinfo.result.nevals <= 66 + @test limited.evalinfo.result.nseed_exhausted == 1 + # A fresh round follows adaptation and adds a proposal batch before selection. + fresh = evalmeasure(target, MolewhackerSampling(nsamples = 32, batchsize = 64, maxiter = 1, + fresh_rounds = 1, nseeds = 0), context()).evalinfo.result + h = fresh.history + @test (fresh.niterations, fresh.nfresh, getproperty.(h, :fresh)) == (1, 1, [false, true]) + @test (h[1].npilot - h[1].drawn, h[2].npilot - h[2].drawn) == (64, h[1].npilot + 64) + end + + @testset "Defensive nonlinear tails" begin + p = PosteriorMeasure(Likelihood(z -> Normal(2tanh(z[1]), 1.0), 0.0), prior) + ε = 0.1 + em = evalmeasure(p, MolewhackerSampling(nsamples = 16, maxiter = 0, + nseeds = 1, init_mode = nothing, init = ExplicitInit([[0.0]]), laplace_seeds = false, exploration_mass = ε), context()) + q = Distribution(em.approx.transformed) + points = [[0.0], [-12.0], [12.0]] + # Fisher variance at zero is 1/5. This local Gaussian alone has infinite IS variance. + logw = logdensityof.(Ref(BAT.batmeasure(p)), points) .- logpdf.(Ref(q), points) + @test all(logw .<= logpdf(Normal(), 0.0) - log(ε)) + + # The local Gaussian alone gives unbounded weights exp(2z²), a heavy Pareto tail. + # Prior mixing bounds them, so the fitted shape turns negative. PSIS changes + # only the largest raw ratios. + n = 2000 + tails(; kw...) = evalmeasure(p, MolewhackerSampling(; nsamples = n, maxiter = 0, nseeds = 1, + init_mode = nothing, init = ExplicitInit([[0.0]]), laplace_seeds = false, kw...), context(75)) + raw, mixed, smoothed = tails(), tails(exploration_mass = ε), tails(smooth_weights = true) + @test mixed.evalinfo.result.pareto_k < 0 < raw.evalinfo.result.pareto_k + logratio(em) = log.(BAT.samplesof(em).weight) .+ em.evalinfo.result.logweight_scale + r, s = logratio(raw), logratio(smoothed) + tail = partialsortperm(r, 1:ceil(Int, 3sqrt(n)), rev = true) + @test r[setdiff(eachindex(r), tail)] ≈ s[setdiff(eachindex(s), tail)] + @test maximum(s) <= maximum(r) && issorted(s[reverse(tail)]) + end + + @testset "Local failure and zero weights" begin + # A singular seed cannot supply Fisher geometry. Keep the prior proposal. + p = PosteriorMeasure(Likelihood(z -> Poisson(z[1] > 0 ? exp(z[1]) : 0.0), 1), prior) + em = evalmeasure(p, MolewhackerSampling(nsamples = 256, maxiter = 0, + nseeds = 1, init_mode = nothing, init = ExplicitInit([[-1.0]])), context()) + s = BAT.samplesof(em) + @test logpdf.(Ref(Distribution(em.approx.transformed)), s.v) ≈ logpdf.(Ref(prior), s.v) + @test all(iszero, s.weight[first.(s.v) .<= 0]) + @test sum(s.weight) > 0 + recovered = evalmeasure(p, MolewhackerSampling(nsamples = 32, maxiter = 0, + nseeds = 3, init_mode = nothing, init = ExplicitInit([[-1.0], [-1.0], [1.0]]), laplace_seeds = false), context()) + q_recovered = Distribution(recovered.approx.transformed) + points = [[0.0], [1.0], [3.0]] + expected = [pdf(Normal(1, sqrt(1 / (1 + exp(1)))), x[1]) for x in points] + @test pdf.(Ref(q_recovered), points) ≈ expected + half_target = PosteriorMeasure(Likelihood(z -> Poisson(z[1] > 0 ? 1.0 : 0.0), 1), prior) + no_mass_round = evalmeasure(half_target, MolewhackerSampling(nsamples = 128, batchsize = 1, + maxiter = 5, nseeds = 0), BATContext(rng = BAT.Random.Xoshiro(6), ad = ForwardDiff)) + info = no_mass_round.evalinfo.result + # One fresh draw follows the stop and finds no mass either. + @test (info.nevals, info.nfresh) == (131, 1) + @test info.stop_reason == :no_finite_candidate + end +end