Repository navigation
Conversation
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.
|
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. |
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
left a comment
There was a problem hiding this comment.
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:
- Rebuilding masks every layer (can we precompute once and just reuse it?)
- No tests for the feature extractor
- 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.
- 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
| | [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 | |
There was a problem hiding this comment.
We'll definitely have to coordinate and move these to GDM or the right owner before launch
There was a problem hiding this comment.
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.
Rocketknight1
left a comment
There was a problem hiding this comment.
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?
|
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 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. |
|
thanks @Raymondlol checking! |
|
@Raymondlol i have added some of your suggestions here now if you want to check |
## 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.
|
[For maintainers] Suggested jobs to run (before merge) run-slow: auto, weathernext2 |
CI recapDashboard: View test results in Grafana |
|
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.
It looks like One smaller thing, in case it's useful. With the kernel, the peak moves to the grid latents in grid→mesh. Under
Happy to share the small patches I used for both, if that helps. |
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-forecastingtask andAutoModelForWeatherForecasting.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:
eager,sdpaandflex_attentionagree 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):
Still to do: push the converted weights to the Hub, and add slow integration tests pinned to them.
🤖 mlinter review state
e80a3732dfbec6bc