From 7800f90e915104a70b6ee118734f5d07b4bf0139 Mon Sep 17 00:00:00 2001 From: Adrian Price-Whelan <583379+adrn@users.noreply.github.com> Date: Mon, 14 Sep 2026 11:22:27 -0400 Subject: [PATCH 1/3] fix KeyError('eccentricity') for EcoswEsinwRV default priors EcoswEsinwRV().default_prior(sigma_K0=...) gives rv_semiamp a PeriodDependentKPrior, whose __call__ reads params["eccentricity"]. That parameterization carries (ecosw, esinw) instead, so the key does not exist and every sampler run raised KeyError -- the documented happy path was unusable. The parameterization already knows the conversion, so rather than duplicating sqrt(ecosw^2 + esinw^2) inside the prior, AbstractParameterization grows a derived_eccentricity() hook returning None by default. That default covers both parameterizations carrying eccentricity outright (nothing to derive) and the Kepler-free Fourier bases (none exists); EcoswEsinwRV overrides it by delegating to its existing eccentricity(). The values dict handed to callable priors is enriched at the point of use, so any future eccentricity-dependent prior works too. Reading .parameterization off the component model at those seven sites also needs it declared on AbstractComponentModel, which previously declared only extensions. Declaring it as an eqx.AbstractVar matches how extensions is already declared and makes the class docstring true: it has always said component models "carry only parameterization and extensions", while only one of the two was actually part of the interface. Shared priors in a JointModel are left alone on purpose: components need not agree on a parameterization, so there is no unambiguous value to derive. The three existing EcoswEsinwRV prior tests only *constructed* the prior and inspected its keys, which is why this survived. The new tests evaluate it through an actual run, and fail with the original KeyError when the fix is reverted. Co-Authored-By: Claude Opus 5 (1M context) --- src/harv/models/_helpers.py | 32 +++++++++ src/harv/models/component.py | 25 ++++++- src/harv/models/joint.py | 15 +++- src/harv/models/parameterizations/_base.py | 27 ++++++++ src/harv/models/parameterizations/rv.py | 11 +++ src/harv/samplers/numpyro.py | 1 + tests/unit/models/test_ecosw_esinw.py | 81 ++++++++++++++++++++++ 7 files changed, 188 insertions(+), 4 deletions(-) diff --git a/src/harv/models/_helpers.py b/src/harv/models/_helpers.py index 3b3a3f0..2f57985 100644 --- a/src/harv/models/_helpers.py +++ b/src/harv/models/_helpers.py @@ -84,11 +84,42 @@ def _needs_explicit_sampling(d: PriorDist | LinearPriorCallable) -> bool: return not (callable(d) and not isinstance(d, dist.Distribution)) +def _with_derived_eccentricity( + param_values: dict[str, Any], + parameterization: Any | None, +) -> dict[str, Any]: + """Add ``eccentricity`` when the parameterization implies one but does not carry it. + + ``LinearPriorCallable`` implementations are written against the standard parameter + names, so ``EcoswEsinwRV`` -- which carries ``(ecosw, esinw)`` instead -- would + otherwise raise ``KeyError: 'eccentricity'`` for the default ``rv_semiamp`` prior. + The parameterization already owns the conversion; this only applies it. + + Parameters + ---------- + param_values + Values that will be handed to a callable prior. + parameterization + The parameterization the values came from, or ``None`` to skip. + + Returns + ------- + ``param_values`` unchanged, or a copy with ``eccentricity`` added. + """ + if parameterization is None or "eccentricity" in param_values: + return param_values + ecc = parameterization.derived_eccentricity(param_values) + if ecc is None: + return param_values + return {**param_values, "eccentricity": ecc} + + def _resolve_prior_to_mvn( prior_dict: dict[str, PriorDist | LinearPriorCallable], nl_values: dict[str, Any], unit_dict: dict[str, str], extra_values: dict[str, Any] | None = None, + parameterization: Any | None = None, ) -> dist.MultivariateNormal: """Build diagonal MVN from per-parameter priors.""" locs: list[Any] = [] @@ -99,6 +130,7 @@ def _resolve_prior_to_mvn( param_values = dict(nl_values) if extra_values: param_values.update(extra_values) + param_values = _with_derived_eccentricity(param_values, parameterization) for name, prior in prior_dict.items(): target_u = unit_dict.get(name, "") resolved = None diff --git a/src/harv/models/component.py b/src/harv/models/component.py index 9324a94..50d2e63 100644 --- a/src/harv/models/component.py +++ b/src/harv/models/component.py @@ -27,8 +27,10 @@ _needs_explicit_sampling, _resolve_prior_to_mvn, _unwrap_dist, + _with_derived_eccentricity, ) from harv.models.extensions.base import AbstractExtension, ParamInfo +from harv.models.parameterizations._base import AbstractParameterization from harv.stats import MarginalizedLinear from harv.stats.linear_op import to_linear_op @@ -78,10 +80,13 @@ class AbstractComponentModel(eqx.Module): Concrete subclasses must also declare: + - ``parameterization``: the AbstractParameterization describing the + parameter set and design matrix. - ``extensions``: tuple of AbstractExtension (model modifiers). """ # Concrete subclasses must declare: + parameterization: eqx.AbstractVar[AbstractParameterization] extensions: eqx.AbstractVar[tuple[AbstractExtension, ...]] # Subclass hooks @@ -408,7 +413,11 @@ def _build_marg_blocks( u = param_units.get(name, "") extra_q[name] = Q(val, u) if u else val lp = _resolve_prior_to_mvn( - prior_dict, nl_values, unit_dict, extra_values=extra_q + prior_dict, + nl_values, + unit_dict, + extra_values=extra_q, + parameterization=self.parameterization, ) return _MargBuildingBlocks( @@ -938,6 +947,7 @@ def model_fn() -> None: nl_values, {name: target_unit}, extra_values=explicit_linear_q, + parameterization=component.parameterization, ) raw = numpyro.sample( name, @@ -1013,11 +1023,20 @@ def model_fn() -> None: if callable(d) and not isinstance( d, dist.Distribution | QuantityDistribution ): - resolved_lp[name] = d(nl_values) + resolved_lp[name] = d( + _with_derived_eccentricity( + dict(nl_values), component.parameterization + ) + ) else: resolved_lp[name] = d gaussian_units = {n: param_units.get(n, "") for n in gaussian_names} - mvn = _resolve_prior_to_mvn(resolved_lp, nl_values, gaussian_units) + mvn = _resolve_prior_to_mvn( + resolved_lp, + nl_values, + gaussian_units, + parameterization=component.parameterization, + ) linear_vec = jnp.atleast_1d(numpyro.sample("_linear", mvn)) for i, lname in enumerate(gaussian_names): numpyro.deterministic(lname, linear_vec[i]) diff --git a/src/harv/models/joint.py b/src/harv/models/joint.py index ae0d1b7..d00348d 100644 --- a/src/harv/models/joint.py +++ b/src/harv/models/joint.py @@ -102,6 +102,7 @@ def _sample_explicit_linear_prior( extra_values: dict[str, Any] | None = None, *, site_name: str | None = None, + parameterization: Any | None = None, ) -> jax.Array: """Sample one explicit (non-marginalized) linear prior in a numpyro model. @@ -133,6 +134,11 @@ def _sample_explicit_linear_prior( values are needed (e.g. for shared explicit-linear priors). site_name Optional site name to use within numpyro. + parameterization + Parameterization the values came from, used to derive ``eccentricity`` for + callable priors that need it. ``None`` for *shared* joint priors, where the + components need not agree on a parameterization and there is no unambiguous + answer; per-component priors pass their own. Returns ------- @@ -145,6 +151,7 @@ def _sample_explicit_linear_prior( nl_values, {name: target_unit}, extra_values=extra_values, + parameterization=parameterization, ) return cast( "jax.Array", @@ -1171,13 +1178,19 @@ def model_fn() -> None: target_u, nl_values, site_name=f"{comp_name}.{name}", + parameterization=joint.components[comp_name].parameterization, ) nl_values[f"{comp_name}.{name}"] = raw explicit_linear_q[name] = Q(raw, target_u) if target_u else raw for name, p in _comp_explicit_callable_lp[comp_name].items(): target_u = pu.get(name, "") raw = _sample_explicit_linear_prior( - name, p, target_u, nl_values, extra_values=explicit_linear_q + name, + p, + target_u, + nl_values, + extra_values=explicit_linear_q, + parameterization=joint.components[comp_name].parameterization, ) nl_values[f"{comp_name}.{name}"] = raw explicit_linear_q[name] = Q(raw, target_u) if target_u else raw diff --git a/src/harv/models/parameterizations/_base.py b/src/harv/models/parameterizations/_base.py index 1ed60b2..c77ca9b 100644 --- a/src/harv/models/parameterizations/_base.py +++ b/src/harv/models/parameterizations/_base.py @@ -71,6 +71,33 @@ def default_prior(self, **kwargs: Any) -> "HarvPrior": ) raise NotImplementedError(msg) + def derived_eccentricity(self, nl_values: dict[str, Any]) -> Any | None: # noqa: ARG002 + """Eccentricity implied by this parameterization, if it is not a parameter. + + ``LinearPriorCallable`` implementations such as + :class:`~harv.models.priors.custom_priors.PeriodDependentKPrior` are written + against the standard parameter names, so a parameterization that encodes + eccentricity indirectly must say how to recover it or those priors cannot be + evaluated at all. + + Returns ``None`` by default, which covers both parameterizations that carry + ``eccentricity`` directly as a nonlinear parameter (nothing to derive) and the + Kepler-free Fourier bases (no eccentricity exists). Override when the value is + derivable but absent -- see + :class:`~harv.models.parameterizations.rv.EcoswEsinwRV`. + + Parameters + ---------- + nl_values + Nonlinear parameter values keyed by bare parameter name. + + Returns + ------- + The derived eccentricity, or ``None`` when this parameterization does not + imply one. + """ + return None + def linear_log_prior_correction( self, linear_map: dict[str, jax.Array], # noqa: ARG002 diff --git a/src/harv/models/parameterizations/rv.py b/src/harv/models/parameterizations/rv.py index c40000d..778815b 100644 --- a/src/harv/models/parameterizations/rv.py +++ b/src/harv/models/parameterizations/rv.py @@ -216,6 +216,17 @@ def eccentricity(self, nl_values: dict[str, Any]) -> Any: esinw = nl_values["esinw"] return jnp.sqrt(ecosw**2 + esinw**2) + def derived_eccentricity(self, nl_values: dict[str, Any]) -> Any: + """Recover ``eccentricity`` for priors that expect the standard names. + + This parameterization carries ``(ecosw, esinw)`` rather than + ``(eccentricity, arg_peri)``, so callable linear priors that depend on + eccentricity -- the default ``rv_semiamp`` prior among them -- have no + ``eccentricity`` key to read. See + :meth:`~harv.models.parameterizations._base.AbstractParameterization.derived_eccentricity`. + """ + return self.eccentricity(nl_values) + def strip_nl_for_design(self, nl_values: dict[str, Any]) -> dict[str, Any]: """Return nl_values with units stripped for ``design_matrix``.""" d = dict(nl_values) diff --git a/src/harv/samplers/numpyro.py b/src/harv/samplers/numpyro.py index 351bfd5..b75072e 100644 --- a/src/harv/samplers/numpyro.py +++ b/src/harv/samplers/numpyro.py @@ -199,6 +199,7 @@ def model_fn() -> None: nl_values, {name: target_unit}, extra_values=explicit_linear_q, + parameterization=component.parameterization, ) raw = numpyro.sample( name, diff --git a/tests/unit/models/test_ecosw_esinw.py b/tests/unit/models/test_ecosw_esinw.py index eabdf77..5332cf4 100644 --- a/tests/unit/models/test_ecosw_esinw.py +++ b/tests/unit/models/test_ecosw_esinw.py @@ -1,13 +1,24 @@ """Tests for the EcoswEsinwRV parameterization and its use in RVModel.""" +from typing import Any + import jax import jax.numpy as jnp +import jax.random as jr import numpyro.distributions as dist +import pytest + +# quaxed for anything touching Quantity: `Samples.__getitem__` returns Q, and +# jax.numpy rejects it. +import quaxed.numpy as qnp from unxt import Q +import harv.models as hm from harv.data import RVData from harv.models.parameterizations.rv import EcoswEsinwRV, StandardRV from harv.models.rv import RVModel +from harv.samplers.rejection import RejectionSampler +from harv.simulate.rv import simulate_rv_sb1_data def _make_rv_data(n_obs=20): @@ -283,3 +294,73 @@ def test_marginalized_equivalence(self): assert jnp.allclose(ll_std, ll_eco, atol=1e-6), ( f"std={float(ll_std)}, eco={float(ll_eco)}" ) + + +class TestEcoswEsinwDefaultPriorIsEvaluable: + """The default prior must survive being *evaluated*, not just constructed. + + ``EcoswEsinwRV.default_prior`` gives ``rv_semiamp`` a + :class:`~harv.models.priors.custom_priors.PeriodDependentKPrior`, which reads + ``params["eccentricity"]`` -- a key this parameterization does not have, since it + carries ``(ecosw, esinw)``. The existing prior tests only inspect keys, so the + whole combination raised ``KeyError: 'eccentricity'`` on every sampler run. + """ + + @staticmethod + def _prior_and_model() -> tuple[Any, RVModel]: + param = hm.EcoswEsinwRV() + prior = param.default_prior( + period_min=Q(10.0, "day"), + period_max=Q(3000.0, "day"), + sigma_K0=Q(30.0, "km/s"), + sigma_v0=Q(50.0, "km/s"), + ) + return prior, RVModel(parameterization=param) + + def test_derived_eccentricity_matches_ecosw_esinw(self): + """The parameterization can recover eccentricity from its own params.""" + param = hm.EcoswEsinwRV() + nl = {"ecosw": jnp.asarray(0.3), "esinw": jnp.asarray(0.4)} + assert param.derived_eccentricity(nl) == pytest.approx(0.5) + + def test_base_parameterizations_derive_nothing(self): + """Parameterizations that carry eccentricity (or have none) return None.""" + assert hm.StandardRV().derived_eccentricity({}) is None + assert hm.FourierRV(n_terms=2).derived_eccentricity({}) is None + + @pytest.mark.filterwarnings("ignore:Under-resolved rejection run:UserWarning") + def test_rejection_sampler_runs(self): + """The regression: this raised KeyError('eccentricity') before the fix.""" + data, _ = simulate_rv_sb1_data(seed=42, n_obs=16) + prior, model = self._prior_and_model() + sampler = RejectionSampler(prior, model, batch_size=2000) + # ignore_non_finite: independent Uniform(-1, 1) priors on ecosw/esinw put + # ~21% of draws outside the unit disk (e >= 1, an unbound orbit), where the + # K prior's (1 - e^2)^(-1/2) is NaN. Those draws must be rejected. + samples = sampler.run_with_samples( + data, + prior.sample(jr.key(0), 2000, model=model), + top_k=8, + seed=0, + ignore_non_finite=True, + ) + assert samples.n_samples == 8 + assert jnp.isfinite(samples.metadata["max_log_likelihood"]) + + @pytest.mark.filterwarnings("ignore:Under-resolved rejection run:UserWarning") + def test_unit_disk_violation_is_rejected_not_propagated(self): + """e >= 1 draws must not poison max_log_likelihood. + + Without ignore_non_finite a single NaN propagates through the max + reduction, leaving every evidence statistic NaN. This pins the sharp edge + documented in docs/sharp-bits.md so a change in either direction is visible. + """ + data, _ = simulate_rv_sb1_data(seed=42, n_obs=16) + prior, model = self._prior_and_model() + cache = prior.sample(jr.key(0), 2000, model=model) + ecc = qnp.sqrt(cache["ecosw"] ** 2 + cache["esinw"] ** 2) + assert qnp.any(ecc >= 1.0), "expected the square prior to escape the unit disk" + + sampler = RejectionSampler(prior, model, batch_size=2000) + without = sampler.run_with_samples(data, cache, top_k=8, seed=0) + assert not jnp.isfinite(without.metadata["max_log_likelihood"]) From 1be6fdf41b114bcc3722cb3e2c6f25c6a5a1c9c0 Mon Sep 17 00:00:00 2001 From: Adrian Price-Whelan <583379+adrn@users.noreply.github.com> Date: Mon, 14 Sep 2026 11:23:11 -0400 Subject: [PATCH 2/3] document that EcoswEsinwRV's default prior admits unbound orbits Fixing the KeyError exposed the next problem on the same path. default_prior puts independent Uniform(-1, 1) priors on ecosw and esinw, so the support is a square while a bound orbit needs the unit disk: ~21% of draws (1 - pi/4) have e = sqrt(ecosw^2 + esinw^2) >= 1, where the K prior's (1 - e^2)^(-1/2) is NaN. NaN does not behave like a rejected sample. It propagates through the max reduction the rejection step normalizes by, so max_log_likelihood, logZ_int and logZ_int_ess all come back NaN and nothing is accepted -- silently. Measured on 4000 draws: max_log_likelihood=nan without ignore_non_finite, finite with it. ignore_non_finite=True is the existing, correct remedy, so this documents the trap in sharp-bits.md and the spec rather than changing the prior -- narrowing the support is a modelling decision, not a bugfix, and is tracked separately. Co-Authored-By: Claude Opus 5 (1M context) --- docs/sharp-bits.md | 26 ++++++++++++++++++++++++++ docs/spec.md | 24 ++++++++++++++++++++++-- 2 files changed, 48 insertions(+), 2 deletions(-) diff --git a/docs/sharp-bits.md b/docs/sharp-bits.md index f1e7384..cf45063 100644 --- a/docs/sharp-bits.md +++ b/docs/sharp-bits.md @@ -118,3 +118,29 @@ observations, for the same reason; sources with different epoch counts will retr Both conditions point the same way as the advice in `frequency_grid`: use one shared grid across a population so the results are directly comparable and the sampler compiles once. + +## `EcoswEsinwRV`'s default prior admits unbound orbits + +`EcoswEsinwRV().default_prior(...)` puts independent `Uniform(-1, 1)` priors on `ecosw` +and `esinw`. That support is a square, but a bound orbit needs the unit disk: + +``` +e = sqrt(ecosw**2 + esinw**2) < 1 +``` + +Roughly 21% of draws (`1 - pi/4`) fall outside it with `e >= 1`, which is not an orbit. +The default `rv_semiamp` prior scales as `(1 - e**2)**(-1/2)`, so it returns `NaN` +there. + +`NaN` does not behave like a rejected sample. It propagates through the `max` reduction +the rejection step normalizes by, so `max_log_likelihood`, `logZ_int`, and +`logZ_int_ess` all come back `NaN` and **no samples are accepted** -- with no error +raised. Pass `ignore_non_finite=True` so those draws are treated as rejections: + +```python +samples = sampler.run(data, n_prior_samples=1_000_000, ignore_non_finite=True) +``` + +This is not specific to `EcoswEsinwRV` -- any prior that can produce a non-finite +likelihood needs the same flag -- but that parameterization's default prior guarantees +it on about a fifth of all draws. diff --git a/docs/spec.md b/docs/spec.md index 5c71be5..c47b819 100644 --- a/docs/spec.md +++ b/docs/spec.md @@ -593,6 +593,15 @@ This is the simplest sensible default for this parameterization; users wanting a matched prior should sample with `StandardRV` and convert (or override via `**kwargs`). +**The default prior admits unbound orbits.** `default_prior` puts independent +`Uniform(-1, 1)` priors on `ecosw` and `esinw`, whose support is a *square*, while a +bound orbit requires the *unit disk* (`e = sqrt(ecosw² + esinw²) < 1`). About 21% of +draws (`1 - π/4`) land outside it with `e >= 1`, where the default `rv_semiamp` prior's +`(1 - e²)^(-1/2)` is `NaN`. Those draws must be rejected: pass +`ignore_non_finite=True`, or a single `NaN` propagates through the `max` reduction and +leaves `max_log_likelihood` and every evidence statistic `NaN`. See +`docs/sharp-bits.md`. + ### `StandardGaiaAstrometry` Standard Gaia epoch-astrometry parameterization: @@ -972,8 +981,19 @@ The callable receives a **plain dict** keyed by bare parameter name and must ret `Normal` (bare, or wrapped in a `QuantityDistribution` to declare its unit). The dict contains: -- every **nonlinear** parameter value sampled so far (`period`, `eccentricity`, …), and -- every **explicit** (non-marginalized) **linear** parameter value sampled so far. +- every **nonlinear** parameter value sampled so far (`period`, `eccentricity`, …), +- every **explicit** (non-marginalized) **linear** parameter value sampled so far, and +- `eccentricity`, even under a parameterization that does not carry it as a + parameter, when the parameterization can derive one. + +The third bullet exists because callables like `PeriodDependentKPrior` are written +against the standard parameter names. `EcoswEsinwRV` carries `(ecosw, esinw)` instead, +so without this the default `rv_semiamp` prior could not be evaluated at all. The value +comes from `AbstractParameterization.derived_eccentricity(nl_values)`, which returns +`None` by default -- covering both parameterizations that already carry `eccentricity` +(nothing to derive) and the Kepler-free Fourier bases (no eccentricity exists) -- and is +overridden by `EcoswEsinwRV`. Shared priors in a `JointModel` are the one exception: +components need not agree on a parameterization, so no derivation is attempted there. The second bullet is why `parallax` is readable by callable priors: with the default `HalfNormal` prior it is classified non-Gaussian, so it is sampled explicitly rather From f90e851b7222941028b38e842d65825e934401c1 Mon Sep 17 00:00:00 2001 From: Adrian Price-Whelan <583379+adrn@users.noreply.github.com> Date: Mon, 31 Aug 2026 10:55:42 -0400 Subject: [PATCH 3/3] record the EcoswEsinwRV unit-disk prior bug in the spec Its default prior's support is a square, so about 21% of draws are unbound orbits (e >= 1). ignore_non_finite=True already stops those from silently NaN-poisoning the evidence statistics, but the fifth of every prior library they waste is real and unfixed. The fix is not local. harv.stats.numpyro_ext.UnitDisk is the right distribution but is two-dimensional (event_shape=(2,)), while HarvPrior.nonlinear_priors holds one scalar prior per name. Supporting it means adding a joint multi-name nonlinear prior to the public API, which the spec has to define before any implementation, and it touches seven load-bearing sites. Recorded as a new "Known bugs" section rather than a separate BUGS.md, so that deferred defects sit beside the design they contradict and under the same rule that governs the rest of the spec. The section states the distinction from "Planned features and known gaps": bugs are behavior that is wrong, gaps are behavior that is absent. Includes the measurements, the seven sites, and the three alternatives already ruled out, so picking this up later does not mean re-deriving the analysis. Co-Authored-By: Claude Opus 5 (1M context) --- docs/spec.md | 79 ++++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 79 insertions(+) diff --git a/docs/spec.md b/docs/spec.md index c47b819..f0879fb 100644 --- a/docs/spec.md +++ b/docs/spec.md @@ -2743,6 +2743,85 @@ if it stays the same, the cached compilation is reused. ______________________________________________________________________ +## Known bugs + +Defects in current behavior, recorded with enough analysis that picking one up does +not require re-deriving it. Fixing a bug removes its entry. Features that are absent +rather than wrong belong in "Planned features and known gaps" below. + +### `EcoswEsinwRV.default_prior` admits unbound orbits + +Status: open, mitigated but not fixed. Affects `EcoswEsinwRV.default_prior` and any +caller relying on it. + +`default_prior` puts independent `Uniform(-1, 1)` priors on `ecosw` and `esinw`. That +support is a square, while a bound orbit requires the unit disk, +`e = sqrt(ecosw**2 + esinw**2) < 1`. About 21% of draws (`1 - pi/4`) fall outside it. +Measured over 4000 draws, `e` ranged from 0.015 to 1.404, with 21.3% at `e >= 1`. + +There are two consequences: + +1. A fifth of every prior library is unphysical and can never be accepted. At + `M = 1e7` that is roughly 2.1M dead draws. +1. The default `rv_semiamp` prior (`PeriodDependentKPrior`) scales as + `(1 - e**2)**(-1/2)`, which is `NaN` for `e >= 1`. That `NaN` propagates through + the `max` reduction the rejection step normalizes by, so `max_log_likelihood`, + `logZ_int`, and `logZ_int_ess` all return `NaN` and no samples are accepted, with + no error raised. + +`ignore_non_finite=True` converts those draws into ordinary rejections and restores +finite evidence statistics: measured, `max_log_likelihood` is `nan` without the flag +and `-1709.60` with it. See §`EcoswEsinwRV` and `docs/sharp-bits.md`. The flag removes +the silent-`NaN` behavior, not the wasted 21%. + +#### Why the fix requires a spec decision first + +`harv.stats.numpyro_ext.UnitDisk` is the correct distribution (uniform on the disk, +`log_prob = -log(pi)`), but it is two-dimensional: `event_shape=(2,)`, and `sample` +returns shape `(*sample_shape, 2)`. `HarvPrior.nonlinear_priors` is +`dict[str, PriorDist]`, exactly one scalar prior per parameter name, and +`HarvPrior.sample_nonlinear` draws each independently at shape `(n_samples,)`. A +single distribution spanning both `ecosw` and `esinw` therefore cannot be dropped in. + +Supporting it means adding a **joint (multi-name) nonlinear prior** to the public API: +for example a tuple-keyed entry `{("ecosw", "esinw"): UnitDisk()}`, whose sample's +trailing axis is split across the named parameters and whose `log_prob` is counted +once for the pair. This section must specify the tuple-key form, how +`sample_nonlinear` splits the event axis, how `ln_prior` accumulates it, and how +prior caches key it on disk, before any implementation. + +Sites that assume one scalar prior per name and would need to handle the joint case: + +- `HarvPrior.nonlinear_priors` type and its `__check_init__` validation +- `HarvPrior.sample_nonlinear` -- the per-name zip over `nonlinear_priors` +- `HarvPrior.sample` -- the `all_nonlinear_priors[name]` lookup in the nonlinear + unit-restoration loop +- `_sample_nonlinear_params` in `harv/models/component.py` -- the numpyro path, one + `numpyro.sample(name, ...)` per key +- `_expected_prior_keys` in `harv/samplers/rejection.py` -- prior-cache key validation +- `ln_prior` accumulation, so the pair contributes one `-log(pi)` rather than two terms +- `EcoswEsinwRV.default_prior` itself + +`nonlinear_priors` has roughly 70 references across 12 files; most are annotations or +docstrings, but the seven above are load-bearing. + +#### Alternatives considered + +- **Conditional scalar priors.** Draw `ecosw` from a semicircle density, then + `esinw | ecosw ~ Uniform(+/- sqrt(1 - ecosw**2))`. Exact, and it keeps one prior per + name, but the nonlinear prior machinery has no equivalent of `LinearPriorCallable`, + so a nonlinear prior cannot depend on another nonlinear parameter. This needs its + own new concept. +- **One 2-vector parameter.** Declare a single nonlinear parameter of shape `(2,)`. + This avoids the joint-prior concept but changes the parameterization's public + parameter names, breaking `Samples["ecosw"]`, `convert_parameterization`, and the + design-matrix contract. +- **`UnitDiskTransform`.** Not applicable. It is the `biject_to` transform for + unconstrained MCMC; pushing `Uniform(-1, 1)**2` through it gives a density + proportional to `1 / sqrt(1 - y0**2)`, which is not uniform on the disk. + +______________________________________________________________________ + ## Planned features and known gaps ### JointModel non-marginalized MCMC