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..f0879fb 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 @@ -2723,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 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"])