Skip to content

feat(metrics): CRPS, PIT and interval coverage in gpjax.metrics #809

Description

@thomaspinder

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

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions