Skip to content

fix: repair three runtime defects in the WeatherMesh and Aurora models - #246

Open
yasumorishima wants to merge 7 commits into
openclimatefix:mainfrom
yasumorishima:fix/runtime-defects-weathermesh-aurora
Open

yasumorishima wants to merge 7 commits into
openclimatefix:mainfrom
yasumorishima:fix/runtime-defects-weathermesh-aurora

Conversation

@yasumorishima

Copy link
Copy Markdown

Pull Request

Description

Three defects that are reachable from the public API. All three surfaced while documenting these modules and were reported at the end of #245; this fixes them.

to_json raises AttributeError on all four WeatherMesh configs. WeatherMeshEncoderConfig, WeatherMeshDecoderConfig, WeatherMeshProcessorConfig and WeatherMeshConfig each implement it as dacite.asdict(self), but dacite does not export asdict — from_dict is its only converter. dataclasses.asdict is the drop-in: it recurses into the nested encoder, processor and decoder configs, so from_json reads the result straight back.

That alone is not enough to make the pair usable, though. json.dumps turns the kernel and kernel_size tuples into lists, and dacite then rejects them with WrongTypeError: wrong value type for field "kernel" - should be "tuple" instead of value "[3, 5, 5]" of type "list". For methods named to_json and from_json that seemed worth closing, so from_dict now gets config=dacite.Config(cast=[tuple]).

WeatherMesh.processors was invisible to torch. Both constructor branches assigned a plain list, so the processors never registered as submodules. On a two-processor model they contribute nothing to state_dict(), parameters() reports zero of their parameters, and .to() and .train()/.eval() do not reach them — a checkpoint saved after training silently drops every processor weight, and moving the model to a GPU leaves the processors on the CPU. nn.ModuleList registers them while keeping len() and iteration, which is all forward uses.

This does change the state_dict layout: processors.* keys now exist where they did not before, so a strict load_state_dict of a checkpoint saved by the old code will report them as missing. Those checkpoints never contained processor weights in the first place, so I think that is the right way round, but it is a visible change and worth calling out.

EarthSystemLoss.spatial_correlation_loss only worked for a batch size of one. It flattened points to (B*N, 2) before torch.cdist, which returns (B*N, B*N), then reshaped that to (B, N, N). The element counts agree only when B == 1, so anything larger raises RuntimeError: shape '[2, 4, 4]' is invalid for input of size 64. torch.cdist is already batched, so calling it on the (B, N, 2) tensor gives the intended per-element distances and stops computing cross-batch pairs that were never wanted. The existing test only ever passed batch_size = 1, which is why this went unnoticed.

Removing the batch_size, num_points, _ = points.shape unpack removed the only shape check the function had. Un-batched (N, 2) points used to raise ValueError from that unpack; against the batched cdist they broadcast instead, and when the feature count happens to equal the point count the function returns a plausible-looking number for input it cannot interpret. Restored as an explicit ValueError naming the expected shape.

Nothing that worked before changes value: the B == 1 loss is bit-identical.

How Has This Been Tested?

Eight tests added — tests/test_weathermesh.py (config round-trips, both as a dict and through real JSON; processor registration for both the supplied and the default branch) and tests/test_aurora.py (batched spatial loss, the B == 1 value pinned against the old body, cross-batch independence, the shape guard, and the full EarthSystemLoss.forward).

Each defect was reproduced on main first, then re-run after the fix:

  • to_json() → AttributeError: module 'dacite' has no attribute 'asdict' before; after, it returns a plain nested dict that from_json reads back equal for all four configs, including after a json.dumps/json.loads round-trip.
  • processors before: type list, zero processors.* keys in state_dict(), parameters() sums to 0 while the processors themselves hold 12, and model.eval() leaves them in training mode. After: four processors.* keys, 12 parameters, and eval() and to(torch.float64) both propagate.
  • spatial loss before: B=1 fine, B=2 RuntimeError. After: B=2 fine; over 20 random B=1 cases max |old - new| = 0.000e+00; and for B of 2, 3 and 5 the batched value matches the mean of the per-element losses to 1.5e-08.

Mutation check — re-introducing each defect on its own makes the new tests fail, so they are not vacuous:

reverted result
to_json back to dacite.asdict test_weathermesh_configs_round_trip fails
from_json back without the tuple cast test_weathermesh_configs_round_trip fails
supplied processors back to a list test_weathermesh_supplied_processors_are_registered fails
default processors back to a list test_weathermesh_default_processors_are_registered fails
flattened cdist back 3 aurora tests fail
shape guard removed test_spatial_correlation_loss_rejects_unbatched_points fails

Commands (Python 3.12, torch 2.10.0, dacite 1.9.2):

  • pytest tests/test_aurora.py tests/test_weathermesh.py -q → 28 passed, 4 failed
  • pytest tests/ -q (deselecting the three test_forecaster_and_loss variants, which need ~16 GB) → 190 passed, 3 skipped, 10 failed, 10 errors. The same command on a clean main in the same environment gives 9 failed and the same 10 errors.

Those failures are pre-existing: I diffed the sorted list of failing node ids before and after, and every entry matches except one. They are the four test_weathermesh cases that need a NATTEN backend unavailable on CPU, three test_gencast assertions, and test_weather_station_reader without SynopticPy installed.

The single extra failure, test_anemoi.py::test_normalization, is flaky rather than a regression: 10 isolated runs of it on this branch gave 5 passes and 5 failures. It declares stds=0.2 for U(0, 1) data whose true std is 1/sqrt(12) ≈ 0.2887, so the normalised std centres on 1.443 against a |std - 1| < 0.5 tolerance and the 12-sample noise crosses it often. It is in a module this change does not touch.

CI on this branch will still be red, for the same reason every PR here is: pytest segfaults during collection at torch_scatter/__init__.py while importing torch_geometric, before any test runs. That is #232, which #241 fixes; this branch is cut from main, so it inherits it.

To check that the new tests really do pass where the suite can run, I merged the torch pin from #241 into a throwaway branch on my fork and let the same workflow run it: ubuntu 213 passed, macOS 212 passed, 0 failed on both. #241 on its own gives 205 and 204, so all eight new tests pass on hosted runners, on both operating systems. Happy to rebase this onto #241 if you would rather see that in the PR itself — note it also touches tests/test_weathermesh.py, so the two will conflict trivially on the import block whichever order they land in.

Checklist:

  • My code follows OCF's coding style guidelines
  • I have performed a self-review of my own code
  • I have made corresponding changes to the documentation
  • I have added tests that prove my fix is effective or that my feature works
  • I have checked my code and corrected any misspellings

…imatefix#232)

The pixi environment resolves pytorch without a version constraint, but
the installpyg task pulls the PyG companion wheels from the torch 2.7.0
index. Once conda-forge moved past 2.7 the two no longer matched, so
importing torch_scatter aborted the interpreter and pytest crashed
during collection.

Pin pytorch to 2.10.* and point installpyg/installpygcuda at the
matching wheel indices so the compiled extensions share an ABI with the
installed torch.

Verified in a clean pixi environment: before the change
`python -c "import torch_scatter"` aborts (SIGABRT) and pytest dies
during collection; after it imports cleanly and the suite runs.
Review point: allowing patch updates leaves the same class of bug open.
data.pyg.org publishes a single torch-2.10.0 index for this line, so a
later 2.10.1 on conda-forge would be picked up by 2.10.* while the
wheels stayed at 2.10.0, recreating the ABI mismatch this change fixes.
Pinning to 2.10.0 makes the coupling between the two exact.

Verified after the change: the environment resolves and torch,
torch_scatter, torch_geometric and graph_weather all import.
With the segfault gone, pytest reaches collection and fails there
instead: every test module that imports the data package hits
"No module named 'nnja_ai'", 25 errors per job on both runners.

pyproject already defines an installnnja task, but the workflow never
calls it, so nnja-ai and its jsonschema dependency were absent. Adding
the step alongside the existing installpyg and installnat calls lets
collection complete: 208 tests are collected after running it in a
clean environment.
…#232)

With pytest able to collect and run again, the suite surfaces failures that
have been masked since at least June 4. None of them are caused by the torch
pin in this PR; they are pre-existing breakage on main.

torch-harmonics 0.9.0 started applying triangular truncation in truncate_sht,
clamping mmax down to lmax. generate_isotropic_noise still sized its
coefficients as (lmax, lmax + 1), so InverseRealSHT.forward asserted on the
last dimension. Read the retained mode counts back off the transform instead.
The dropped coefficients have m > l, where the associated Legendre functions
vanish, so the generated noise is unchanged.

pandas removed the uppercase offset aliases, so every freq="1H" and freq="H"
raises ValueError. Lowercased them in weather_station_reader, dataloader and
the two training entrypoints.

The WeatherMesh encoder test still expected the latent depth from before
"Fix keeping the correct depth": the pressure path uses stride=(1, 2, 2), so
the depth is the 25 pressure levels plus the surface level, not 5. The
processor, decoder and end-to-end tests asked NATTEN for a head dimension of
4, which no backend supports on CPU, and the processor's 26x32x64 grid made
the flex fallback materialise a 22.7 GB attention mask - that is what killed
the macOS job with exit 137. Kept latent_dim // num_heads >= 8 and shrank the
grids.

