Skip to content

feat(xarray): input transforms, chunked GridSpec.predict, and a surface-temperature infilling example - #808

Open
thomaspinder wants to merge 4 commits into
mainfrom
feat/xarray-transforms-predict
Open

thomaspinder wants to merge 4 commits into
mainfrom
feat/xarray-transforms-predict

Conversation

@thomaspinder

@thomaspinder thomaspinder commented Oct 4, 2026 •

Copy link
Copy Markdown
Collaborator

Summary

Input transforms and chunked prediction for gpjax.xarray. The xarray introduction moves to Getting started, and a new applied example infills a real, complete reanalysis field.

 from_xarray(ds, target, inputs,
+            transforms=[UnitSphere(), Cyclic("time", 365.25), Standardise([...])])
 GridSpec
+  .columns                      # names of the X columns after the transforms
   .inputs_for(grid)             # applies the same fitted transforms
+  .predict(predict_fn, grid, chunk_size=4096)   # mean + variance, in chunks
   .to_xarray(dist)
spec.predict(fn, grid)
  broadcast the inputs onto the grid (lazy if Dask)
  xr.apply_ufunc(dask="parallelized"), per block:
    encode the inputs (time origins, then the fitted transforms)
    jit(fn) on chunk_size rows, last chunk padded  -> 1 compilation
    put the results back into a NaN-masked grid
  -> {target}_mean, {target}_variance, target attrs, squared units
 docs/
 ├── index.md                                  # toctree: intro -> Getting started, new page -> Applied modelling
 └── examples/
-    ├── xarray_workflow.py                    # "Gridded Data with xarray", Applied modelling
+    ├── xarray_workflow.py                    # "Working with Gridded Data", Getting started (same URL);
+    │                                         #   synthetic field, now with Standardise, predict, UnitSphere/Cyclic
+    ├── infilling_surface_temperature.py      # "Infilling Global Surface Temperature": NCEP-NCAR R1 2024
+    │                                         #   anomaly, random vs value-dependent removal
     └── data/
+        ├── ncep_air_anomaly_2024.nc          # 44 KB, netCDF3 (SciPy reads it), CF attrs
         └── _pull_reference_datasets.py       # + pull_ncep_anomaly(); docs builds stay offline
  • UnitSphere maps lat/lon to 3D points on the unit sphere, so a stationary kernel uses the chord distance. This removes the seam at the antimeridian and the distortion at the poles.
  • Standardise fits on the kept training cells only. Cyclic encodes periodic inputs, for example the seasonal cycle.
  • Transforms are opt-in. Existing calls give the same results.
  • dask is added to the dev group only, for the lazy-grid test.

Evidence

  • Before: no transforms, and prediction needed the full (M, D) input matrix and the full distribution in memory.
    After: uv run poe test: 3256 passed, 1 skipped. The 23 new tests in tests/test_xarray.py include:

    • predict == to_xarray(fn(inputs_for(grid))) for chunk_size in {1, 5, 12, 100}, and on a fitted posterior
    • one trace for any number of chunks
    • a Dask grid stays lazy (predict_fn is not called) until .compute(), and then matches the eager result
    • UnitSphere: unit-norm rows, lon 0 ≡ 360, and all longitudes at a pole give one point
    • the error paths
  • uv run poe lint and uv run poe docstrings pass. A fresh docs-ci build (sphinx-build -E -W) passes with 0 warnings.

  • Infilling Global Surface Temperature results (fit on 5°, scored on the 2.5° truth, 2024 global mean truth 0.64 K):

    Random removal (keep 25%) Warm cells removed more often (keep 21%)
    RMSE: GP / mean of kept cells 0.33 K / 1.05 K 0.97 K / 1.37 K
    Within 2 sd 94% 81%
    Global mean from joint samples 0.63 ± 0.04 K 0.34 ± 0.05 K (interval does not hold the truth)

Merge Danger

Door: two-way

The change only adds to the public API: new classes, a new keyword with a default, and new methods. GridSpec gets a new last field with a default, so the existing constructor calls still work.

Blast Radius: small

  • Only users of gpjax.xarray see the new API.
  • examples/xarray_workflow keeps its URL and its synthetic content, with the new API. The page moves from Applied modelling to Getting started, and the new applied page has a new URL.
  • The repo gains one 44 KB data file.

🤖 Generated with Claude Code

