Skip to content

Model Hook Injection: Perturbation

Adding model noise by using custom hooks.

This example will demonstrate how to run an ensemble inference workflow to generate a perturbed ensemble forecast. This perturbation is done by injecting code into the model front and rear hooks. These hooks are applied to the tensor data before/after the model forward call.

This example also illustrates how you can subselect data for IO. In this example we will only output two variables: total column water vapor (tcwv) and 500 hPa geopotential (z500). To run this, make sure that the model selected predicts these variables are change appropriately.

In this example you will learn:

  • How to instantiate a built in prognostic model
  • Creating a data source and IO object
  • Changing the model forward/rear hooks
  • Choose a subselection of coordinates to save to an IO object.
  • Post-processing results

Creating an Ensemble Workflow

To start let's begin with creating an ensemble workflow to use. We encourage users to explore and experiment with their own custom workflows that borrow ideas from built in workflows inside earth2studio.run or the examples.

Creating our own generalizable ensemble workflow is easy when we rely on the component interfaces defined in Earth2Studio (use dependency injection). Here we create a run method that accepts the following:

  • time: Input list of datetimes / strings to run inference for
  • nsteps: Number of forecast steps to predict
  • nensemble: Number of ensembles to run for
  • prognostic: Our initialized prognostic model
  • data: Initialized data source to fetch initial conditions from
  • io: io store that data is written to.
  • output_coords: CoordSystem of output coordinates that should be saved. Should be a proper subset of model output coordinates.

Set Up

With the ensemble workflow defined, we now need to create the individual components.

We need the following:

We will first run the ensemble workflow using an unmodified function, that is a model that has the default (identity) forward and rear hooks. Then we will define new hooks for the model and rerun the inference request.

import os

os.makedirs("outputs", exist_ok=True)
from dotenv import load_dotenv

load_dotenv()  # TODO: make common example prep function

import numpy as np

from earth2studio.data import GFS
from earth2studio.io import ZarrBackend
from earth2studio.models.px import DLWP
from earth2studio.perturbation import Gaussian
from earth2studio.run import ensemble

# Load the default model package which downloads the check point from NGC
package = DLWP.load_default_package()
model = DLWP.load_model(package)

# Create the data source
data = GFS()

# Create the IO handler, store in memory
chunks = {"ensemble": 1, "time": 1, "lead_time": 1}
io_unperturbed = ZarrBackend(
    file_name="outputs/05_ensemble.zarr",
    chunks=chunks,
    backend_kwargs={"overwrite": True},
)
Console output9 lines
/__w/earth2studio/earth2studio/.venv/lib/python3.13/site-packages/torch/cuda/__init__.py:64: FutureWarning: The pynvml package is deprecated. Please install nvidia-ml-py instead. If you did not install pynvml directly, please report this to the maintainers of the package that installed pynvml for you.
  import pynvml  # type: ignore[import]
