Strata is a deep-learning weather emulator for the SCREAM global atmosphere model. It trains transformer-based neural networks to emulate SCREAM atmospheric physics on the cubed-sphere grid and supports multi-day global forecasting.
The Strata architecture (a two-stage 3D neighborhood-attention transformer
with stereographic rotary position embeddings) lives in
NVIDIA PhysicsNeMo as
physicsnemo.experimental.models.strata. This repository adds the
SCREAM-specific pieces around it: spherical tile geometry, tile-local wind
rotation, data pipelines, training configs, and rollout/inference tooling
(see screamcast/strata_wrappers.py).
The public release is named Strata, but the Python package, command-line
scripts, and environment variables keep the project's original name: you import
screamcast, run the screamcast-prefixed entry points, and configure
SCREAM_* variables. Strata and screamcast refer to the same project.
This code is provided for research and development purposes only.
The models need a recent NVIDIA GPU software stack that cannot come from PyPI
alone: a CUDA-tuned PyTorch build, NVIDIA DALI (data pipelines), and a NATTEN
build that matches the installed torch. The supported
base environment is the NGC PyTorch container
nvcr.io/nvidia/pytorch:26.01-py3 (anonymous pulls generally work; a free
NGC account may be required depending on NGC policy).
Build the training/inference image from the repository root:
docker build -f docker/Dockerfile -t strata .
The image contains the complete environment, including a from-source NATTEN
build (set CUDA_ARCH in docker/build_natten.sh
for non-Hopper GPUs) and PhysicsNeMo pinned to a commit that includes the
Strata models.
Outside NGC containers the package also installs as a normal Python project. The one extra flag is NATTEN's wheel index, which lets pip pick the CUDA wheel matching your torch build (without it, pip falls back to the PyPI sdist and a long from-source compile):
pip install -e . -f https://whl.natten.org
Optional extras: .[s3] (S3 data access), .[analysis] (plotting),
.[cu12]/.[cu13] (DALI for your CUDA major version), .[dev] (tests),
.[notebooks] (marimo). earth2grid installs automatically from its pinned
GitHub archive (it has no PyPI release). Note this path resolves the newest
compatible versions from PyPI rather than reproducing the container
environment above — use the Docker image for the verified stack.
Create a .env file at the repository root that points to your data and output
locations. Copy the template and fill in the values:
cp envs/example.env .env
The key variables are PROJECT_ROOT (training and rollout outputs are written
here), ZARR_ROOT (the SCREAM zarr dataset), AUX_DATA_ROOT (auxiliary files,
below), and WANDB_API_KEY (experiment logging).
The SCREAM training dataset (cubed-sphere zarr v3 stores) is published on Hugging Face as three dataset repositories — a 14-day simulation (sdecadal) plus two 7-day simulations (sdy1, sdy2) used as the additional sources in the three-source training configs:
- nvidia/STRATA-SCREAM-sdecadal — 14 days, 10-minute steps
- nvidia/STRATA-SCREAM-sdy1 — 7 days, 10-minute steps
- nvidia/STRATA-SCREAM-sdy2 — 7 days, 10-minute steps
Each repository is the zarr store itself, so it can be read directly over
HTTP via huggingface_hub's fsspec filesystem (installed in the Docker
image) — convenient for exploration without downloading ~TB of data:
import xarray as xr
ds = xr.open_zarr("hf://datasets/nvidia/STRATA-SCREAM-sdy1")
ds["T_2m"].isel(time=500, ncol=slice(0, 100)).values # streams just the chunks it needsFor training, download the stores instead (e.g. with huggingface-cli download) — the dataloader reads tiles at a rate that HTTP streaming cannot
sustain — and place each under ZARR_ROOT with the directory name the
training configs expect (see _ZARR_SDECADAL / _ZARR_SDY1 / _ZARR_SDY2
in train_configs.py), e.g.
$ZARR_ROOT/sdecadal.ne1024pg2_ne1024pg2.F20TR-SCREAMv1.c10-sep11.out10min.cubesphere.zarr.
(Setting ZARR_ROOT itself to an hf:// prefix does not work: the configs
join it with the original store names, which differ from the Hugging Face
repository names.)
The auxiliary files used for training and cubed-sphere inference are
latlon_ne1024pg2.nc, ne1024pg2_scrip.nc, ne1024halo256pg2_scrip.nc,
and scream_vertical_coordinate.nc; place them under AUX_DATA_ROOT.
scream_vertical_coordinate.nc ships in data/; the other three are
generated from scratch (pure numpy + netCDF4, ~15 minutes) by
bash data_prep/scrip_generation/generate_all.sh --output-dir <AUX_DATA_ROOT>
(see data_prep/scrip_generation/; the
derived latlon_ne1024pg2.nc is bit-identical to the file the shipped
configs were trained with).
Availability: trained model checkpoints are not yet publicly distributed, so the rollout-from-checkpoint workflows below cannot currently run end to end outside NVIDIA. Training from scratch is fully reproducible: dataset (Hugging Face, above), auxiliary files (generated, above), and environment (Docker, above) are all public.
Training settings are named Python configs in train_configs.py; see
screamcast/config.py for the full config reference. Add an experiment
(optionally branching off an existing one with dataclasses.replace()), then
launch it directly in an interactive GPU session.
Single GPU:
python train.py <config_name>
Multiple GPUs on one node:
torchrun --nproc_per_node=<num_gpus> train.py <config_name>
Checkpoints (best.pth / latest.pth) and logs are written to the config's
rundir (default output/). train.py uses PyTorch Lightning Fabric and
auto-detects the launch environment, so the same entry point also runs under
SLURM srun (via SLURM_* env vars) for multi-node jobs.
Global rollouts run scripts/ace/run_screamcast_nudged.py under SLURM.
slurm/submit_inference.sh is a public example
launcher — edit the checkpoint/output paths at the top and submit it with
sbatch. Run python3 scripts/ace/run_screamcast_nudged.py --help for the full
CLI (checkpoint, number of steps, tile/halo size, omega filtering, output
levels, initial time, ...). --output-levels selects which vertical levels are
written to the output zarr.
For a quick rollout on a single tile, open
notebooks/run_screamcast.py as a
marimo notebook (marimo edit notebooks/run_screamcast.py).
It loads a checkpoint via ScreamcastModel and rolls out on one tile so you can
inspect predictions interactively.
The optional ACE→SCREAM forecast-residual workflow is documented in
docs/ace2scream_finetuning.md; its scripts
live under scripts/ace/.
make lint # SPDX license headers, black, and ruff
pytest
This project will download and install additional third-party open-source software. Review the license terms of those projects before use.