thomaspinder and others added 2 commits October 4, 2026 12:27
from_xarray takes transforms (Standardise, UnitSphere, Cyclic) that turn
the named inputs into the columns of X. The GridSpec keeps them, fitted on
the training cells, and applies them to every new grid.

GridSpec.predict gives the predictive mean and variance on a grid in
chunks of fixed size, compiled once, and keeps a Dask-backed grid lazy.
Dask is a dev dependency for the lazy-grid test only.

Co-Authored-By: Claude Opus 5.5 <[email protected]>
The example now reads a netCDF file of the 2024 NCEP-NCAR Reanalysis 1
temperature anomaly, which is complete, removes cells, and scores the GP
infill against the truth. Random removal gives a global mean whose
interval holds the truth; value-dependent removal shows the bias of a GP
when data are missing not at random.

The pull script records how the 44 KB file was made, so docs builds stay
offline.

Co-Authored-By: Claude Opus 5.5 <[email protected]>
@github-actions github-actions Bot added documentation Improvements or additions to documentation ci Continuous Integration dependencies Pull requests that update a dependency file tests examples dev-tools formatting release size/xl labels Oct 4, 2026
@github-actions

github-actions Bot commented Oct 4, 2026

Copy link
Copy Markdown

📖 Docs preview: https://pr-808--endearing-crepe-c2d5fe.netlify.app

Smoke render — the expensive notebooks run with reduced budgets, so
figures are not publication fidelity. /render-mode.txt says smoke.

Working with Gridded Data (examples/xarray_workflow, same URL) is the
synthetic introduction to gpjax.xarray again, now under Getting started.
It uses Standardise and GridSpec.predict, and shows UnitSphere and Cyclic.

Infilling Global Surface Temperature (examples/infilling_surface_temperature)
holds the reanalysis study, under Applied modelling. It also drops an empty
markdown cell at the top of the notebook.

Co-Authored-By: Claude Opus 5.5 <[email protected]>
@thomaspinder thomaspinder changed the title feat(xarray): input transforms, chunked GridSpec.predict, and a real-data example feat(xarray): input transforms, chunked GridSpec.predict, and a surface-temperature infilling example Oct 4, 2026
@thomaspinder

Copy link
Copy Markdown
Collaborator Author

@claude-pr-review

@claude

claude Bot commented Oct 4, 2026 •

Copy link
Copy Markdown

Claude finished @thomaspinder's task in 25s —— View job


PR Review

This PR adds input transforms (Standardise, UnitSphere, Cyclic) and a chunked GridSpec.predict to gpjax.xarray. It also adds an infilling example. The code is correct and follows the AGENTS.md rules I checked. I found one documentation problem.

In-scope

  • P2 The docs overstate the memory bound of GridSpec.predict. gpjax/xarray.py#L381
    An eager grid is a single block. predict_block builds the full (M, D) raw input matrix and then a second encoded copy. Only the predict_fn call is chunked. The docstring says the memory use does not grow with grid size. That is true only for Dask-backed grids. Fix: state that chunk_size bounds the memory of predict_fn, and that Dask input is needed to bound the input matrix too. Make the same change in the PR description.

Out-of-scope

  • Cache the jitted moments function in GridSpec.predict. predict makes a new jax.jit wrapper on each call, so each call compiles again. This is fine for now. If users call predict in a loop, a cache keyed on predict_fn would help. File: gpjax/xarray.py.

I checked these points and found no problem: val() use (not relevant to this diff), no new paramax.unwrap, tests next to their module, the new GridSpec field has a default, and UnitSphere handles NaN and out-of-range latitude. I did not run the tests or the docs build.

Comment thread gpjax/xarray.py
_encode(name, block.ravel(), self.time_origins)
for name, block in zip(self.inputs, blocks, strict=True)
]
input_matrix = np.stack(columns, axis=1)

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 For an eager (non-Dask) grid, the whole grid is one block. This line builds the full (M, D) raw input matrix, and _encode then builds a second full copy. Only the predict_fn call is chunked. The docstring says "the memory use does not grow with the size of the grid". That is true only for Dask grids. Change the docstring to say this, for example: "Memory for predict_fn is bounded by chunk_size. For eager input, the (M, D) input matrix is still held in memory. Use Dask-backed input to bound it too." The PR description makes the same claim.

This branch has not been deployed

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

Labels

ci Continuous Integration dependencies Pull requests that update a dependency file dev-tools documentation Improvements or additions to documentation examples formatting release size/xl tests

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant