Feature Request
Describe the Feature Request
GPJax has no tools to score probabilistic predictions. Weather and climate verification uses proper scoring rules and calibration checks, mainly the continuous ranked probability score (CRPS), probability integral transform (PIT) histograms, and the coverage of central intervals. At the moment users must write these themselves. For example, the Gridded Data with xarray example (#808) computes the 2-sd coverage by hand.
Describe Preferred Solution
A new module, gpjax.metrics, with pure JAX functions that work under jit and vmap:
from gpjax.metrics import crps, pit, interval_coverage
crps(dist, y) # per-point CRPS; Gaussian closed form
crps(samples, y, sample_axis=0) # CRPS from samples (ensemble form)
pit(dist, y) # F(y), per point
interval_coverage(dist, y, level=0.95) # fraction of y in the central interval
-
Gaussian CRPS, closed form: $\mathrm{CRPS}(\mathcal{N}(\mu,\sigma^2), y) = \sigma\left[z,(2\Phi(z)-1) + 2\varphi(z) - 1/\sqrt{\pi}\right]$, where $z = (y-\mu)/\sigma$. It takes the marginals of a
GaussianDistribution, so a diagonal covariance is enough.
-
Sample CRPS: $\mathbb{E}|X-y| - \tfrac12,\mathbb{E}|X-X'|$. Use the $O(m \log m)$ sorted form, and offer the "fair" (unbiased) version as an option. This covers non-Gaussian likelihoods and joint samples.
-
PIT: the marginal CDF at the observation. For discrete likelihoods, use the randomised PIT.
-
Interval coverage: for any level.
-
xarray helpers: functions that take the
{target}_mean / {target}_variance datasets from GridSpec.predict, or the sample dim from GridSpec.to_xarray. They return labelled scores and keep NaN cells out of the averages. They live in gpjax.xarray, or behind the same optional import.
Describe Alternatives
properscoring and scoringrules. These are NumPy-based (scoringrules has JAX backends for some scores), and adding them as dependencies only for CRPS is heavy.
xskillscore works on xarray, but not under JAX transforms.
Related Code
gpjax/distributions.py (GaussianDistribution.mean, .variance)
gpjax/xarray.py (GridSpec.predict, GridSpec.to_xarray)
docs/examples/xarray_workflow.py: change the hand-written coverage to use the new functions.
Additional Context
Acceptance criteria:
- The closed-form and sample CRPS agree for Gaussians to Monte Carlo error. Test with hypothesis.
- PIT values are uniform for calibrated synthetic data.
- Every function works under
jit and vmap, and is parametrised over float32/float64 in the tests.
- A short reference page, and use in the xarray example.
References: Gneiting & Raftery (2007), Strictly proper scoring rules, prediction, and estimation, JASA 102:359–378. Hersbach (2000), Decomposition of the continuous ranked probability score for ensemble prediction systems, Weather and Forecasting 15:559–570.
If the feature request is approved, would you be willing to submit a PR?
Yes
Feature Request
Describe the Feature Request
GPJax has no tools to score probabilistic predictions. Weather and climate verification uses proper scoring rules and calibration checks, mainly the continuous ranked probability score (CRPS), probability integral transform (PIT) histograms, and the coverage of central intervals. At the moment users must write these themselves. For example, the Gridded Data with xarray example (#808) computes the 2-sd coverage by hand.
Describe Preferred Solution
A new module,
gpjax.metrics, with pure JAX functions that work underjitandvmap:GaussianDistribution, so a diagonal covariance is enough.{target}_mean/{target}_variancedatasets fromGridSpec.predict, or thesampledim fromGridSpec.to_xarray. They return labelled scores and keep NaN cells out of the averages. They live ingpjax.xarray, or behind the same optional import.Describe Alternatives
properscoringandscoringrules. These are NumPy-based (scoringrules has JAX backends for some scores), and adding them as dependencies only for CRPS is heavy.xskillscoreworks on xarray, but not under JAX transforms.Related Code
gpjax/distributions.py(GaussianDistribution.mean,.variance)gpjax/xarray.py(GridSpec.predict,GridSpec.to_xarray)docs/examples/xarray_workflow.py: change the hand-written coverage to use the new functions.Additional Context
Acceptance criteria:
jitandvmap, and is parametrised over float32/float64 in the tests.References: Gneiting & Raftery (2007), Strictly proper scoring rules, prediction, and estimation, JASA 102:359–378. Hersbach (2000), Decomposition of the continuous ranked probability score for ensemble prediction systems, Weather and Forecasting 15:559–570.
If the feature request is approved, would you be willing to submit a PR?
Yes