WARNING[XFORMERS]: xFormers can't load C++/CUDA extensions. xFormers was built for:
    PyTorch 2.10.0+cu128 with CUDA 1208 (you have 2.13.0+cu130)
    Python  3.10.19 (you have 3.13.13)
  Please reinstall xformers (see https://github.com/facebookresearch/xformers#installing-xformers)
  Memory-efficient attention, SwiGLU, sparse and more won't be available.
  Set XFORMERS_MORE_DETAILS=1 for more details
CuPy distance computation test failed with error: cuVS >= 24.12 or pylibraft < 24.12 should be installed to use this feature

Execute the Workflow

First, we will run the ensemble workflow but with a earth2studio.perturbation.Gaussian perturbation as the control.

The workflow will return the provided IO object back to the user, which can be used to then post process. Some have additional APIs that can be handy for post-processing or saving to file. Check the API docs for more information.

nsteps = 4 * 12
nensemble = 16
batch_size = 4
forecast_date = "2024-01-01"
output_coords = {
    "lat": np.arange(25.0, 60.0, 0.25),
    "lon": np.arange(230.0, 300.0, 0.25),
    "variable": np.array(["tcwv", "z500"]),
}

# First run with no model perturbation
io_unperturbed = ensemble(
    [forecast_date],
    nsteps,
    nensemble,
    model,
    data,
    io_unperturbed,
    Gaussian(noise_amplitude=0.01),
    output_coords=output_coords,
    batch_size=batch_size,
)
Console output209 lines
2026-08-15 04:39:41.746 | INFO     | earth2studio.run:ensemble:394 - Running ensemble inference!
2026-08-15 04:39:41.747 | INFO     | earth2studio.run:ensemble:401 - Inference device: cuda

Fetching GFS data:   0%|          | 0/7 [00:00<?, ?it/s]
Fetching GFS data:  14%|โ–ˆโ–        | 1/7 [00:00<00:01,  4.48it/s]
Fetching GFS data: 100%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆ| 7/7 [00:00<00:00, 30.08it/s]

Fetching GFS data:   0%|          | 0/7 [00:00<?, ?it/s]
Fetching GFS data:  14%|โ–ˆโ–        | 1/7 [00:00<00:00,  6.14it/s]
Fetching GFS data: 100%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆ| 7/7 [00:00<00:00, 40.40it/s]
2026-08-15 04:39:42.255 | SUCCESS  | earth2studio.run:ensemble:422 - Fetched data from GFS
2026-08-15 04:39:42.311 | INFO     | earth2studio.run:ensemble:466 - Starting 16 Member Ensemble Inference with             4 number of batches.



Total Ensemble Batches:   0%|          | 0/4 [00:00<?, ?it/s]

Running batch 0 inference:   0%|          | 0/49 [00:00<?, ?it/s]

Running batch 0 inference:   2%|โ–         | 1/49 [00:00<00:06,  7.44it/s]

Running batch 0 inference:   4%|โ–         | 2/49 [00:00<00:11,  4.17it/s]

Running batch 0 inference:   8%|โ–Š         | 4/49 [00:00<00:05,  7.92it/s]

Running batch 0 inference:  12%|โ–ˆโ–        | 6/49 [00:00<00:04, 10.33it/s]

Running batch 0 inference:  16%|โ–ˆโ–‹        | 8/49 [00:00<00:03, 11.41it/s]

Running batch 0 inference:  20%|โ–ˆโ–ˆ        | 10/49 [00:00<00:03, 12.88it/s]

Running batch 0 inference:  24%|โ–ˆโ–ˆโ–       | 12/49 [00:01<00:02, 13.13it/s]

Running batch 0 inference:  29%|โ–ˆโ–ˆโ–Š       | 14/49 [00:01<00:02, 13.58it/s]

Running batch 0 inference:  33%|โ–ˆโ–ˆโ–ˆโ–Ž      | 16/49 [00:01<00:02, 14.90it/s]

Running batch 0 inference:  39%|โ–ˆโ–ˆโ–ˆโ–‰      | 19/49 [00:01<00:01, 17.08it/s]

Running batch 0 inference:  45%|โ–ˆโ–ˆโ–ˆโ–ˆโ–     | 22/49 [00:01<00:01, 17.69it/s]

Running batch 0 inference:  51%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆ     | 25/49 [00:01<00:01, 18.92it/s]

Running batch 0 inference:  57%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–‹    | 28/49 [00:01<00:01, 18.29it/s]

Running batch 0 inference:  63%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Ž   | 31/49 [00:02<00:00, 19.42it/s]

Running batch 0 inference:  69%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–‰   | 34/49 [00:02<00:00, 19.21it/s]

Running batch 0 inference:  76%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Œ  | 37/49 [00:02<00:00, 19.94it/s]

Running batch 0 inference:  82%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ– | 40/49 [00:02<00:00, 18.79it/s]

Running batch 0 inference:  86%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Œ | 42/49 [00:02<00:00, 18.99it/s]

Running batch 0 inference:  92%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–| 45/49 [00:02<00:00, 19.74it/s]

Running batch 0 inference:  98%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Š| 48/49 [00:02<00:00, 19.42it/s]




Total Ensemble Batches:  25%|โ–ˆโ–ˆโ–Œ       | 1/4 [00:03<00:09,  3.06s/it]

Running batch 4 inference:   0%|          | 0/49 [00:00<?, ?it/s]

Running batch 4 inference:   4%|โ–         | 2/49 [00:00<00:02, 18.20it/s]

Running batch 4 inference:   8%|โ–Š         | 4/49 [00:00<00:02, 19.19it/s]

Running batch 4 inference:  12%|โ–ˆโ–        | 6/49 [00:00<00:02, 19.52it/s]

Running batch 4 inference:  16%|โ–ˆโ–‹        | 8/49 [00:00<00:02, 19.45it/s]

Running batch 4 inference:  22%|โ–ˆโ–ˆโ–       | 11/49 [00:00<00:01, 20.64it/s]

Running batch 4 inference:  29%|โ–ˆโ–ˆโ–Š       | 14/49 [00:00<00:01, 20.19it/s]

Running batch 4 inference:  35%|โ–ˆโ–ˆโ–ˆโ–      | 17/49 [00:00<00:01, 20.93it/s]

Running batch 4 inference:  41%|โ–ˆโ–ˆโ–ˆโ–ˆ      | 20/49 [00:00<00:01, 20.43it/s]

Running batch 4 inference:  47%|โ–ˆโ–ˆโ–ˆโ–ˆโ–‹     | 23/49 [00:01<00:01, 20.41it/s]

Running batch 4 inference:  53%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Ž    | 26/49 [00:01<00:01, 19.50it/s]

Running batch 4 inference:  57%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–‹    | 28/49 [00:01<00:01, 19.09it/s]

Running batch 4 inference:  63%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Ž   | 31/49 [00:01<00:00, 20.14it/s]

Running batch 4 inference:  69%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–‰   | 34/49 [00:01<00:00, 19.99it/s]

Running batch 4 inference:  76%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Œ  | 37/49 [00:01<00:00, 20.41it/s]

Running batch 4 inference:  82%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ– | 40/49 [00:02<00:00, 16.89it/s]

Running batch 4 inference:  88%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Š | 43/49 [00:02<00:00, 17.99it/s]

Running batch 4 inference:  92%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–| 45/49 [00:02<00:00, 18.07it/s]

Running batch 4 inference:  96%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Œ| 47/49 [00:02<00:00, 18.24it/s]




Total Ensemble Batches:  50%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆ     | 2/4 [00:05<00:05,  2.76s/it]

Running batch 8 inference:   0%|          | 0/49 [00:00<?, ?it/s]

Running batch 8 inference:   6%|โ–Œ         | 3/49 [00:00<00:02, 22.11it/s]

Running batch 8 inference:  12%|โ–ˆโ–        | 6/49 [00:00<00:02, 20.30it/s]

Running batch 8 inference:  18%|โ–ˆโ–Š        | 9/49 [00:00<00:01, 20.38it/s]

Running batch 8 inference:  24%|โ–ˆโ–ˆโ–       | 12/49 [00:00<00:01, 19.49it/s]

Running batch 8 inference:  29%|โ–ˆโ–ˆโ–Š       | 14/49 [00:00<00:01, 17.99it/s]

Running batch 8 inference:  33%|โ–ˆโ–ˆโ–ˆโ–Ž      | 16/49 [00:00<00:01, 18.05it/s]

Running batch 8 inference:  37%|โ–ˆโ–ˆโ–ˆโ–‹      | 18/49 [00:00<00:01, 17.36it/s]

Running batch 8 inference:  41%|โ–ˆโ–ˆโ–ˆโ–ˆ      | 20/49 [00:01<00:01, 16.50it/s]

Running batch 8 inference:  45%|โ–ˆโ–ˆโ–ˆโ–ˆโ–     | 22/49 [00:01<00:01, 17.13it/s]

Running batch 8 inference:  49%|โ–ˆโ–ˆโ–ˆโ–ˆโ–‰     | 24/49 [00:01<00:01, 17.13it/s]

Running batch 8 inference:  53%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Ž    | 26/49 [00:01<00:01, 17.14it/s]

Running batch 8 inference:  57%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–‹    | 28/49 [00:01<00:01, 15.95it/s]

Running batch 8 inference:  61%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆ    | 30/49 [00:01<00:01, 14.48it/s]

Running batch 8 inference:  65%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Œ   | 32/49 [00:01<00:01, 15.76it/s]

Running batch 8 inference:  69%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–‰   | 34/49 [00:01<00:00, 15.90it/s]

Running batch 8 inference:  73%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Ž  | 36/49 [00:02<00:00, 16.11it/s]

Running batch 8 inference:  78%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Š  | 38/49 [00:02<00:00, 15.36it/s]

Running batch 8 inference:  82%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ– | 40/49 [00:02<00:00, 14.22it/s]

Running batch 8 inference:  86%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Œ | 42/49 [00:02<00:00, 13.59it/s]

Running batch 8 inference:  90%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–‰ | 44/49 [00:02<00:00, 13.52it/s]

Running batch 8 inference:  94%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–| 46/49 [00:02<00:00, 14.25it/s]

Running batch 8 inference:  98%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Š| 48/49 [00:02<00:00, 14.02it/s]




Total Ensemble Batches:  75%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Œ  | 3/4 [00:08<00:02,  2.89s/it]

Running batch 12 inference:   0%|          | 0/49 [00:00<?, ?it/s]

Running batch 12 inference:   4%|โ–         | 2/49 [00:00<00:02, 17.70it/s]

Running batch 12 inference:   8%|โ–Š         | 4/49 [00:00<00:02, 17.89it/s]

Running batch 12 inference:  12%|โ–ˆโ–        | 6/49 [00:00<00:02, 18.21it/s]

Running batch 12 inference:  16%|โ–ˆโ–‹        | 8/49 [00:00<00:02, 17.91it/s]

Running batch 12 inference:  20%|โ–ˆโ–ˆ        | 10/49 [00:00<00:02, 17.94it/s]

Running batch 12 inference:  24%|โ–ˆโ–ˆโ–       | 12/49 [00:00<00:02, 18.20it/s]

Running batch 12 inference:  29%|โ–ˆโ–ˆโ–Š       | 14/49 [00:00<00:01, 17.81it/s]

Running batch 12 inference:  33%|โ–ˆโ–ˆโ–ˆโ–Ž      | 16/49 [00:00<00:01, 17.95it/s]

Running batch 12 inference:  37%|โ–ˆโ–ˆโ–ˆโ–‹      | 18/49 [00:00<00:01, 18.34it/s]

Running batch 12 inference:  41%|โ–ˆโ–ˆโ–ˆโ–ˆ      | 20/49 [00:01<00:01, 18.70it/s]

Running batch 12 inference:  45%|โ–ˆโ–ˆโ–ˆโ–ˆโ–     | 22/49 [00:01<00:01, 18.76it/s]

Running batch 12 inference:  51%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆ     | 25/49 [00:01<00:01, 20.08it/s]

Running batch 12 inference:  57%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–‹    | 28/49 [00:01<00:01, 19.54it/s]

Running batch 12 inference:  63%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Ž   | 31/49 [00:01<00:00, 20.46it/s]

Running batch 12 inference:  69%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–‰   | 34/49 [00:01<00:00, 19.71it/s]

Running batch 12 inference:  73%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Ž  | 36/49 [00:01<00:00, 19.51it/s]

Running batch 12 inference:  78%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Š  | 38/49 [00:02<00:00, 19.07it/s]

Running batch 12 inference:  82%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ– | 40/49 [00:02<00:00, 19.03it/s]

Running batch 12 inference:  86%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Œ | 42/49 [00:02<00:00, 19.17it/s]

Running batch 12 inference:  92%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–| 45/49 [00:02<00:00, 20.27it/s]

Running batch 12 inference:  98%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Š| 48/49 [00:02<00:00, 19.66it/s]




Total Ensemble Batches: 100%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆ| 4/4 [00:11<00:00,  2.75s/it]
Total Ensemble Batches: 100%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆ| 4/4 [00:11<00:00,  2.80s/it]
2026-08-15 04:39:53.509 | SUCCESS  | earth2studio.run:ensemble:536 - 
Inference complete

Now let's introduce slight model perturbation using the prognostic model hooks defined in [earth2studio.models.px.utils.PrognosticMixin][]. Note that center.unsqueeze(-1) is DLWP specific since it operates on a cubed sphere with grid dimensions (nface, lat, lon) instead of just (lat,lon). To switch out the model, consider removing the unsqueeze .

model.front_hook = lambda x, coords: (
    x
    - 0.1
    * x.var(dim=0)
    * (x - model.center.unsqueeze(-1))
    / (model.scale.unsqueeze(-1)) ** 2
    + 0.1 * (x - x.mean(dim=0)),
    coords,
)
# Also could use model.rear_hook = ...

io_perturbed = ZarrBackend(
    file_name="outputs/05_ensemble_model_perturbation.zarr",
    chunks=chunks,
    backend_kwargs={"overwrite": True},
)
io_perturbed = ensemble(
    [forecast_date],
    nsteps,
    nensemble,
    model,
    data,
    io_perturbed,
    Gaussian(noise_amplitude=0.01),
    output_coords=output_coords,
    batch_size=batch_size,
)
Console output209 lines
2026-08-15 04:39:54.091 | INFO     | earth2studio.run:ensemble:394 - Running ensemble inference!
2026-08-15 04:39:54.091 | INFO     | earth2studio.run:ensemble:401 - Inference device: cuda

Fetching GFS data:   0%|          | 0/7 [00:00<?, ?it/s]
Fetching GFS data:  14%|โ–ˆโ–        | 1/7 [00:00<00:00,  8.05it/s]
Fetching GFS data: 100%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆ| 7/7 [00:00<00:00, 56.09it/s]

Fetching GFS data:   0%|          | 0/7 [00:00<?, ?it/s]
Fetching GFS data:  29%|โ–ˆโ–ˆโ–Š       | 2/7 [00:00<00:00, 13.60it/s]
Fetching GFS data: 100%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆ| 7/7 [00:00<00:00, 44.83it/s]
2026-08-15 04:39:54.443 | SUCCESS  | earth2studio.run:ensemble:422 - Fetched data from GFS
2026-08-15 04:39:54.505 | INFO     | earth2studio.run:ensemble:466 - Starting 16 Member Ensemble Inference with             4 number of batches.



Total Ensemble Batches:   0%|          | 0/4 [00:00<?, ?it/s]

Running batch 0 inference:   0%|          | 0/49 [00:00<?, ?it/s]

Running batch 0 inference:   4%|โ–         | 2/49 [00:00<00:03, 11.90it/s]

Running batch 0 inference:  10%|โ–ˆ         | 5/49 [00:00<00:02, 17.24it/s]

Running batch 0 inference:  16%|โ–ˆโ–‹        | 8/49 [00:00<00:02, 18.17it/s]

Running batch 0 inference:  22%|โ–ˆโ–ˆโ–       | 11/49 [00:00<00:01, 19.57it/s]

Running batch 0 inference:  29%|โ–ˆโ–ˆโ–Š       | 14/49 [00:00<00:01, 18.62it/s]

Running batch 0 inference:  35%|โ–ˆโ–ˆโ–ˆโ–      | 17/49 [00:00<00:01, 19.78it/s]

Running batch 0 inference:  41%|โ–ˆโ–ˆโ–ˆโ–ˆ      | 20/49 [00:01<00:01, 19.45it/s]

Running batch 0 inference:  47%|โ–ˆโ–ˆโ–ˆโ–ˆโ–‹     | 23/49 [00:01<00:01, 20.19it/s]

Running batch 0 inference:  53%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Ž    | 26/49 [00:01<00:01, 19.84it/s]

Running batch 0 inference:  59%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–‰    | 29/49 [00:01<00:00, 20.57it/s]

Running batch 0 inference:  65%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Œ   | 32/49 [00:01<00:00, 20.13it/s]

Running batch 0 inference:  71%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–  | 35/49 [00:01<00:00, 20.59it/s]

Running batch 0 inference:  78%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Š  | 38/49 [00:01<00:00, 19.98it/s]

Running batch 0 inference:  84%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Ž | 41/49 [00:02<00:00, 20.79it/s]

Running batch 0 inference:  90%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–‰ | 44/49 [00:02<00:00, 20.03it/s]

Running batch 0 inference:  96%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Œ| 47/49 [00:02<00:00, 19.88it/s]




Total Ensemble Batches:  25%|โ–ˆโ–ˆโ–Œ       | 1/4 [00:02<00:07,  2.50s/it]

Running batch 4 inference:   0%|          | 0/49 [00:00<?, ?it/s]

Running batch 4 inference:   6%|โ–Œ         | 3/49 [00:00<00:02, 21.59it/s]

Running batch 4 inference:  12%|โ–ˆโ–        | 6/49 [00:00<00:02, 19.89it/s]

Running batch 4 inference:  16%|โ–ˆโ–‹        | 8/49 [00:00<00:02, 19.28it/s]

Running batch 4 inference:  20%|โ–ˆโ–ˆ        | 10/49 [00:00<00:02, 19.33it/s]

Running batch 4 inference:  24%|โ–ˆโ–ˆโ–       | 12/49 [00:00<00:01, 19.31it/s]

Running batch 4 inference:  29%|โ–ˆโ–ˆโ–Š       | 14/49 [00:00<00:01, 18.55it/s]

Running batch 4 inference:  33%|โ–ˆโ–ˆโ–ˆโ–Ž      | 16/49 [00:00<00:01, 18.32it/s]

Running batch 4 inference:  37%|โ–ˆโ–ˆโ–ˆโ–‹      | 18/49 [00:00<00:01, 18.23it/s]

Running batch 4 inference:  41%|โ–ˆโ–ˆโ–ˆโ–ˆ      | 20/49 [00:01<00:01, 18.62it/s]

Running batch 4 inference:  45%|โ–ˆโ–ˆโ–ˆโ–ˆโ–     | 22/49 [00:01<00:01, 18.11it/s]

Running batch 4 inference:  51%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆ     | 25/49 [00:01<00:01, 19.26it/s]

Running batch 4 inference:  55%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Œ    | 27/49 [00:01<00:01, 18.46it/s]

Running batch 4 inference:  59%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–‰    | 29/49 [00:01<00:01, 18.83it/s]

Running batch 4 inference:  63%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Ž   | 31/49 [00:01<00:00, 19.08it/s]

Running batch 4 inference:  67%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–‹   | 33/49 [00:01<00:00, 19.24it/s]

Running batch 4 inference:  73%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Ž  | 36/49 [00:01<00:00, 19.22it/s]

Running batch 4 inference:  80%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–‰  | 39/49 [00:02<00:00, 20.18it/s]

Running batch 4 inference:  86%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Œ | 42/49 [00:02<00:00, 18.16it/s]

Running batch 4 inference:  90%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–‰ | 44/49 [00:02<00:00, 18.55it/s]

Running batch 4 inference:  94%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–| 46/49 [00:02<00:00, 18.13it/s]

Running batch 4 inference:  98%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Š| 48/49 [00:02<00:00, 18.59it/s]




Total Ensemble Batches:  50%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆ     | 2/4 [00:05<00:05,  2.55s/it]

Running batch 8 inference:   0%|          | 0/49 [00:00<?, ?it/s]

Running batch 8 inference:   4%|โ–         | 2/49 [00:00<00:02, 18.92it/s]

Running batch 8 inference:  10%|โ–ˆ         | 5/49 [00:00<00:02, 20.87it/s]

Running batch 8 inference:  16%|โ–ˆโ–‹        | 8/49 [00:00<00:02, 19.98it/s]

Running batch 8 inference:  22%|โ–ˆโ–ˆโ–       | 11/49 [00:00<00:01, 20.81it/s]

Running batch 8 inference:  29%|โ–ˆโ–ˆโ–Š       | 14/49 [00:00<00:01, 19.54it/s]

Running batch 8 inference:  33%|โ–ˆโ–ˆโ–ˆโ–Ž      | 16/49 [00:00<00:01, 19.23it/s]

Running batch 8 inference:  37%|โ–ˆโ–ˆโ–ˆโ–‹      | 18/49 [00:00<00:01, 19.04it/s]

Running batch 8 inference:  41%|โ–ˆโ–ˆโ–ˆโ–ˆ      | 20/49 [00:01<00:01, 18.67it/s]

Running batch 8 inference:  45%|โ–ˆโ–ˆโ–ˆโ–ˆโ–     | 22/49 [00:01<00:01, 18.83it/s]

Running batch 8 inference:  49%|โ–ˆโ–ˆโ–ˆโ–ˆโ–‰     | 24/49 [00:01<00:01, 19.00it/s]

Running batch 8 inference:  53%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Ž    | 26/49 [00:01<00:01, 18.96it/s]

Running batch 8 inference:  57%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–‹    | 28/49 [00:01<00:01, 18.61it/s]

Running batch 8 inference:  61%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆ    | 30/49 [00:01<00:01, 17.38it/s]

Running batch 8 inference:  65%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Œ   | 32/49 [00:01<00:00, 17.45it/s]

Running batch 8 inference:  69%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–‰   | 34/49 [00:01<00:00, 18.02it/s]

Running batch 8 inference:  73%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Ž  | 36/49 [00:01<00:00, 18.50it/s]

Running batch 8 inference:  78%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Š  | 38/49 [00:02<00:00, 18.83it/s]

Running batch 8 inference:  82%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ– | 40/49 [00:02<00:00, 19.05it/s]

Running batch 8 inference:  88%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Š | 43/49 [00:02<00:00, 19.79it/s]

Running batch 8 inference:  92%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–| 45/49 [00:02<00:00, 19.84it/s]

Running batch 8 inference:  96%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Œ| 47/49 [00:02<00:00, 19.54it/s]

Running batch 8 inference: 100%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆ| 49/49 [00:02<00:00, 18.77it/s]




Total Ensemble Batches:  75%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Œ  | 3/4 [00:07<00:02,  2.56s/it]

Running batch 12 inference:   0%|          | 0/49 [00:00<?, ?it/s]

Running batch 12 inference:   4%|โ–         | 2/49 [00:00<00:02, 17.51it/s]

Running batch 12 inference:   8%|โ–Š         | 4/49 [00:00<00:02, 16.30it/s]

Running batch 12 inference:  12%|โ–ˆโ–        | 6/49 [00:00<00:02, 17.34it/s]

Running batch 12 inference:  16%|โ–ˆโ–‹        | 8/49 [00:00<00:02, 17.40it/s]

Running batch 12 inference:  20%|โ–ˆโ–ˆ        | 10/49 [00:00<00:02, 16.43it/s]

Running batch 12 inference:  24%|โ–ˆโ–ˆโ–       | 12/49 [00:00<00:02, 17.18it/s]

Running batch 12 inference:  29%|โ–ˆโ–ˆโ–Š       | 14/49 [00:00<00:02, 17.47it/s]

Running batch 12 inference:  33%|โ–ˆโ–ˆโ–ˆโ–Ž      | 16/49 [00:00<00:01, 17.88it/s]

Running batch 12 inference:  37%|โ–ˆโ–ˆโ–ˆโ–‹      | 18/49 [00:01<00:01, 18.04it/s]

Running batch 12 inference:  41%|โ–ˆโ–ˆโ–ˆโ–ˆ      | 20/49 [00:01<00:01, 18.13it/s]

Running batch 12 inference:  45%|โ–ˆโ–ˆโ–ˆโ–ˆโ–     | 22/49 [00:01<00:01, 17.99it/s]

Running batch 12 inference:  49%|โ–ˆโ–ˆโ–ˆโ–ˆโ–‰     | 24/49 [00:01<00:01, 17.38it/s]

Running batch 12 inference:  53%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Ž    | 26/49 [00:01<00:01, 17.45it/s]

Running batch 12 inference:  57%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–‹    | 28/49 [00:01<00:01, 16.57it/s]

Running batch 12 inference:  61%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆ    | 30/49 [00:01<00:01, 15.97it/s]

Running batch 12 inference:  65%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Œ   | 32/49 [00:01<00:01, 15.84it/s]

Running batch 12 inference:  69%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–‰   | 34/49 [00:02<00:00, 16.21it/s]

Running batch 12 inference:  76%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Œ  | 37/49 [00:02<00:00, 18.31it/s]

Running batch 12 inference:  82%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ– | 40/49 [00:02<00:00, 17.66it/s]

Running batch 12 inference:  86%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Œ | 42/49 [00:02<00:00, 17.13it/s]

Running batch 12 inference:  90%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–‰ | 44/49 [00:02<00:00, 15.60it/s]

Running batch 12 inference:  96%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–Œ| 47/49 [00:02<00:00, 17.34it/s]




Total Ensemble Batches: 100%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆ| 4/4 [00:10<00:00,  2.67s/it]
Total Ensemble Batches: 100%|โ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆโ–ˆ| 4/4 [00:10<00:00,  2.63s/it]
2026-08-15 04:40:05.016 | SUCCESS  | earth2studio.run:ensemble:536 - 
Inference complete

Post Processing

The last step is to post process our results. Here we plot and compare the ensemble mean and standard deviation from using an unperturbed/perturbed model.

Notice that the Zarr IO function has additional APIs to interact with the stored data.

import cartopy.crs as ccrs
import matplotlib.pyplot as plt
from matplotlib.colors import LogNorm

levels_unperturbed = np.linspace(0, io_unperturbed["tcwv"][:].max())
levels_perturbed = np.linspace(0, io_perturbed["tcwv"][:].max())


std_levels_perturbed = np.linspace(0, io_perturbed["tcwv"][:].std(axis=0).max())

plt.close("all")
fig = plt.figure(figsize=(20, 10), tight_layout=True)
ax0 = fig.add_subplot(2, 2, 1, projection=ccrs.PlateCarree())
ax1 = fig.add_subplot(2, 2, 2, projection=ccrs.PlateCarree())
ax2 = fig.add_subplot(2, 2, 3, projection=ccrs.PlateCarree())
ax3 = fig.add_subplot(2, 2, 4, projection=ccrs.PlateCarree())


def update(frame):
    """This function updates the frame with a new lead time for animation."""
    import warnings

    warnings.filterwarnings("ignore")
    ax0.clear()
    ax1.clear()
    ax2.clear()
    ax3.clear()

    ## Update unperturbed image
    im0 = ax0.contourf(
        io_unperturbed["lon"][:],
        io_unperturbed["lat"][:],
        io_unperturbed["tcwv"][:, 0, frame].mean(axis=0),
        transform=ccrs.PlateCarree(),
        cmap="Blues",
        levels=levels_unperturbed,
    )
    ax0.coastlines()
    ax0.gridlines()

    im1 = ax1.contourf(
        io_unperturbed["lon"][:],
        io_unperturbed["lat"][:],
        io_unperturbed["tcwv"][:, 0, frame].std(axis=0),
        transform=ccrs.PlateCarree(),
        cmap="RdPu",
        levels=std_levels_perturbed,
        norm=LogNorm(vmin=1e-1, vmax=std_levels_perturbed[-1]),
    )
    ax1.coastlines()
    ax1.gridlines()

    im2 = ax2.contourf(
        io_perturbed["lon"][:],
        io_perturbed["lat"][:],
        io_perturbed["tcwv"][:, 0, frame].mean(axis=0),
        transform=ccrs.PlateCarree(),
        cmap="Blues",
        levels=levels_perturbed,
    )
    ax2.coastlines()
    ax2.gridlines()

    im3 = ax3.contourf(
        io_perturbed["lon"][:],
        io_perturbed["lat"][:],
        io_perturbed["tcwv"][:, 0, frame].std(axis=0),
        transform=ccrs.PlateCarree(),
        cmap="RdPu",
        levels=std_levels_perturbed,
        norm=LogNorm(vmin=1e-1, vmax=std_levels_perturbed[-1]),
    )
    ax3.coastlines()
    ax3.gridlines()

    for i in range(16):
        ax0.contour(
            io_unperturbed["lon"][:],
            io_unperturbed["lat"][:],
            io_unperturbed["z500"][i, 0, frame] / 100.0,
            transform=ccrs.PlateCarree(),
            levels=np.arange(485, 580, 15),
            colors="black",
            linestyle="dashed",
        )

        ax2.contour(
            io_perturbed["lon"][:],
            io_perturbed["lat"][:],
            io_perturbed["z500"][i, 0, frame] / 100.0,
            transform=ccrs.PlateCarree(),
            levels=np.arange(485, 580, 15),
            colors="black",
            linestyle="dashed",
        )
    plt.suptitle(
        f'Forecast Starting on {forecast_date} - Lead Time - {io_perturbed["lead_time"][frame]}'
    )

    ax0.set_title("Unperturbed Ensemble Mean - tcwv + z500 countors")
    ax1.set_title("Unperturbed Ensemble Std - tcwv")
    ax2.set_title("Perturbed Ensemble Mean - tcwv + z500 contours")
    ax3.set_title("Perturbed Ensemble Std - tcwv")

    if frame == 0:
        plt.colorbar(
            im0, ax=ax0, shrink=0.75, pad=0.04, label="kg m^-2", format="%2.1f"
        )
        plt.colorbar(
            im1, ax=ax1, shrink=0.75, pad=0.04, label="kg m^-2", format="%1.2e"
        )
        plt.colorbar(
            im2, ax=ax2, shrink=0.75, pad=0.04, label="kg m^-2", format="%2.1f"
        )
        plt.colorbar(
            im3, ax=ax3, shrink=0.75, pad=0.04, label="kg m^-2", format="%1.2e"
        )


# Uncomment this for animation
# import matplotlib.animation as animation
# update(0)
# ani = animation.FuncAnimation(
# fig=fig, func=update, frames=range(1, nsteps), cache_frame_data=False
# )
# ani.save(f"outputs/05_model_perturbation_{forecast_date}.gif", dpi=300)


for lt in [10, 20, 30, 40]:
    update(lt)
    plt.savefig(
        f"outputs/05_model_perturbation_{forecast_date}_leadtime_{lt}.png",
        dpi=300,
        bbox_inches="tight",
    )

Output from Model Hook Injection: Perturbation

Output from Model Hook Injection: Perturbation

Output from Model Hook Injection: Perturbation

Output from Model Hook Injection: Perturbation


Execution profile

Runtime telemetry

Total runtime1m 10s

Execution environment

CPUAMD EPYC 7313P 16-Core Processor
GPUNVIDIA H100 PCIe ยท 79.6 GiB
System RAM58.5 GiB
PlatformLinux 6.8.0-136-generic
Python3.13.13
GPU driver / CUDADriver 595.84 ยท CUDA support 13.2