Skip to content

StormCastSDA

NADANWC202440 GBNVIDIAPyTorch

Import path: earth2studio.models.da.StormCastSDA

View source on GitHub View install commands

Documentation

Bases: Module, AutoModelMixin

StormCast with score-based data assimilation (SDA) using diffusion posterior sampling for convection-allowing regional forecasts. Combines a regression and diffusion model with DPS guidance to assimilate observations during inference. Model time step size is 1 hour, taking as input:

  • High-resolution (3km) HRRR state over the central United States (99 vars)
  • High-resolution land-sea mask and orography invariants
  • Coarse resolution (25km) global state (26 vars)
  • Point observations for data assimilation

The high-resolution grid is the HRRR Lambert conformal projection. Coarse-resolution inputs are regridded to the HRRR grid internally.

Note

For more information see the following references:

Parameters:

  • regression_model (Module) –

    Deterministic model used to make an initial prediction

  • diffusion_model (Module) –

    Generative model correcting the deterministic prediciton

  • means (Tensor) –

    Mean value of each input high-resolution variable

  • stds (Tensor) –

    Standard deviation of each input high-resolution variable

  • invariants (Tensor) –

    Static invariant quantities

  • hrrr_lat_lim (tuple[int, int], default: (273, 785) ) –

    HRRR grid latitude limits, defaults to be the StormCastV1 region in central United States, by default (273, 785)

  • hrrr_lon_lim (tuple[int, int], default: (579, 1219) ) –

    HRRR grid longitude limits, defaults to be the StormCastV1 region in central United States,, by default (579, 1219)

  • variables (array, default: array(VARIABLES) ) –

    High-resolution variables, by default np.array(VARIABLES)

  • conditioning_means (Tensor | None, default: None ) –

    Means to normalize conditioning data, by default None

  • conditioning_stds (Tensor | None, default: None ) –

    Standard deviations to normalize conditioning data, by default None

  • conditioning_variables (array, default: array(CONDITIONING_VARIABLES) ) –

    Global variables for conditioning, by default np.array(CONDITIONING_VARIABLES)

  • conditioning_data_source (DataSource | ForecastSource | None, default: None ) –

    Data Source to use for global conditioning. Required for running in iterator mode, by default None

  • time_tolerance (TimeTolerance, default: timedelta64(10, 'm') ) –

    Time tolerance for filtering observations. Observations within the tolerance window around each requested time will be used for data assimilation, by default np.timedelta64(10, "m")

  • sampler_steps (int, default: 36 ) –

    Number of diffusion sampler steps, by default 36

  • sampler_args (dict[str, float | int] | None, default: None ) –

    Arguments to pass to the diffusion sampler, by default None

  • sda_std_obs (float, default: 0.1 ) –

    Observation noise standard deviation for DPS guidance, by default 0.1

  • sda_gamma (float, default: 0.001 ) –

    SDA scaling factor for DPS guidance, by default 0.001

__call__

__call__(x: DataArray, obs: DataFrame | None) -> DataArray

Runs assimilation model 1 step.

Parameters:

  • x (DataArray) –

    Input state on the HRRR curvilinear grid

  • obs (DataFrame | None) –

    Sparse observations DataFrame, or None for no assimilation

Returns:

  • DataArray –

    Output state one time-step into the future

Raises:

  • RuntimeError –

    If conditioning data source is not initialized

create_generator

create_generator(
    x: DataArray,
) -> Generator[DataArray, DataFrame | None, None]

Creates a generator for iterative forecast with data assimilation.

The generator yields forecast states and receives observation DataFrames via send(). At each step, conditioning data is fetched, observations are mapped to the HRRR grid, and the diffusion model produces the next forecast step.

Parameters:

  • x (DataArray) –

    Initial state on the HRRR curvilinear grid

Yields:

  • DataArray –

    Forecast state at each time step

Receives:

  • DataFrame | None –

    Observations sent via generator.send(). Pass None for steps without assimilation.

Example
>>> gen = model.create_generator(x0)
>>> state = next(gen)           # yields initial state x0
>>> state = gen.send(obs_df)    # step 1 with observations
>>> state = gen.send(None)      # step 2 without observations

load_default_package classmethod

load_default_package() -> Package

Load assimilation package

load_model classmethod

load_model(
    package: Package,
    conditioning_data_source: (
        DataSource | ForecastSource
    ) = GFS_FX(verbose=False),
    time_tolerance: TimeTolerance = timedelta64(10, "m"),
    sampler_steps: int = 36,
    sda_std_obs: float = 0.1,
    sda_gamma: float = 0.001,
) -> AssimilationModel

Load assimilation from package

Parameters:

  • package (Package) –

    Package to load model from

  • conditioning_data_source (DataSource | ForecastSource, default: GFS_FX(verbose=False) ) –

    Data source to use for global conditioning, by default GFS_FX

  • time_tolerance (TimeTolerance, default: timedelta64(10, 'm') ) –

    Time tolerance for filtering observations. Observations within the tolerance window around each requested time will be used for data assimilation, by default np.timedelta64(10, "m")

  • sampler_steps (int, default: 36 ) –

    Number of diffusion sampler steps, by default 36

  • sda_std_obs (float, default: 0.1 ) –

    Observation noise standard deviation for DPS guidance, by default 0.1

  • sda_gamma (float, default: 0.001 ) –

    SDA scaling factor for DPS guidance, by default 0.001

Returns:

  • AssimilationModel –

    Assimilation model