DLESyM¶
GlobalS2S202540 GBNVIDIAPyTorch
Import path: earth2studio.models.px.DLESyM
View source on GitHub View install commands
Documentation¶
Bases: Module, AutoModelMixin, PrognosticMixin
DLESyM-V1-ERA5 prognostic model. This is an ensemble forecast model for global earth system modeling. This model includes an atmosphere and ocean component, using atmospheric variables as well as the sea-surface temperature on a HEALPix nside=64 (approximately 1 degree) resolution grid. The model architecture is a U-Net with padding operations modified to support using the HEALPix grid. Because the atmosphere and ocean models are predicted at different times, not all entries in the output tensor are valid. As a result, we provide convenience methods for retrieving the valid atmospheric and oceanic outputs.
The default package provided for this model contains the checkpoints used
in the ECMWF AI Weather Quest S2S competition. These checkpoints are trained
with CRPS loss and use sampled random noise to produce ensemble variability
each forward pass, seeded with the set_rng method.
Parameters:
-
atmos_model(Module) –Atmosphere model
-
ocean_model(Module) –Ocean model
-
hpx_lat(ndarray) –HEALPix latitude coordinates, shape (12, nside, nside)
-
hpx_lon(ndarray) –HEALPix longitude coordinates, shape (12, nside, nside)
-
nside(int) –HEALPix nside
-
center(ndarray) –Means of the full output variable set (prognostics + diagnostics, in the same order as
output_coords'svariableaxis), shape (1, 1, 1, num_output_variables, 1, 1, 1) -
scale(ndarray) –Standard deviations of the full output variable set, same shape and ordering as
center -
atmos_constants(ndarray) –Constants for the atmosphere model, shape (12, num_atmos_constants, nside, nside)
-
ocean_constants(ndarray) –Constants for the ocean model, shape (12, num_ocean_constants, nside, nside)
-
atmos_input_times(ndarray) –Atmospheric input times, shape (num_atmos_input_times,)
-
ocean_input_times(ndarray) –Ocean input times, shape (num_ocean_input_times,)
-
atmos_output_times(ndarray) –Atmospheric output times, shape (num_atmos_output_times,)
-
ocean_output_times(ndarray) –Ocean output times, shape (num_ocean_output_times,)
-
atmos_variables(list[str]) –Atmospheric variables
-
ocean_variables(list[str]) –Ocean variables
-
atmos_coupling_variables(list[str]) –Atmospheric coupling variables
-
ocean_coupling_variables(list[str]) –Ocean coupling variables
-
atmos_diagnostic_variables(list[str], default:None) –Atmospheric diagnostic output variables. These are produced by the atmos model but, unlike
atmos_variables, are not fed back in as input to the next autoregressive step, by default [] -
ocean_diagnostic_variables(list[str], default:None) –Ocean diagnostic output variables, analogous to
atmos_diagnostic_variables, by default [] -
use_cln(bool, default:False) –Whether the atmos/ocean models use conditional layer norm, which requires sampling and passing noise to the model to produce ensemble variability from a single set of weights, by default False
-
condition_shape(int, default:None) –Dimension of the conditional layer norm noise vector. Required if
use_clnis True, by default None
Note
For more information about this model see:
For more information about the HEALPix grid see:
Example
pkg = DLESyM.load_default_package()
model = DLESyM.load_model(pkg)
# Create iterator
iterator = model.create_iterator(x, coords)
for step, (x, coords) in enumerate(iterator):
if step > 0:
# Valid atmos and ocean predictions with their respective coordinates extracted below
atmos_outputs, atmos_coords = model.retrieve_valid_atmos_outputs(x, coords)
ocean_outputs, ocean_coords = model.retrieve_valid_ocean_outputs(x, coords)
```pycon
...
```
__call__ ¶
create_iterator ¶
Creates a iterator which can be used to perform time-integration of the prognostic model. Will return the initial condition first (0th step).
Parameters:
-
x(Tensor) –Input tensor
-
coords(CoordSystem) –Input coordinate system
Yields:
load_default_package
classmethod
¶
load_default_package() -> Package
Default DLESyM model package on NGC
The package's top-level config.yaml lists the available
checkpoint versions; see load_model for how to select
between them.
load_model
classmethod
¶
load_model(
package: Package,
atmos_model_idx: int = 0,
ocean_model_idx: int = 0,
version: Literal["v1.0", "v1.1"] = "v1.1",
) -> PrognosticModel
Load prognostic from package
Parameters:
-
package(Package) –Package to load model from
-
version(('v1.0', 'v1.1'), default:"v1.0") –Checkpoint version to load; see each version's entry in the package's
config.yamlfor details.v1.1is the checkpoint submitted to the ECMWF AI Weather Quest competition (aiweatherquest.ecmwf.int/).v1.0is the previous checkpoint, kept for reproducibility; loading it logs a deprecation warning, by default "v1.1" -
atmos_model_idx(int, default:0) –Index of atmos model weights to load. Only meaningful for checkpoint versions that ship multiple atmos checkpoints (used to build ensembles without conditional layer norm), by default 0
-
ocean_model_idx(int, default:0) –Index of ocean model weights to load. Only meaningful for checkpoint versions that ship multiple ocean checkpoints, by default 0
Returns:
-
PrognosticModel–Prognostic model