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__ ¶
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 ¶
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(). PassNonefor steps without assimilation.
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