A full GenCast predictor pipeline involves wrapping the base gencast.GenCast model with normalization and NaN cleaning layers.
- Initialize
gencast.GenCast with the loaded configs. - Wrap with
normalization.InputsAndResiduals using provided statistics (diffs_stddev_by_level, mean_by_level, stddev_by_level). - Wrap with
nan_cleaning.NaNCleaner to handle missing values (e.g., for sea_surface_temperature) using min_by_level.
def construct_wrapped_gencast():
"""Constructs and wraps the GenCast Predictor."""
predictor = gencast.GenCast(
sampler_config=sampler_config,
task_config=task_config,
denoiser_architecture_config=denoiser_architecture_config,
noise_config=noise_config,
noise_encoder_config=noise_encoder_config,
)
predictor = normalization.InputsAndResiduals(
predictor,
diffs_stddev_by_level=diffs_stddev_by_level,
mean_by_level=mean_by_level,
stddev_by_level=stddev_by_level,
)
predictor = nan_cleaning.NaNCleaner(
predictor=predictor,
reintroduce_nans=True,
fill_value=min_by_level,
var_to_clean='sea_surface_temperature',
)
return predictor