feat: add MOSAIC block-sparse attention and native-grid processing (#217) - #242
yasumorishima wants to merge 9 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.
…penclimatefix#217) Adds the two components requested in the issue: block-sparse attention and processing at the native grid resolution, from "(Sparse) Attention to the Details: Preserving Spectral Fidelity in ML-based Weather Forecasting Models" (arXiv:2604.16429). BlockSparseAttention implements the three branches of Section 4.2 and Eq. 8-11: block-level mean pooling with dense attention between block representations, top-n key block selection shared by all queries in a block, and full attention within each block. The branch outputs are combined by learned gating (Eq. 2). Everything operates on point sets of shape (batch, n_tokens, dim) rather than images, so the modules compose with the graph-based models here. Blocks are contiguous index ranges, which is a spatial neighbourhood under HEALPix NESTED ordering; healpix.nested_order provides the same property for irregular point sets and is the only place healpy is used, behind a guarded import. No new required dependencies: the implementation is pure PyTorch. Written from the paper text and equations only. The reference implementation is published without a license, so no code from it was consulted.
The fine-grained branch gathered from a view expanded to (n_blocks, n_blocks); the backward pass of gather allocates a buffer with that shape, so training memory grew quadratically with the sequence length and defeated the point of sparse attention. Measured with resource.ru_maxrss at dim=256, block=128: doubling from 4096 to 8192 tokens took the forward-plus-backward delta from 265MB to 790MB (3.0x). Selecting with advanced indexing instead keeps the gradient buffer the size of the key tensor: the same measurement now reads 241MB to 413MB (1.7x). A regression test asserts the growth stays sub-quadratic. Also from review: - The padding test perturbed a real token, so it never checked that padded positions stay out of the result. It now poisons the internal padding and asserts real-token outputs are bit-identical. - The dense-equivalence test only covered top_n == n_blocks, where the softmax is permutation invariant and any bijective index error still passes. Added a sparse case checked against a naive loop reference. - HealpixCoarsen and HealpixRefine now require positions when built with use_positions=True, so the projection can never sit unused. - MosaicProcessor validates the token count up front and names both the input size and the required divisor. - knn_indices computes distances in chunks instead of materialising the full pairwise matrix. - Documented that axial rotary embeddings are discontinuous across the date line and at the poles, and that block selection is shared across heads.
Review follow-ups: - CrossAttentionInterpolator normalised the relative position with norm().clamp(min=1e-6). Clamping the norm does not tame the derivative, which passes through norm() at zero: with a target sitting exactly on a source point the gradient reached 3.9e4 on torch 2.10. Softening the length under the square root instead keeps it below 1e3, and a test covers the coincident-point case. A NaN was reported for this path but does not reproduce on torch 2.10, which returns a large finite value. - Documented that the compression branch attends between all block representations and is therefore quadratic in the block count, with the guidance to grow block_size alongside the token count. The selection and local branches remain linear. - healpix helpers now reject an nside that is not a power of two instead of letting healpy fail later; NESTED ordering requires one. - test_interpolator_reproduces_constant_field asserted only the shape and finiteness despite its name. It now feeds a constant field and checks the convex combination reproduces it.
…tions A second review pass showed the previous round left two problems. The memory regression test had no power: replaying it against the old quadratic implementation passed three times out of three, because at 2048 and 4096 tokens the linear terms dominate the process footprint. It now measures the largest single allocation during forward and backward with the autograd profiler and compares it against the size of the (n_blocks, n_blocks) buffer, which is what the old formulation allocated. Restoring the gather formulation makes the new test fail, so it detects the regression it is named for. The interpolator normalised relative positions with an epsilon added under the square root. That bounded the gradient but shrank short vectors: neighbouring points on a 0.25 degree grid lost 2.5 percent of their length, and at 1e-3 apart the direction was no longer unit length at all. Points closer than sqrt(eps) are now treated as coincident and given a zero direction while everything else is normalised exactly, so distortion is zero at every separation measured and the gradient at a coincident point is zero rather than 3.9e4. The padding test only poisoned the feature tensor; coordinates are padded along a different axis, so the rotary path was never covered. It now locates padded rows from the shape change and poisons both.
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.
|
CI on this branch fails with the segfault described in #232 — it dies while importing torch_scatter, before any test runs, exactly as it does on main. Nothing here can turn it green on its own. #241 fixes that. On that branch CI reaches the suite and reports 187 passed, 8 failed. Once it lands I will rebase this branch so you can see these tests run on your own infrastructure rather than taking my local numbers on faith. For what it is worth, locally in an environment with the matched wheels: main is 182 passed / 10 failed and this branch is 198 passed / 10 failed, with an identical failing set. |
…#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.
|
Note: this branch now includes the commits from #241 (PyTorch pin + the test fixes it unblocked). Without them the CI job here segfaults during collection, so the pytest checks on this PR could not run at all. With #241 merged in, both jobs are green on this branch: run 31069305350 — ubuntu-latest 221 passed / 3 skipped, macos-latest 220 passed / 4 skipped. As a result the diff temporarily shows #241's files as well. Once #241 lands I'll rebase this branch onto |
Addresses #217 — the block-sparse attention and native-grid processing parts of (Sparse) Attention to the Details, not a reproduction of the full forecaster.
What is here
graph_weather/models/mosaic/:BlockSparseAttentionRotaryEmbedding2DCrossAttentionInterpolatorHealpixCoarsen/HealpixRefineMosaicTransformerBlock/MosaicProcessorEverything takes point sets of shape
(batch, n_tokens, dim)rather than images, so it composes with the graph-based models here. Blocks are contiguous index ranges, which is a spatial neighbourhood under HEALPix NESTED ordering;healpix.nested_ordergives the same property for irregular point sets and is the only place healpy is used, behind a guarded import. No new required dependencies — the implementation is pure PyTorch.Testing
Same environment and command, only the branch differs:
The failing set is identical on both (the sorted
FAILEDlists diff empty), so nothing existing is affected; the 16 extra passes aretests/test_mosaic.py. Three tests intest_model.pythat exceed 8GB were deselected in both runs — see #241, which this branch does not depend on but which is what currently keeps CI from running at all.Two of the tests are worth calling out because they were written after the first versions of them turned out to be worthless:
test_selection_branch_never_materialises_the_block_squaremeasures the largest single allocation during forward and backward. An earlier version compared process memory at two token counts and passed against the quadratic implementation three times out of three, because the linear terms dominate at test scale. Restoring the old formulation makes the current test fail.test_padding_never_influences_real_tokenspoisons the internal padding and asserts real-token outputs are unchanged. The first version perturbed a real token, so it checked nothing.Choices worth flagging
block_sizewith the token count.MosaicProcessorbuilds the pooling layers without their relative-position term, since it carries features only. CallingHealpixCoarsendirectly with positions uses the full Eq. 12.Provenance
Written from the paper text and equations only. The reference implementation is published without a license, so no code from it was consulted and no numerical parity check against it was possible; the equivalence test against a naive loop reference stands in for that.
Happy to split this up or drop pieces if you would rather land it in smaller parts.