fix: repair three runtime defects in the WeatherMesh and Aurora models - #246
Open
yasumorishima wants to merge 7 commits into
Open
yasumorishima wants to merge 7 commits into
yasumorishima wants to merge 7 commits into
Conversation
…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.
Author
|
pytest was failing here for the same reason as #232 — 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. |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
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_jsonraisesAttributeErroron all four WeatherMesh configs.WeatherMeshEncoderConfig,WeatherMeshDecoderConfig,WeatherMeshProcessorConfigandWeatherMeshConfigeach implement it asdacite.asdict(self), butdacitedoes not exportasdict—from_dictis its only converter.dataclasses.asdictis the drop-in: it recurses into the nested encoder, processor and decoder configs, sofrom_jsonreads the result straight back.That alone is not enough to make the pair usable, though.
json.dumpsturns thekernelandkernel_sizetuples into lists, anddacitethen rejects them withWrongTypeError: wrong value type for field "kernel" - should be "tuple" instead of value "[3, 5, 5]" of type "list". For methods namedto_jsonandfrom_jsonthat seemed worth closing, sofrom_dictnow getsconfig=dacite.Config(cast=[tuple]).WeatherMesh.processorswas invisible to torch. Both constructor branches assigned a plainlist, so the processors never registered as submodules. On a two-processor model they contribute nothing tostate_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.ModuleListregisters them while keepinglen()and iteration, which is allforwarduses.This does change the
state_dictlayout:processors.*keys now exist where they did not before, so a strictload_state_dictof 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_lossonly worked for a batch size of one. It flattenedpointsto(B*N, 2)beforetorch.cdist, which returns(B*N, B*N), then reshaped that to(B, N, N). The element counts agree only whenB == 1, so anything larger raisesRuntimeError: shape '[2, 4, 4]' is invalid for input of size 64.torch.cdistis 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 passedbatch_size = 1, which is why this went unnoticed.Removing the
batch_size, num_points, _ = points.shapeunpack removed the only shape check the function had. Un-batched(N, 2)points used to raiseValueErrorfrom that unpack; against the batchedcdistthey 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 explicitValueErrornaming the expected shape.Nothing that worked before changes value: the
B == 1loss 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) andtests/test_aurora.py(batched spatial loss, theB == 1value pinned against the old body, cross-batch independence, the shape guard, and the fullEarthSystemLoss.forward).Each defect was reproduced on
mainfirst, then re-run after the fix:to_json()→AttributeError: module 'dacite' has no attribute 'asdict'before; after, it returns a plain nesteddictthatfrom_jsonreads back equal for all four configs, including after ajson.dumps/json.loadsround-trip.list, zeroprocessors.*keys instate_dict(),parameters()sums to 0 while the processors themselves hold 12, andmodel.eval()leaves them in training mode. After: fourprocessors.*keys, 12 parameters, andeval()andto(torch.float64)both propagate.B=1fine,B=2RuntimeError. After:B=2fine; over 20 randomB=1casesmax |old - new| = 0.000e+00; and forBof 2, 3 and 5 the batched value matches the mean of the per-element losses to1.5e-08.Mutation check — re-introducing each defect on its own makes the new tests fail, so they are not vacuous:
to_jsonback todacite.asdicttest_weathermesh_configs_round_tripfailsfrom_jsonback without the tuple casttest_weathermesh_configs_round_tripfailstest_weathermesh_supplied_processors_are_registeredfailstest_weathermesh_default_processors_are_registeredfailscdistbacktest_spatial_correlation_loss_rejects_unbatched_pointsfailsCommands (Python 3.12, torch 2.10.0, dacite 1.9.2):
pytest tests/test_aurora.py tests/test_weathermesh.py -q→ 28 passed, 4 failedpytest tests/ -q(deselecting the threetest_forecaster_and_lossvariants, which need ~16 GB) → 190 passed, 3 skipped, 10 failed, 10 errors. The same command on a cleanmainin 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_weathermeshcases that need a NATTEN backend unavailable on CPU, threetest_gencastassertions, andtest_weather_station_readerwithout 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 declaresstds=0.2forU(0, 1)data whose true std is1/sqrt(12) ≈ 0.2887, so the normalised std centres on 1.443 against a|std - 1| < 0.5tolerance and the 12-sample noise crosses it often. It is in a module this change does not touch.ruff check .with the pinnedv0.15.15→ 156 errors both before and after this change, so nothing is added (that backlog is pre-commit.ci has been red on main since at least 2026-07-13 (156 ruff errors, none auto-fixable) #244)black --line-length 100 --check .→ 115 files unchangedCI on this branch will still be red, for the same reason every PR here is:
pytestsegfaults during collection attorch_scatter/__init__.pywhile importingtorch_geometric, before any test runs. That is #232, which #241 fixes; this branch is cut frommain, 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: