Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions .github/workflows/workflows.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,7 @@ jobs:
pixi run -e ${{ matrix.environment }} installpyg
pixi run -e ${{ matrix.environment }} pip install coverage==7.4.3 pytest-cov
pixi run -e ${{ matrix.environment }} installnat
pixi run -e ${{ matrix.environment }} installnnja
pixi run -e ${{ matrix.environment }} install
- name: Setup with pytest-cov
run: |
Expand Down
4 changes: 2 additions & 2 deletions graph_weather/data/dataloader.py
Original file line number Diff line number Diff line change
Expand Up @@ -99,7 +99,7 @@ def __getitem__(self, item):
np.cos(day_of_year)
solar_times = [np.array([extraterrestrial_irrad(date, lat, lon) for lat, lon in lat_lons])]
for when in pd.date_range(
date - pd.Timedelta("12 hours"), date + pd.Timedelta("12 hours"), freq="1H"
date - pd.Timedelta("12 hours"), date + pd.Timedelta("12 hours"), freq="1h"
):
solar_times.append(
np.array([extraterrestrial_irrad(when, lat, lon) for lat, lon in lat_lons])
Expand All @@ -112,7 +112,7 @@ def __getitem__(self, item):
np.array([extraterrestrial_irrad(end_date, lat, lon) for lat, lon in lat_lons])
]
for when in pd.date_range(
end_date - pd.Timedelta("12 hours"), end_date + pd.Timedelta("12 hours"), freq="1H"
end_date - pd.Timedelta("12 hours"), end_date + pd.Timedelta("12 hours"), freq="1h"
):
end_solar_times.append(
np.array([extraterrestrial_irrad(when, lat, lon) for lat, lon in lat_lons])
Expand Down
4 changes: 2 additions & 2 deletions graph_weather/data/weather_station_reader.py
Original file line number Diff line number Diff line change
Expand Up @@ -681,14 +681,14 @@ def interpolate_missing_data(
return interpolated

def resample_observations(
self, observations: Optional[xr.Dataset], freq: str = "1H", aggregation: str = "mean"
self, observations: Optional[xr.Dataset], freq: str = "1h", aggregation: str = "mean"
) -> Optional[xr.Dataset]:
"""
Resample observations to a different time frequency.

Args:
observations: Dataset with observations.
freq: Target frequency ('1H', '1D', etc.).
freq: Target frequency ('1h', '1D', etc.).
aggregation: Aggregation method ('mean', 'sum', 'min', 'max').

Returns:
Expand Down
11 changes: 8 additions & 3 deletions graph_weather/models/gencast/utils/noise.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,12 +38,17 @@ def generate_isotropic_noise(num_lon: int, num_lat: int, num_samples=1, isotropi
if isotropic:
lmax = num_lat - 1 if extend else num_lat
mmax = lmax + 1
coeffs = torch.randn(num_samples, lmax, mmax, dtype=torch.complex64) / np.sqrt(
(num_lat**2) // 2
)
isht = th.InverseRealSHT(
nlat=num_lat, nlon=num_lon, lmax=lmax, mmax=mmax, grid="equiangular"
)
# torch-harmonics >= 0.9.0 applies triangular truncation, which clamps mmax down
# to lmax, so the transform can keep fewer modes than were requested. Read the
# retained mode counts back off the transform instead of assuming lmax/mmax are
# honoured verbatim. The discarded coefficients have m > l, where the associated
# Legendre functions vanish, so this does not change the generated noise.
coeffs = torch.randn(num_samples, isht.lmax, isht.mmax, dtype=torch.complex64) / np.sqrt(
(num_lat**2) // 2
)
noise = isht(coeffs) * np.sqrt(2 * np.pi)
noise = einops.rearrange(noise, "b lat lon -> lon lat b").numpy()
else:
Expand Down
42 changes: 42 additions & 0 deletions graph_weather/models/mosaic/README.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,42 @@
# MOSAIC

Unofficial implementation of components from
[(Sparse) Attention to the Details: Preserving Spectral Fidelity in ML-based Weather Forecasting Models](https://arxiv.org/abs/2604.16429)
(Zhdanov et al.).

This subpackage provides the two parts requested in issue #217: block-sparse
attention and processing at the native grid resolution. It is not a
reproduction of the full forecaster.

## Components

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

## Token ordering

Block-sparse attention assumes contiguous index blocks are spatial
neighbourhoods. On a HEALPix grid this holds in NESTED ordering, where pixel
`p` subdivides into children `4p ... 4p+3` (Section 3.2). For irregular point
sets, `healpix.nested_order` returns a permutation that provides the same
property; any other locality-preserving ordering works as well, so the
attention module never requires `healpy`.

## Scope

Implemented: the attention mechanism, grid transfer, hierarchical pooling and
a processor that can be used standalone.

Not implemented: the full forecaster and its 82 channel input/output, CRPS
training, ensemble noise injection, the Triton kernel used for large
sequences on GPU, and pretrained weights.

## Provenance

Written from the paper text and equations only. The reference implementation
is published without a license, so no code from it was consulted.
26 changes: 26 additions & 0 deletions graph_weather/models/mosaic/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,26 @@
"""MOSAIC block-sparse attention and native-grid processing.

Modules from "(Sparse) Attention to the Details: Preserving Spectral
Fidelity in ML-based Weather Forecasting Models" (arXiv:2604.16429).

Everything here is pure PyTorch and operates on point sets of shape
(batch, n_tokens, dim), so the components can be reused by the graph-based
models in this repository. See README.md in this directory for scope.
"""

from .block_sparse_attention import BlockSparseAttention, RotaryEmbedding2D
from .coarsen import HealpixCoarsen, HealpixRefine
from .interpolate import CrossAttentionInterpolator, knn_indices
from .layers import MosaicProcessor, MosaicTransformerBlock, SwiGLU

__all__ = [
"BlockSparseAttention",
"CrossAttentionInterpolator",
"HealpixCoarsen",
"HealpixRefine",
"MosaicProcessor",
"MosaicTransformerBlock",
"RotaryEmbedding2D",
"SwiGLU",
"knn_indices",
]
Loading
Loading