test_gencast_graph cross-checks against torch_geometric's TwoHop, which
multiplies two sparse CSR matrices. torch only implements that on CPU when
built against MKL, which the macOS arm64 build is not, so that comparison is
now skipped there. graph_weather's own khop builder uses torch.sparse.mm and
is still exercised everywhere.

test_normalization declared stds=0.2 for data drawn uniformly from [0, 1),
whose standard deviation is 1/sqrt(12). That put the expected normalised
spread at 1.44 against an asserted bound of 1.5, and the twelve-value sample
crossed it often enough to be flaky - locally it failed 2 runs in 5. Declared
the true standard deviation and seeded the generator.

Verified on a Linux CPU box with the pixi default environment (torch 2.10.0,
torch-harmonics 0.9.1, natten 0.21.7, pandas 3.0.5): 10 failed / 182 passed
before, 201 passed / 4 skipped after, with the three test_forecaster_and_loss
variants deselected as they need more memory than the box has. black is clean
and this adds no new ruff findings.
All three are reachable from the public API. They were found while documenting
these modules and were reported in openclimatefix#245; this change fixes them.

WeatherMeshEncoderConfig, WeatherMeshDecoderConfig, WeatherMeshProcessorConfig
and WeatherMeshConfig each implement to_json as dacite.asdict(self), but dacite
does not export asdict - from_dict is its only converter - so all four raise
AttributeError. dataclasses.asdict is the drop-in, and it recurses into the
nested encoder, processor and decoder configs, so from_json reads the result
back unchanged.

WeatherMesh kept its processors in a plain list in both constructor branches,
so torch never saw them. They were absent from state_dict(), invisible to
parameters(), and untouched by .to() and .train()/.eval(): a checkpoint saved
after training dropped every processor weight, and moving the model to a GPU
left the processors behind on the CPU. nn.ModuleList registers them while
keeping len() and iteration, which is all forward uses.

EarthSystemLoss.spatial_correlation_loss flattened points to (B*N, 2) before
torch.cdist, which returns (B*N, B*N), and then reshaped that to (B, N, N).
The element counts only agree when B is 1, so any larger batch raised
"RuntimeError: shape '[2, 4, 4]' is invalid for input of size 64". torch.cdist
is already batched, so calling it on the (B, N, 2) tensor gives the intended
per-element distances and stops computing cross-batch pairs that were never
wanted. The B == 1 value is bit-identical to before, and the existing loss test
only ever used a batch size of one, which is why this went unnoticed.

Each defect gets a test: the config round-trips, processor registration for
both the supplied and the default branch, and the batched loss checked against
the mean of the per-element losses. Reverting any one fix on its own fails at
least one of them.
…ad shapes

Two gaps found in self-review of the previous commit.

to_json no longer raises, but the pair still could not survive actual JSON:
json.dumps turns the tuple fields into lists, and dacite then rejects them
with WrongTypeError: wrong value type for field "kernel" - should be "tuple"
instead of value "[3, 5, 5]" of type "list". Methods called to_json and
from_json should round-trip through JSON, so from_dict now gets
config=dacite.Config(cast=[tuple]) and the test goes through
json.loads(json.dumps(...)) rather than handing the dict straight back.

Dropping the (batch, num_points, _) unpack also dropped the only shape check
spatial_correlation_loss had. Un-batched (N, 2) points used to raise
ValueError from the unpack; with the batched cdist they broadcast instead, and
when the feature count happens to equal the point count the function returns a
plausible-looking number for input it cannot interpret. Restored as an explicit
ValueError naming the expected shape.
…o pytest can run

pytest segfaults on main because pixi resolves torch 2.10.0 while installpyg fetches torch-2.7.0 wheels, so torch_scatter aborts at import. openclimatefix#241 fixes that; this branch needs it to get a green run.

The only conflict was at the top of tests/test_weathermesh.py: this branch added an import, openclimatefix#241 added the module docstring. Kept both, docstring first.
@yasumorishima

Copy link
Copy Markdown
Author

pytest was failing here for the same reason as #232 — torch_scatter aborts at import because pixi resolves torch 2.10.0 while installpyg fetches torch-2.7.0 wheels. That is fixed in #241, so I merged its head into this branch to get a run that actually exercises the changes: ubuntu 213 passed / 3 skipped, macOS 212 passed / 4 skipped, 0 failed.

The merge commit is only here to make CI meaningful — happy to rebase it away once #241 lands, or to drop it if you would rather review this branch against main as-is.

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

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant