Skip to content

Add WeatherNext 2 weather forecasting model - #47874

Open
kashif wants to merge 104 commits into
huggingface:mainfrom
kashif:add-weathernext2
Open

kashif wants to merge 104 commits into
huggingface:mainfrom
kashif:add-weathernext2

Conversation

@kashif

@kashif kashif commented Aug 10, 2026 •

Copy link
Copy Markdown
Contributor

CPU CI GPU run-slow

Adds WeatherNext 2 (FGN), Google DeepMind's global medium-range weather model. This is the first weather model in the library, so it also adds a weather-forecasting task and AutoModelForWeatherForecasting.

The model is an encode-process-decode graph network: the lat/lon grid is encoded, projected onto an icosahedral mesh, processed by a transformer with attention restricted to a k-hop neighbourhood on that mesh, and projected back. It's made probabilistic in an unusual way — one 32-dim noise vector per ensemble member modulates the scale and offset of every norm layer, so an ensemble is just several noise draws (the batch dimension here).

Two things worth a look during review:

  • The mesh, both grid↔mesh graphs and the attention mask are deterministic functions of the config, so they're built at init into non-persistent buffers and cached on disk rather than shipped in the checkpoint. Mesh, RCM permutation, k-hop mask and ball-query edges are bit-identical to the original.
  • After the RCM ordering the attention mask is banded, so attention runs over three block-diagonals instead of the full node set. That keeps a mask that would be 1.7 GB dense at 0.25° down to 1.2 GB and avoids materializing the scores. eager, sdpa and flex_attention agree to 3e-8; Flash Attention can't express this mask.

Checked against the original JAX/Haiku implementation with fixed noise on the published sample data — normalized inputs match to 1e-6 and predictions correlate at 0.999+, though there's still a ~0.04% output difference I haven't tracked down yet.

Converted checkpoints run end-to-end and beat 6h persistence comfortably (0.25°, verified against the HRES analysis):

field model persistence
2m_temperature 0.796 K 2.622 K
mean_sea_level_pressure 47.3 Pa 260 Pa
geopotential @500hPa 25.4 m²/s² 224.5 m²/s²

Still to do: push the converted weights to the Hub, and add slow integration tests pinned to them.

🤖 mlinter review state

kashif added 4 commits August 10, 2026 11:03
Ports Google DeepMind's WeatherNext 2 (FGN) to PyTorch: encode-process-decode
graph network over an icosahedral mesh, made probabilistic by a single global
noise vector that modulates every norm layer.

Adds a weather-forecasting task and AutoModelForWeatherForecasting, since this
is the first weather model in the library.
The fiddle configs store it as an int, which tripped the strict float|None
validation. Only the 0.25 degree checkpoints set it, so the mini model never
hit this.
scipy is only needed to build the mesh at init, so importing it at module
level broke the doc build. Also swaps @DataClass below @auto_docstring, and
points the paper link at hf.co/papers.
Adds PipelineTesterMixin, a create_and_check_model on the tester, and a slow
integration test against the released checkpoint.
@HuggingFaceDocBuilderDev

Copy link
Copy Markdown

The docs for this PR live here. All of your documentation changes will be reflected on that endpoint. The docs are available until 30 days after the last update.

kashif added 7 commits August 10, 2026 12:07
It imports scipy at module level, so exporting it made `from transformers
import *` require scipy.
check_repo wants every public object in the docs, and every model in at least
one auto mapping.
Adds a checkpoint link to the config docstring, which check_config_docstrings
requires, and switches the examples and the slow test to the 1 degree Mini so
they run on modest hardware.
They are read through the channel-layout properties rather than directly in
the modeling file, which is what check_config_attributes scans for.
- Fix base_model_tp_plan: the paths were copied from Llama and matched nothing,
  which is what the test I had skipped was telling me.
- Wire output capturing the way the library does it, so output_attentions
  works, and let the head propagate what the base model recorded.
- Rename WeatherNext2Processor to WeatherNext2FeatureExtractor: it subclasses
  FeatureExtractionMixin, and it was the only such class named Processor.
