Skip to content

feat: add MOSAIC block-sparse attention and native-grid processing (#217) - #242

Open
yasumorishima wants to merge 9 commits into
openclimatefix:mainfrom
yasumorishima:feature/217-mosaic-block-sparse-attention
Open

yasumorishima wants to merge 9 commits into
openclimatefix:mainfrom
yasumorishima:feature/217-mosaic-block-sparse-attention

Conversation

@yasumorishima

Copy link
Copy Markdown

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/:

Module Paper Purpose
BlockSparseAttention Eq. 2, 8-11, §4.2 Compression, block-level top-n selection, and within-block local attention, combined by learned gating
RotaryEmbedding2D §4.3 2D axial rotary embedding over (lat, lon)
CrossAttentionInterpolator Eq. 6-7, §4.1 Moves features between point sets using relative-position queries
HealpixCoarsen / HealpixRefine Eq. 12-13, §4.3 Learnable pooling and unpooling of four sibling pixels
MosaicTransformerBlock / MosaicProcessor Eq. 14, §4.3 Pre-norm block, and a processor whose first stage runs at native resolution

Everything 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_order gives 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:

branch passed failed
main (a23d8c5) 182 10
this branch 198 10

The failing set is identical on both (the sorted FAILED lists diff empty), so nothing existing is affected; the 16 extra passes are tests/test_mosaic.py. Three tests in test_model.py that 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_square measures 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_tokens poisons the internal padding and asserts real-token outputs are unchanged. The first version perturbed a real token, so it checked nothing.

Choices worth flagging

  • The paper does not state the gate nonlinearity for Eq. 2. I used a per-token sigmoid per branch and said so in the docstring rather than implying the paper prescribes it.
  • Block selection is shared across heads, which the docstring states.
  • The compression branch is quadratic in the block count, as in the paper's own cost expression. That is documented on the class along with the guidance to grow block_size with the token count.
  • MosaicProcessor builds the pooling layers without their relative-position term, since it carries features only. Calling HealpixCoarsen directly 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.

…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.
@yasumorishima

Copy link
Copy Markdown
Author

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.
@yasumorishima

Copy link
Copy Markdown
Author

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 main so the diff is limited to the MOSAIC changes.

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