- Override test_attention_outputs instead of skipping it, and give every
  remaining skip a precise reason.

@Rocketknight1 Rocketknight1 left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

There are some issues here that are the same ones I encountered with ESMFold2! I used that PR as a reference since they're both science models with unusual modalities. I commented in a few locations, but other issues are:

  1. Rebuilding masks every layer (can we precompute once and just reuse it?)
  2. No tests for the feature extractor
  3. Do we document anywhere where users can actually get weather data for this? Since it's an unusual modality users might have no idea how to actually use it.
  4. Are we careful to either actually support a batch dim, or raise an error if it's passed? What we want to avoid is looking like we half-support it in places, but crashing because it's untested/unsupported in reality. Raising an error if there's a batch dim (or a batch dim > 1) might be good

Comment thread docs/source/en/model_doc/weathernext2.md Outdated
Comment thread docs/source/en/model_doc/weathernext2.md
Comment thread src/transformers/models/auto/modeling_auto.py Outdated
Comment on lines +43 to +45
| [kashif/weathernext2-mini](https://huggingface.co/kashif/weathernext2-mini) | 1° | 56.7M | runs on modest hardware; `main` and `<2023` revisions |
| [kashif/weathernext2](https://huggingface.co/kashif/weathernext2) | 0.25° | 183.8M | also predicts 100m winds |
| [kashif/weathernext-cyclones](https://huggingface.co/kashif/weathernext-cyclones) | 0.25° | 183.8M | operational cyclone model; `main`, `<2024`, `<2023` revisions |

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

We'll definitely have to coordinate and move these to GDM or the right owner before launch

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Whats the state here btw?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Code and validation are ready. The Hub transfer is still outstanding and needs coordination with GDM/the final owner before launch, so I am leaving this thread open.

Comment thread src/transformers/models/weathernext2/configuration_weathernext2.py Outdated
Comment thread docs/source/en/model_doc/weathernext2.md Outdated
Comment thread src/transformers/models/weathernext2/geometry_weathernext2.py Outdated
Comment thread src/transformers/models/weathernext2/geometry_weathernext2.py Outdated

@Rocketknight1 Rocketknight1 left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

It's looking better! I added a few more comments. Also, do we have any integration testing that checks actual values, and is there any support for bf16? If so, do we need to make sure some models are added to the keep_in_fp32 attributes?

@Raymondlol

Copy link
Copy Markdown

Hi @kashif, great to see WeatherNext 2 coming to Transformers. I've been working on the JAX side (faster-weathernext: the official code, patched to run at 0.25° on consumer GPUs). Two things that might help here:

The ~0.04% output difference. It's worth checking which precision the JAX reference ran at. On Ampere and newer GPUs, JAX's default matmul precision uses TF32; on TPU it uses bf16 passes. With the official WeatherNext 2 modules on the 0.25° Cyclones checkpoint, two implementations that agree to 2e-5 relative RMS at jax.default_matmul_precision("highest") differed by up to 4.7e-3 at the default precision (numbers). So a 4e-4 gap could come from the reference side. If the reference ran on CPU, this doesn't apply. Either way, I can dump strict-fp32 JAX activations for the sample input and a fixed noise draw to help bisect: the encoder output, each of the 24 transformer layers, and the decoder output.

Memory (~50 GB per member in the model card). In the JAX code, the peak came from per-edge tensors in the grid↔mesh GNNs: the mesh→grid graph has 3.1 M edges, and each per-edge tensor is about 8.9 GiB. Processing edges and receivers in blocks, plus running the 24 layers as a loop, brought one step down to about 6.4 GiB, with results unchanged beyond rounding. The decoder blocks cleanly because each grid point has exactly 3 incoming edges, from its enclosing mesh triangle. The same blocking should carry over to the PyTorch port; there's a write-up in how-it-works.md, and I'm happy to share details.

@kashif

kashif commented Sep 29, 2026

Copy link
Copy Markdown
Contributor Author

thanks @Raymondlol checking!

@kashif

kashif commented Oct 2, 2026

Copy link
Copy Markdown
Contributor Author

@Raymondlol i have added some of your suggestions here now if you want to check

github-actions Bot pushed a commit to huggingface/kernels-community that referenced this pull request Oct 5, 2026
## Related issue

Follow-up to #1091; companion model PR:
huggingface/transformers#47874.
No separate issue was opened for this follow-up.

## What does this PR do?

Adds chunked grid encoding and the forecast head to the existing
WeatherNext2 kernel. The approach is inspired by
[Faster-WeatherNext](https://github.com/Raymondlol/Faster-WeatherNext).

## Motivation

The full 0.25° model creates large grid intermediates. Processing them
in blocks makes inference fit more comfortably without changing the
weights or precision.

## Changes

- Add inference-only forward replacements and bump the API to v2.
- Use Triton for conditioning/output writes and shifted-sigmoid
selection. Keep PyTorch LayerNorm to preserve its rounding.
- Keep the original forwards for training, gradients, autocast and lower
precision.

## Testing

- CPU and Radeon 890M/ROCm: 31 kernel tests pass; Ruff passes.
- The Mini forecast matches the original forward within 1e-4. Original
JAX component comparisons also pass at that tolerance.
- With the actual 0.25° checkpoint weights, encoder additional peak
allocation drops from 12.64 to 3.28 GiB, and the head from 5.94 to 0.60
GiB. These are separate component measurements with synthetic inputs,
not a full-forecast peak.
- Full 0.25° forecast validation, CUDA/XPU testing and the production
build are still pending, so keeping this draft.

## Checklist

- [ ] This PR is linked to an issue that was discussed and approved
- [x] I have tested these changes locally
- [x] New/changed functionality has test coverage
- LLM disclosure:
  - [ ] I did not use an LLM to create this PR.
  - [x] I used an LLM for assistance while creating this PR.
  - [ ] This PR was mostly or completely generated by an LLM.
@github-actions

github-actions Bot commented Oct 9, 2026

Copy link
Copy Markdown
Contributor

[For maintainers] Suggested jobs to run (before merge)

run-slow: auto, weathernext2

@github-actions

github-actions Bot commented Oct 9, 2026

Copy link
Copy Markdown
Contributor

CI recap

Dashboard: View test results in Grafana
Latest run: 37911272135:1
Result: failure | Jobs: 16 | Tests: 198,544 | Failures: 2 | Duration: 15h 58m

@Raymondlol

Copy link
Copy Markdown

Thanks @kashif, and for the credit in the kernel docs!

Since CUDA testing was still open in huggingface/kernels-community#1197, I tried kernel v2 on an RTX 5090. Setup: 0.25°, fp32, batch 1, torch 2.14; PR at f7ccc35, kernel at db674afe.

step peak
sdpa 4.2 s 17.5 GiB
use_kernels=True (v2 as released) 40.7 s 15.1 GiB
v2 with tf32x3 instead of ieee 2.5 s 15.1 GiB

It looks like precision="ieee" in layers.py is the cause. On NVIDIA, Triton's IEEE dot doesn't use tensor cores; on ROCm IEEE is the default anyway, so it wouldn't show up there. With tf32x3, the outputs differ from sdpa by 1.1e-5 max rel RMS, about as much as two sdpa runs differ from each other (1.2e-5). Maybe the kernel could use tf32x3 on CUDA when torch.get_float32_matmul_precision() is "highest"?

One smaller thing, in case it's useful. With the kernel, the peak moves to the grid latents in grid→mesh. Under no_grad, three changes brought it to 9.5 GiB in fp32 and 5.9 GiB in bf16 here, with no change in speed or outputs:

  • update grid_states in place;
  • drop projected_senders after the edge loop;
  • write the MLP chunks into a preallocated output instead of torch.cat.

Happy to share the small patches I used for both, if that helps.

This branch has not been deployed

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

Projects

None yet

Development

Successfully merging this pull request may close these issues.

6 participants