| """Copyright (c) Microsoft Corporation. Licensed under the MIT license.""" |
|
|
| import contextlib |
| import dataclasses |
| import warnings |
| from datetime import timedelta |
| from typing import Optional |
|
|
| import numpy as np |
| import torch |
| import torch.nn as nn |
| from huggingface_hub import hf_hub_download |
| from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import ( |
| apply_activation_checkpointing, |
| ) |
|
|
| from .aurora_batch import Batch |
| from .aurora_insolation import insolation |
| from .aurora_compat import ( |
| _adapt_checkpoint_air_pollution, |
| _adapt_checkpoint_pretrained, |
| _adapt_checkpoint_v1p5, |
| _adapt_checkpoint_wave, |
| ) |
| from .aurora_decoder import Perceiver3DDecoder |
| from .aurora_encoder import Perceiver3DEncoder |
| from .aurora_lora import LoRAMode |
| from .aurora_normalisation import log_transform, log_untransform |
| from .aurora_perceiver import PerceiverAttention |
| from .aurora_swin3d import Swin3DTransformerBackbone, WindowAttention |
|
|
| __all__ = [ |
| "Aurora", |
| "AuroraPretrained", |
| "AuroraSmallPretrained", |
| "AuroraSmall", |
| "Aurora12hPretrained", |
| "AuroraHighRes", |
| "AuroraAirPollution", |
| "AuroraWave", |
| "AuroraV1p5", |
| "AuroraV1p5Ensemble", |
| ] |
|
|
|
|
| class Aurora(torch.nn.Module): |
| """The Aurora model. |
| |
| Defaults to the 1.3 B parameter configuration. |
| |
| Also supports ensemble forecasts. |
| """ |
|
|
| default_checkpoint_repo = "microsoft/aurora" |
| """str: Name of the HuggingFace repository to load the default checkpoint from.""" |
|
|
| default_checkpoint_name = "aurora-0.25-finetuned.ckpt" |
| """str: Name of the default checkpoint.""" |
|
|
| default_checkpoint_revision = "0be7e57c685dac86b78c4a19a3ab149d13c6a3dd" |
| """str: Commit hash of the default checkpoint.""" |
|
|
| def __init__( |
| self, |
| *, |
| surf_vars: tuple[str, ...] = ("2t", "10u", "10v", "msl"), |
| static_vars: tuple[str, ...] = ("lsm", "z", "slt"), |
| atmos_vars: tuple[str, ...] = ("z", "u", "v", "t", "q"), |
| window_size: tuple[int, int, int] = (2, 6, 12), |
| encoder_depths: tuple[int, ...] = (6, 10, 8), |
| encoder_num_heads: tuple[int, ...] = (8, 16, 32), |
| decoder_depths: tuple[int, ...] = (8, 10, 6), |
| decoder_num_heads: tuple[int, ...] = (32, 16, 8), |
| latent_levels: int = 4, |
| patch_size: int = 4, |
| embed_dim: int = 512, |
| num_heads: int = 16, |
| mlp_ratio: float = 4.0, |
| drop_path: float = 0.0, |
| drop_rate: float = 0.0, |
| enc_depth: int = 1, |
| dec_depth: int = 1, |
| dec_mlp_ratio: float = 2.0, |
| perceiver_ln_eps: float = 1e-5, |
| max_history_size: int = 2, |
| timestep: timedelta = timedelta(hours=6), |
| stabilise_level_agg: bool = False, |
| use_lora: bool = True, |
| lora_steps: int = 40, |
| lora_mode: LoRAMode = "single", |
| surf_stats: Optional[dict[str, tuple[float, float]]] = None, |
| autocast: bool = False, |
| autocast_dtype: torch.dtype = torch.bfloat16, |
| bf16_mode: bool = False, |
| use_fp16_safe_attention: bool = False, |
| level_condition: Optional[tuple[int | float, ...]] = None, |
| dynamic_vars: bool = False, |
| atmos_static_vars: bool = False, |
| separate_perceiver: tuple[str, ...] = (), |
| modulation_heads: tuple[str, ...] = (), |
| positive_surf_vars: tuple[str, ...] = (), |
| positive_atmos_vars: tuple[str, ...] = (), |
| clamp_at_first_step: bool = False, |
| simulate_indexing_bug: bool = False, |
| stochastic: bool = False, |
| use_updated_lead_time_embedding: bool = False, |
| variable_lead_time: bool = False, |
| rollout_input_clipping: Optional[dict[str, dict[str, Optional[float]]]] = None, |
| output_only_surf_vars: tuple[str, ...] = (), |
| output_only_atmos_vars: tuple[str, ...] = (), |
| ) -> None: |
| """Construct an instance of the model. |
| |
| Args: |
| surf_vars (tuple[str, ...], optional): All surface-level variables supported by the |
| model. |
| static_vars (tuple[str, ...], optional): All static variables supported by the |
| model. |
| atmos_vars (tuple[str, ...], optional): All atmospheric variables supported by the |
| model. |
| window_size (tuple[int, int, int], optional): Vertical height, height, and width of the |
| window of the underlying Swin transformer. |
| encoder_depths (tuple[int, ...], optional): Number of blocks in each encoder layer. |
| encoder_num_heads (tuple[int, ...], optional): Number of attention heads in each encoder |
| layer. The dimensionality doubles after every layer. To keep the dimensionality of |
| every head constant, you want to double the number of heads after every layer. The |
| dimensionality of attention head of the first layer is determined by `embed_dim` |
| divided by the value here. For all cases except one, this is equal to `64`. |
| decoder_depths (tuple[int, ...], optional): Number of blocks in each decoder layer. |
| Generally, you want this to be the reversal of `encoder_depths`. |
| decoder_num_heads (tuple[int, ...], optional): Number of attention heads in each decoder |
| layer. Generally, you want this to be the reversal of `encoder_num_heads`. |
| latent_levels (int, optional): Number of latent pressure levels. |
| patch_size (int, optional): Patch size. |
| embed_dim (int, optional): Patch embedding dimension. |
| num_heads (int, optional): Number of attention heads in the aggregation and |
| deaggregation blocks. The dimensionality of these attention heads will be equal to |
| `embed_dim` divided by this value. |
| mlp_ratio (float, optional): Hidden dim. to embedding dim. ratio for MLPs. |
| drop_rate (float, optional): Drop-out rate. |
| drop_path (float, optional): Drop-path rate. |
| enc_depth (int, optional): Number of Perceiver blocks in the encoder. |
| dec_depth (int, optional): Number of Perceiver blocks in the decoder. |
| dec_mlp_ratio (float, optional): Hidden dim. to embedding dim. ratio for MLPs in the |
| decoder. The embedding dimensionality here is different, which is why this is a |
| separate parameter. |
| perceiver_ln_eps (float, optional): Epsilon in the perceiver layer norm. layers. Used |
| to stabilise the model. |
| max_history_size (int, optional): Maximum number of history steps. You can load |
| checkpoints with a smaller `max_history_size`, but you cannot load checkpoints |
| with a larger `max_history_size`. |
| timestep (timedelta, optional): Timestep of the model. Defaults to 6 hours. |
| stabilise_level_agg (bool, optional): Stabilise the level aggregation by inserting an |
| additional layer normalisation. Defaults to `False`. |
| use_lora (bool, optional): Use LoRA adaptation. |
| lora_steps (int, optional): Use different LoRA adaptation for the first so-many roll-out |
| steps. |
| lora_mode (str, optional): LoRA mode. `"single"` uses the same LoRA for all roll-out |
| steps, `"from_second"` uses the same LoRA from the second roll-out step on, and |
| `"all"` uses a different LoRA for every roll-out step. Defaults to `"single"`. |
| surf_stats (dict[str, tuple[float, float]], optional): For these surface-level |
| variables, adjust the normalisation to the given tuple consisting of a new location |
| and scale. |
| autocast (bool, optional): To reduce memory usage, `torch.autocast` only the backbone |
| to a lower-precision dtype. This is critical to enable fine-tuning. |
| autocast_dtype (torch.dtype, optional): Data type to use when `autocast` is enabled. |
| Defaults to `torch.bfloat16`. |
| use_fp16_safe_attention (bool, optional): Replace |
| :func:`torch.nn.functional.scaled_dot_product_attention` with a manual |
| implementation that clamps intermediate values to prevent float16 overflow. |
| Recommended when running with `autocast_dtype=torch.float16`. |
| Defaults to `False`. |
| level_condition (tuple[int | float, ...], optional): Make the patch embeddings dependent |
| on pressure level. If you want to enable this feature, provide a tuple of all |
| possible pressure levels. |
| dynamic_vars (bool, optional): Use dynamically generated static variables, like time |
| of day. Defaults to `False`. |
| atmos_static_vars (bool, optional): Also concatenate the static variables to the |
| atmospheric variables. Defaults to `False`. |
| separate_perceiver (tuple[str, ...], optional): In the decoder, use a separate Perceiver |
| for specific atmospheric variables. This can be helpful at fine-tuning time to deal |
| with variables that have a significantly different behaviour. If you want to enable |
| this features, set this to the collection of variables that should be run on a |
| separate Perceiver. |
| modulation_heads (tuple[str, ...], optional): Names of every variable for which to |
| enable an additional head, the so-called modulation head, that can be used to |
| predict the difference. |
| positive_surf_vars (tuple[str, ...], optional): Mark these surface-level variables as |
| positive. Clamp them before running them through the encoder, and also clamp them |
| when autoregressively rolling out the model. The variables are not clamped for the |
| first roll-out step. |
| positive_atmos_vars (tuple[str, ...], optional): Mark these atmospheric variables as |
| positive. Clamp them before running them through the encoder, and also clamp them |
| when autoregressively rolling out the model. The variables are not clamped for the |
| first roll-out step. |
| clamp_at_first_step (bool, optional): Clamp the positive variables for the first |
| roll-out step. Should only be used for inference. Defaults to `False`. |
| simulate_indexing_bug (bool, optional): Simulate an indexing bug that's present for the |
| air pollution version of Aurora. This is necessary to obtain numerical equivalence |
| to the original implementation. Defaults to `False`. |
| stochastic (bool, optional): If `True`, enable stochastic mode with noise injection. |
| Defaults to `False`. |
| use_updated_lead_time_embedding (bool, optional): Whether to use the updated lead time |
| embedding with a minimum wavelength of 2 hours. Defaults to `False`. |
| variable_lead_time (bool, optional): If `True`, use per-sample lead times passed |
| via the `lead_times` argument to `forward` (a tensor of shape `(batch,)` in |
| hours) instead of the fixed `timestep`. When enabled, `lead_times` must be |
| provided. |
| Defaults to `False`. |
| rollout_input_clipping (dict[str, dict[str, float]], optional): Per-variable |
| clipping bounds applied to predictions during autoregressive rollout before they |
| become the next input. Keys are variable names (must match `surf_vars` or |
| `atmos_vars`). Values are dicts with optional `"min"` and `"max"` keys. |
| Example: `{"tcc": {"min": 0, "max": 1}}`. Defaults to `None`. |
| output_only_surf_vars (tuple[str, ...], optional): Surface-level variables that the |
| model predicts but that are not present in real input data. These will be |
| zero-padded in the input batch during rollout. Defaults to `()`. |
| output_only_atmos_vars (tuple[str, ...], optional): Atmospheric variables that the |
| model predicts but that are not present in real input data. These will be |
| zero-padded in the input batch during rollout. Defaults to `()`. |
| """ |
| super().__init__() |
| self.surf_vars = surf_vars |
| self.static_vars = static_vars |
| self.atmos_vars = atmos_vars |
| self.patch_size = patch_size |
| self.surf_stats = surf_stats or dict() |
| self.max_history_size = max_history_size |
| self.timestep = timestep |
| self.use_lora = use_lora |
| self.positive_surf_vars = positive_surf_vars |
| self.positive_atmos_vars = positive_atmos_vars |
| self.clamp_at_first_step = clamp_at_first_step |
| self.variable_lead_time = variable_lead_time |
| self.rollout_input_clipping = rollout_input_clipping |
| self.output_only_surf_vars = output_only_surf_vars |
| self.output_only_atmos_vars = output_only_atmos_vars |
|
|
| if self.surf_stats: |
| warnings.warn( |
| f"The normalisation statics for the following surface-level variables are manually " |
| f"adjusted: {', '.join(sorted(self.surf_stats.keys()))}. " |
| f"Please ensure that this is right!", |
| stacklevel=2, |
| ) |
|
|
| self.encoder = Perceiver3DEncoder( |
| surf_vars=surf_vars, |
| static_vars=static_vars, |
| atmos_vars=atmos_vars, |
| patch_size=patch_size, |
| embed_dim=embed_dim, |
| num_heads=num_heads, |
| drop_rate=drop_rate, |
| mlp_ratio=mlp_ratio, |
| head_dim=embed_dim // num_heads, |
| depth=enc_depth, |
| latent_levels=latent_levels, |
| max_history_size=max_history_size, |
| perceiver_ln_eps=perceiver_ln_eps, |
| stabilise_level_agg=stabilise_level_agg, |
| level_condition=level_condition, |
| dynamic_vars=dynamic_vars, |
| atmos_static_vars=atmos_static_vars, |
| simulate_indexing_bug=simulate_indexing_bug, |
| use_updated_lead_time_embedding=use_updated_lead_time_embedding, |
| ) |
|
|
| self.backbone = Swin3DTransformerBackbone( |
| window_size=window_size, |
| encoder_depths=encoder_depths, |
| encoder_num_heads=encoder_num_heads, |
| decoder_depths=decoder_depths, |
| decoder_num_heads=decoder_num_heads, |
| embed_dim=embed_dim, |
| mlp_ratio=mlp_ratio, |
| drop_path_rate=drop_path, |
| attn_drop_rate=drop_rate, |
| drop_rate=drop_rate, |
| use_lora=use_lora, |
| lora_steps=lora_steps, |
| lora_mode=lora_mode, |
| stochastic=stochastic, |
| use_updated_lead_time_embedding=use_updated_lead_time_embedding, |
| ) |
|
|
| self.decoder = Perceiver3DDecoder( |
| surf_vars=surf_vars, |
| atmos_vars=atmos_vars, |
| patch_size=patch_size, |
| |
| embed_dim=embed_dim * 2, |
| head_dim=embed_dim * 2 // num_heads, |
| num_heads=num_heads, |
| depth=dec_depth, |
| |
| |
| mlp_ratio=dec_mlp_ratio, |
| perceiver_ln_eps=perceiver_ln_eps, |
| level_condition=level_condition, |
| separate_perceiver=separate_perceiver, |
| modulation_heads=modulation_heads, |
| ) |
|
|
| if bf16_mode and not autocast: |
| warnings.warn( |
| "`bf16_mode` was removed, because it caused serious issues for gradient " |
| "computation. `bf16_mode` now automatically activates `autocast`, which will not " |
| "save as much memory, but should be much more stable.", |
| stacklevel=2, |
| ) |
| autocast = True |
|
|
| self.autocast = autocast |
| self.autocast_dtype = autocast_dtype |
| self.autocast_encoder = False |
| self.autocast_backbone = autocast |
| self.autocast_decoder = False |
|
|
| |
| if use_fp16_safe_attention: |
| for m in self.modules(): |
| if isinstance(m, (WindowAttention, PerceiverAttention)): |
| m.use_fp16_safe_attention = True |
|
|
| def reset_noise(self) -> None: |
| """Flush the backbone noise cache. |
| |
| See :meth:`Swin3DTransformerBackbone.reset_noise`.""" |
| self.backbone.reset_noise() |
|
|
| def set_noise_accumulation(self, n: int = 0) -> None: |
| """Enable or disable noise caching in the backbone. |
| |
| See :meth:`Swin3DTransformerBackbone.set_noise_accumulation`. |
| |
| Args: |
| n (int): Number of steps for noise accumulation. Disables accumulation if `n=0`. |
| """ |
| self.backbone.set_noise_accumulation(n) |
|
|
| def forward(self, batch: Batch, lead_times: Optional[torch.Tensor] = None) -> Batch: |
| """Forward pass. |
| |
| Args: |
| batch (:class:`aurora.Batch`): Batch to run the model on. |
| lead_times (:class:`torch.Tensor`, optional): Per-sample lead times of shape |
| `(batch,)` in hours. Required when the model was configured with |
| `variable_lead_time=True`. Ignored otherwise. |
| |
| Returns: |
| :class:`Batch`: Prediction for the batch. |
| """ |
| batch = self.batch_transform_hook(batch) |
|
|
| |
| p = next(self.parameters()) |
| batch = batch.type(p.dtype) |
| batch = self._pre_norm_hook(batch) |
| batch = batch.normalise(surf_stats=self.surf_stats) |
| batch = batch.crop(patch_size=self.patch_size) |
| batch = batch.to(p.device) |
|
|
| H, W = batch.spatial_shape |
| patch_res = ( |
| self.encoder.latent_levels, |
| H // self.encoder.patch_size, |
| W // self.encoder.patch_size, |
| ) |
|
|
| |
| B, T = next(iter(batch.surf_vars.values())).shape[:2] |
| batch = dataclasses.replace( |
| batch, |
| static_vars={k: v[None, None].repeat(B, T, 1, 1) for k, v in batch.static_vars.items()}, |
| ) |
|
|
| |
| |
| transformed_batch = batch |
|
|
| |
| if self.positive_surf_vars: |
| transformed_batch = dataclasses.replace( |
| transformed_batch, |
| surf_vars={ |
| k: v.clamp(min=0) if k in self.positive_surf_vars else v |
| for k, v in batch.surf_vars.items() |
| }, |
| ) |
| if self.positive_atmos_vars: |
| transformed_batch = dataclasses.replace( |
| transformed_batch, |
| atmos_vars={ |
| k: v.clamp(min=0) if k in self.positive_atmos_vars else v |
| for k, v in batch.atmos_vars.items() |
| }, |
| ) |
|
|
| transformed_batch = self._pre_encoder_hook(transformed_batch) |
|
|
| |
| if self.variable_lead_time: |
| if lead_times is None: |
| raise ValueError( |
| "`variable_lead_time=True` but `lead_times` is `None`. " |
| "Please provide a `lead_times` tensor of shape `(batch,)` in hours." |
| ) |
| lead_times = lead_times.to(device=p.device, dtype=p.dtype) |
| else: |
| lead_hours = self.timestep.total_seconds() / 3600 |
| lead_times = torch.full((B,), lead_hours, device=p.device, dtype=p.dtype) |
|
|
| if torch.cuda.is_available(): |
| device_type = "cuda" |
| elif torch.xpu.is_available(): |
| device_type = "xpu" |
| else: |
| device_type = "cpu" |
| autocast = torch.autocast(device_type=device_type, dtype=self.autocast_dtype) |
| context_encoder = autocast if self.autocast_encoder else contextlib.nullcontext() |
| context_backbone = autocast if self.autocast_backbone else contextlib.nullcontext() |
| context_decoder = autocast if self.autocast_decoder else contextlib.nullcontext() |
|
|
| with context_encoder: |
| x = self.encoder( |
| transformed_batch, |
| lead_times=lead_times, |
| ) |
| with context_backbone: |
| x = self.backbone( |
| x, |
| lead_times=lead_times, |
| patch_res=patch_res, |
| rollout_step=batch.metadata.rollout_step, |
| ) |
| with context_decoder: |
| pred = self.decoder( |
| x, |
| batch, |
| lead_times=lead_times, |
| patch_res=patch_res, |
| ) |
|
|
| |
| pred = dataclasses.replace( |
| pred, |
| static_vars={k: v[0, 0] for k, v in batch.static_vars.items()}, |
| ) |
|
|
| |
| pred = dataclasses.replace( |
| pred, |
| surf_vars={k: v[:, None] for k, v in pred.surf_vars.items()}, |
| atmos_vars={k: v[:, None] for k, v in pred.atmos_vars.items()}, |
| ) |
|
|
| pred = self._post_decoder_hook(batch, pred) |
|
|
| |
| clamp_at_rollout_step = ( |
| pred.metadata.rollout_step >= 1 |
| if self.clamp_at_first_step |
| else pred.metadata.rollout_step > 1 |
| ) |
| if self.positive_surf_vars and clamp_at_rollout_step: |
| pred = dataclasses.replace( |
| pred, |
| surf_vars={ |
| k: v.clamp(min=0) if k in self.positive_surf_vars else v |
| for k, v in pred.surf_vars.items() |
| }, |
| ) |
| if self.positive_atmos_vars and clamp_at_rollout_step: |
| pred = dataclasses.replace( |
| pred, |
| atmos_vars={ |
| k: v.clamp(min=0) if k in self.positive_atmos_vars else v |
| for k, v in pred.atmos_vars.items() |
| }, |
| ) |
|
|
| |
| pred = pred.type(torch.float32) |
| pred = pred.unnormalise(surf_stats=self.surf_stats) |
|
|
| pred = self._post_unnorm_hook(batch, pred) |
|
|
| return pred |
|
|
| def batch_transform_hook(self, batch: Batch) -> Batch: |
| """Transform the batch right after receiving it and before normalisation. |
| |
| This function should be idempotent. |
| """ |
| return batch |
|
|
| def _pre_encoder_hook(self, batch: Batch) -> Batch: |
| """Transform the batch before it goes through the encoder.""" |
| return batch |
|
|
| def _pre_norm_hook(self, batch: Batch) -> Batch: |
| """Transform the batch before normalisation. |
| |
| This is called automatically in :meth:`forward` right before :meth:`Batch.normalise`. Unlike |
| :meth:`batch_transform_hook`, this hook is *not* called separately in rollout, so |
| non-idempotent transforms (e.g. log-scaling) belong here. |
| """ |
| return batch |
|
|
| def _post_decoder_hook(self, batch: Batch, pred: Batch) -> Batch: |
| """Transform the prediction right after the decoder.""" |
| return pred |
|
|
| def _post_unnorm_hook(self, batch: Batch, pred: Batch) -> Batch: |
| """Transform the prediction after un-normalisation, in physical space. |
| |
| Subclasses can override this to apply post-processing that must operate on un-normalised |
| (physical) values, such as inverse log-scaling or recomputing prescribed channels. |
| """ |
| return pred |
|
|
| def apply_rollout_input_clipping(self, pred: Batch) -> Batch: |
| """Clamp specified variables according to `rollout_input_clipping`. |
| |
| This is intended to be called during autoregressive rollout *before* feeding a prediction |
| back as input, so that the unclipped prediction remains available for loss computation |
| during training. To minimize any other changes to the data flow from models prior to V1p5, |
| this is not called automatically in :meth:`forward`. |
| """ |
| if not self.rollout_input_clipping: |
| return pred |
|
|
| clipped_surf = dict(pred.surf_vars) |
| clipped_atmos = dict(pred.atmos_vars) |
|
|
| for var_name, bounds in self.rollout_input_clipping.items(): |
| lo = bounds.get("min") |
| hi = bounds.get("max") |
| if var_name in clipped_surf: |
| v = clipped_surf[var_name] |
| if lo is not None: |
| v = v.clamp(min=lo) |
| if hi is not None: |
| v = v.clamp(max=hi) |
| clipped_surf[var_name] = v |
| if var_name in clipped_atmos: |
| v = clipped_atmos[var_name] |
| if lo is not None: |
| v = v.clamp(min=lo) |
| if hi is not None: |
| v = v.clamp(max=hi) |
| clipped_atmos[var_name] = v |
|
|
| return dataclasses.replace(pred, surf_vars=clipped_surf, atmos_vars=clipped_atmos) |
|
|
| def load_checkpoint( |
| self, |
| repo: Optional[str] = None, |
| name: Optional[str] = None, |
| revision: Optional[str] = None, |
| strict: bool = True, |
| ) -> None: |
| """Load a checkpoint from HuggingFace. |
| |
| Args: |
| repo (str, optional): Name of the repository of the form `user/repo`. |
| name (str, optional): Path to the checkpoint relative to the root of the repository, |
| e.g. `checkpoint.cpkt`. |
| revision (str, optional): Version hash of the Huggingface git repository commit. |
| strict (bool, optional): Error if the model parameters are not exactly equal to the |
| parameters in the checkpoint. Defaults to `True`. |
| """ |
| repo = repo or self.default_checkpoint_repo |
| name = name or self.default_checkpoint_name |
| revision = revision or self.default_checkpoint_revision |
| path = hf_hub_download(repo_id=repo, filename=name, revision=revision) |
| self.load_checkpoint_local(path, strict=strict) |
|
|
| def load_checkpoint_local(self, path: str, strict: bool = True) -> None: |
| """Load a checkpoint directly from a file. |
| |
| Args: |
| path (str): Path to the checkpoint. |
| strict (bool, optional): Error if the model parameters are not exactly equal to the |
| parameters in the checkpoint. Defaults to `True`. |
| """ |
| |
| device = next(self.parameters()).device |
| d = torch.load(path, map_location=device, weights_only=True) |
|
|
| d = self._adapt_checkpoint(d) |
|
|
| |
| current_history_size = d["encoder.surf_token_embeds.weights.2t"].shape[2] |
| if self.max_history_size > current_history_size: |
| self.adapt_checkpoint_max_history_size(d) |
| elif self.max_history_size < current_history_size: |
| raise AssertionError( |
| f"Cannot load checkpoint with `max_history_size` {current_history_size} " |
| f"into model with `max_history_size` {self.max_history_size}." |
| ) |
|
|
| self.load_state_dict(d, strict=strict) |
|
|
| def _adapt_checkpoint(self, d: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]: |
| """Adapt an existing checkpoint to make it compatible with the current version of the model. |
| |
| Args: |
| d (dict[str, torch.Tensor]): Checkpoint. |
| |
| Return: |
| dict[str, torch.Tensor]: Adapted checkpoint. |
| """ |
| return _adapt_checkpoint_pretrained(self.patch_size, d) |
|
|
| def adapt_checkpoint_max_history_size(self, checkpoint: dict[str, torch.Tensor]) -> None: |
| """Adapt a checkpoint with smaller `max_history_size` to a model with a larger |
| `max_history_size` than the current model. |
| |
| If a checkpoint was trained with a larger `max_history_size` than the current model, |
| this function will assert fail to prevent loading the checkpoint. This is to |
| prevent loading a checkpoint which will likely cause the checkpoint to degrade its |
| performance. |
| |
| This implementation copies weights from the checkpoint to the model and fills zeros |
| for the new history width dimension. It mutates `checkpoint`. |
| """ |
| for name, weight in list(checkpoint.items()): |
| |
| enc_surf_embedding = name.startswith("encoder.surf_token_embeds.weights.") |
| enc_atmos_embedding = name.startswith("encoder.atmos_token_embeds.weights.") |
| if enc_surf_embedding or enc_atmos_embedding: |
| |
| |
| if not (weight.shape[2] <= self.max_history_size): |
| raise AssertionError( |
| f"Cannot load checkpoint with `max_history_size` {weight.shape[2]} " |
| f"into model with `max_history_size` {self.max_history_size}." |
| ) |
|
|
| |
| new_weight = torch.zeros( |
| (weight.shape[0], 1, self.max_history_size, weight.shape[3], weight.shape[4]), |
| device=weight.device, |
| dtype=weight.dtype, |
| ) |
| |
| |
| new_weight[:, :, : weight.shape[2]] = weight |
|
|
| checkpoint[name] = new_weight |
|
|
| def configure_activation_checkpointing( |
| self, |
| module_names: tuple[str, ...] = ( |
| "Basic3DDecoderLayer", |
| "Basic3DEncoderLayer", |
| "LinearPatchReconstruction", |
| "Perceiver3DDecoder", |
| "Perceiver3DEncoder", |
| "Swin3DTransformerBackbone", |
| "Swin3DTransformerBlock", |
| ), |
| ) -> None: |
| """Configure activation checkpointing. |
| |
| This is required in order to compute gradients without running out of memory. |
| |
| Args: |
| module_names (tuple[str, ...], optional): Names of the modules to checkpoint |
| on. |
| |
| Raises: |
| RuntimeError: If any module specifies in `module_names` was not found and |
| thus could not be checkpointed. |
| """ |
|
|
| found: set[str] = set() |
|
|
| def check(x: torch.nn.Module) -> bool: |
| name = x.__class__.__name__ |
| if name in module_names: |
| found.add(name) |
| return True |
| else: |
| return False |
|
|
| apply_activation_checkpointing(self, check_fn=check) |
|
|
| if found != set(module_names): |
| raise RuntimeError( |
| f"Could not checkpoint on the following modules: " |
| f"{', '.join(sorted(set(module_names) - found))}." |
| ) |
|
|
|
|
| class AuroraPretrained(Aurora): |
| """Pretrained version of Aurora.""" |
|
|
| default_checkpoint_name = "aurora-0.25-pretrained.ckpt" |
| default_checkpoint_revision = "0be7e57c685dac86b78c4a19a3ab149d13c6a3dd" |
|
|
| def __init__( |
| self, |
| *, |
| use_lora: bool = False, |
| **kw_args, |
| ) -> None: |
| super().__init__( |
| use_lora=use_lora, |
| **kw_args, |
| ) |
|
|
|
|
| class AuroraSmallPretrained(Aurora): |
| """Small pretrained version of Aurora. |
| |
| Should only be used for debugging. |
| """ |
|
|
| default_checkpoint_name = "aurora-0.25-small-pretrained.ckpt" |
| default_checkpoint_revision = "0be7e57c685dac86b78c4a19a3ab149d13c6a3dd" |
|
|
| def __init__( |
| self, |
| *, |
| encoder_depths: tuple[int, ...] = (2, 6, 2), |
| encoder_num_heads: tuple[int, ...] = (4, 8, 16), |
| decoder_depths: tuple[int, ...] = (2, 6, 2), |
| decoder_num_heads: tuple[int, ...] = (16, 8, 4), |
| embed_dim: int = 256, |
| num_heads: int = 8, |
| use_lora: bool = False, |
| **kw_args, |
| ) -> None: |
| super().__init__( |
| encoder_depths=encoder_depths, |
| encoder_num_heads=encoder_num_heads, |
| decoder_depths=decoder_depths, |
| decoder_num_heads=decoder_num_heads, |
| embed_dim=embed_dim, |
| num_heads=num_heads, |
| use_lora=use_lora, |
| **kw_args, |
| ) |
|
|
|
|
| AuroraSmall = AuroraSmallPretrained |
|
|
|
|
| class Aurora12hPretrained(Aurora): |
| """Pretrained version of Aurora with time step 12 hours.""" |
|
|
| default_checkpoint_name = "aurora-0.25-12h-pretrained.ckpt" |
| default_checkpoint_revision = "15e76e47b65bf4b28fd2246b7b5b951d6e2443b9" |
|
|
| def __init__( |
| self, |
| *, |
| timestep: timedelta = timedelta(hours=12), |
| use_lora: bool = False, |
| **kw_args, |
| ) -> None: |
| super().__init__( |
| timestep=timestep, |
| use_lora=use_lora, |
| **kw_args, |
| ) |
|
|
|
|
| class AuroraHighRes(Aurora): |
| """High-resolution version of Aurora.""" |
|
|
| default_checkpoint_name = "aurora-0.1-finetuned.ckpt" |
| default_checkpoint_revision = "0be7e57c685dac86b78c4a19a3ab149d13c6a3dd" |
|
|
| def __init__( |
| self, |
| *, |
| patch_size: int = 10, |
| encoder_depths: tuple[int, ...] = (6, 8, 8), |
| decoder_depths: tuple[int, ...] = (8, 8, 6), |
| **kw_args, |
| ) -> None: |
| super().__init__( |
| patch_size=patch_size, |
| encoder_depths=encoder_depths, |
| decoder_depths=decoder_depths, |
| **kw_args, |
| ) |
|
|
|
|
| class AuroraAirPollution(Aurora): |
| """Fine-tuned version of Aurora for air pollution.""" |
|
|
| default_checkpoint_name = "aurora-0.4-air-pollution.ckpt" |
| default_checkpoint_revision = "1764d5630a53d3d7a7d169ca335236fc343e4bfc" |
|
|
| _predict_difference_history_dim_lookup = { |
| "pm1": 0, |
| "pm2p5": 0, |
| "pm10": 0, |
| "co": 1, |
| "tcco": 1, |
| "no": 0, |
| "tc_no": 0, |
| "no2": 0, |
| "tcno2": 0, |
| "so2": 1, |
| "tcso2": 1, |
| "go3": 1, |
| "gtco3": 1, |
| } |
| """dict[str, int]: For every variable that we want to predict the difference for, the index |
| into the history dimension that should be used when predicting the difference.""" |
|
|
| def __init__( |
| self, |
| *, |
| surf_vars: tuple[str, ...] = ( |
| ("2t", "10u", "10v", "msl") |
| + ("pm1", "pm2p5", "pm10", "tcco", "tc_no", "tcno2", "gtco3", "tcso2") |
| ), |
| static_vars: tuple[str, ...] = ( |
| ("lsm", "z", "slt") |
| + ("static_ammonia", "static_ammonia_log", "static_co", "static_co_log") |
| + ("static_nox", "static_nox_log", "static_so2", "static_so2_log") |
| ), |
| atmos_vars: tuple[str, ...] = ("z", "u", "v", "t", "q", "co", "no", "no2", "go3", "so2"), |
| patch_size: int = 3, |
| timestep: timedelta = timedelta(hours=12), |
| level_condition: Optional[tuple[int | float, ...]] = ( |
| (50, 100, 150, 200, 250, 300, 400, 500, 600, 700, 850, 925, 1000) |
| ), |
| dynamic_vars: bool = True, |
| atmos_static_vars: bool = True, |
| separate_perceiver: tuple[str, ...] = ("co", "no", "no2", "go3", "so2"), |
| modulation_heads: tuple[str, ...] = tuple(_predict_difference_history_dim_lookup.keys()), |
| positive_surf_vars: tuple[str, ...] = ( |
| ("pm1", "pm2p5", "pm10", "tcco", "tc_no", "tcno2", "gtco3", "tcso2") |
| ), |
| positive_atmos_vars: tuple[str, ...] = ("co", "no", "no2", "go3", "so2"), |
| simulate_indexing_bug: bool = True, |
| **kw_args, |
| ) -> None: |
| super().__init__( |
| surf_vars=surf_vars, |
| static_vars=static_vars, |
| atmos_vars=atmos_vars, |
| patch_size=patch_size, |
| timestep=timestep, |
| level_condition=level_condition, |
| dynamic_vars=dynamic_vars, |
| atmos_static_vars=atmos_static_vars, |
| separate_perceiver=separate_perceiver, |
| modulation_heads=modulation_heads, |
| positive_surf_vars=positive_surf_vars, |
| positive_atmos_vars=positive_atmos_vars, |
| simulate_indexing_bug=simulate_indexing_bug, |
| **kw_args, |
| ) |
|
|
| self.surf_feature_combiner = torch.nn.ParameterDict( |
| {v: nn.Linear(2, 1, bias=True) for v in self.positive_surf_vars} |
| ) |
| self.atmos_feature_combiner = torch.nn.ParameterDict( |
| {v: nn.Linear(2, 1, bias=True) for v in self.positive_atmos_vars} |
| ) |
| for p in (*self.surf_feature_combiner.values(), *self.atmos_feature_combiner.values()): |
| nn.init.constant_(p.weight, 0.5) |
| nn.init.zeros_(p.bias) |
|
|
| def _pre_encoder_hook(self, batch: Batch) -> Batch: |
| |
| |
|
|
| eps = 1e-4 |
| divisor = -np.log(eps) |
|
|
| def _transform(z: torch.Tensor, feature_combiner: nn.Module) -> torch.Tensor: |
| return feature_combiner( |
| torch.stack( |
| [ |
| z.clamp(min=0, max=2.5), |
| (torch.log(z.clamp(min=eps)) - np.log(eps)) / divisor, |
| ], |
| dim=-1, |
| ) |
| )[..., 0] |
|
|
| return dataclasses.replace( |
| batch, |
| surf_vars={ |
| k: _transform(v, self.surf_feature_combiner[k]) |
| if k in self.surf_feature_combiner |
| else v |
| for k, v in batch.surf_vars.items() |
| }, |
| atmos_vars={ |
| k: _transform(v, self.atmos_feature_combiner[k]) |
| if k in self.atmos_feature_combiner |
| else v |
| for k, v in batch.atmos_vars.items() |
| }, |
| ) |
|
|
| def _post_decoder_hook(self, batch: Batch, pred: Batch) -> Batch: |
| |
| |
| |
|
|
| dim_lookup = AuroraAirPollution._predict_difference_history_dim_lookup |
|
|
| def _transform( |
| prev: dict[str, torch.Tensor], |
| model: dict[str, torch.Tensor], |
| name: str, |
| ) -> torch.Tensor: |
| if name in dim_lookup: |
| return model[name] + (1 + model[f"{name}_mod"]) * prev[name][:, dim_lookup[name]] |
| else: |
| return model[name] |
|
|
| pred = dataclasses.replace( |
| pred, |
| surf_vars={k: _transform(batch.surf_vars, pred.surf_vars, k) for k in batch.surf_vars}, |
| atmos_vars={ |
| k: _transform(batch.atmos_vars, pred.atmos_vars, k) for k in batch.atmos_vars |
| }, |
| ) |
|
|
| |
| |
| if self.use_lora: |
| parts: list[torch.Tensor] = [] |
| for i, level in enumerate(pred.metadata.atmos_levels): |
| section = pred.atmos_vars["so2"][..., i, :, :] |
| if level >= 850: |
| section = section.clamp(max=1) |
| parts.append(section) |
| pred.atmos_vars["so2"] = torch.stack(parts, dim=-3) |
|
|
| return pred |
|
|
| def _adapt_checkpoint(self, d: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]: |
| d = Aurora._adapt_checkpoint(self, d) |
| d = _adapt_checkpoint_air_pollution(self.patch_size, d) |
| return d |
|
|
|
|
| class AuroraWave(Aurora): |
| """Version of Aurora fined-tuned to HRES-WAM ocean wave data.""" |
|
|
| default_checkpoint_name = "aurora-0.25-wave.ckpt" |
| default_checkpoint_revision = "74598e8c65d53a96077c08bb91acdfa5525340c9" |
|
|
| def __init__( |
| self, |
| *, |
| surf_vars: tuple[str, ...] = ( |
| ("2t", "10u", "10v", "msl") |
| + ("swh", "mwd", "mwp", "pp1d", "shww", "mdww", "mpww", "shts", "mdts", "mpts") |
| + ("swh1", "mwd1", "mwp1", "swh2", "mwd2", "mwp2", "wind", "10u_wave", "10v_wave") |
| ), |
| static_vars: tuple[str, ...] = ("lsm", "z", "slt", "wmb", "lat_mask"), |
| lora_mode: LoRAMode = "from_second", |
| stabilise_level_agg: bool = True, |
| density_channel_surf_vars: tuple[str, ...] = ( |
| ("swh", "mwd", "mwp", "pp1d", "shww", "mdww", "mpww", "shts", "mdts", "mpts") |
| + ("swh1", "mwd1", "mwp1", "swh2", "mwd2", "mwp2", "wind", "10u_wave", "10v_wave") |
| ), |
| angle_surf_vars: tuple[str, ...] = ("mwd", "mdww", "mdts", "mwd1", "mwd2"), |
| **kw_args, |
| ) -> None: |
| |
| supplemented_surf_vars: tuple[str, ...] = () |
| for name in surf_vars: |
| if name in angle_surf_vars: |
| supplemented_surf_vars += (f"{name}_sin", f"{name}_cos") |
| else: |
| supplemented_surf_vars += (name,) |
| if name in density_channel_surf_vars: |
| supplemented_surf_vars += (f"{name}_density",) |
|
|
| super().__init__( |
| surf_vars=supplemented_surf_vars, |
| static_vars=static_vars, |
| lora_mode=lora_mode, |
| stabilise_level_agg=stabilise_level_agg, |
| **kw_args, |
| ) |
|
|
| self.density_channel_surf_vars = density_channel_surf_vars |
| self.angle_surf_vars = angle_surf_vars |
|
|
| def _adapt_checkpoint(self, d: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]: |
| d = Aurora._adapt_checkpoint(self, d) |
| d = _adapt_checkpoint_wave(self.patch_size, d) |
| return d |
|
|
| def batch_transform_hook(self, batch: Batch) -> Batch: |
| |
| batch = dataclasses.replace(batch, surf_vars=dict(batch.surf_vars)) |
|
|
| |
| |
| if "dwi" in batch.surf_vars and "wind" in batch.surf_vars: |
| |
| u_wave = -batch.surf_vars["wind"] * torch.sin(torch.deg2rad(batch.surf_vars["dwi"])) |
| v_wave = -batch.surf_vars["wind"] * torch.cos(torch.deg2rad(batch.surf_vars["dwi"])) |
|
|
| |
| batch.surf_vars["10u_wave"] = u_wave |
| batch.surf_vars["10v_wave"] = v_wave |
| del batch.surf_vars["dwi"] |
|
|
| |
| |
| if batch.metadata.rollout_step == 0: |
| for name_sh, other_wave_components in [ |
| ("swh", ("mwd", "mwp", "pp1d")), |
| ("shww", ("mdww", "mpww")), |
| ("shts", ("mdts", "mdts")), |
| ("swh1", ("mwd1", "mwp1")), |
| ("swh2", ("mwd2", "mwp2")), |
| ]: |
| mask = batch.surf_vars[name_sh] < 1e-4 |
| if mask.sum() > 0: |
| for name in (name_sh,) + other_wave_components: |
| x = batch.surf_vars[name].clone() |
| x[mask] = np.nan |
| batch.surf_vars[name] = x |
| |
| if name not in {"mwd", "mdww", "mdts", "mwd1", "mwd2"}: |
| assert (batch.surf_vars[name] < 1e-4).sum() == 0 |
|
|
| return batch |
|
|
| def _pre_encoder_hook(self, batch: Batch) -> Batch: |
| for name in list(batch.surf_vars): |
| x = batch.surf_vars[name] |
|
|
| |
| if name in self.density_channel_surf_vars and f"{name}_density" not in batch.surf_vars: |
| batch.surf_vars[f"{name}_density"] = (~torch.isnan(x)).float() |
| batch.surf_vars[name] = x.nan_to_num(0) |
|
|
| |
| sin_cos_present = f"{name}_sin" in batch.surf_vars and f"{name}_cos" in batch.surf_vars |
| if name in self.angle_surf_vars and not sin_cos_present: |
| batch.surf_vars[f"{name}_sin"] = torch.sin(torch.deg2rad(x)).nan_to_num(0) |
| batch.surf_vars[f"{name}_cos"] = torch.cos(torch.deg2rad(x)).nan_to_num(0) |
| del batch.surf_vars[name] |
|
|
| return batch |
|
|
| def _post_decoder_hook(self, batch: Batch, pred: Batch) -> Batch: |
| wmb_mask = pred.static_vars["wmb"] > 0 |
|
|
| |
| for name in self.angle_surf_vars: |
| if f"{name}_sin" in pred.surf_vars and f"{name}_cos" in pred.surf_vars: |
| sin = pred.surf_vars[f"{name}_sin"] |
| cos = pred.surf_vars[f"{name}_cos"] |
| pred.surf_vars[name] = torch.rad2deg(torch.atan2(sin, cos)) % 360 |
| del pred.surf_vars[f"{name}_sin"] |
| del pred.surf_vars[f"{name}_cos"] |
|
|
| |
| |
| for name in self.density_channel_surf_vars: |
| if name in pred.surf_vars: |
| density = torch.sigmoid(pred.surf_vars[f"{name}_density"]) * wmb_mask |
| data = pred.surf_vars[name] * wmb_mask |
| data[density < 0.5] = np.nan |
| pred.surf_vars[name] = data |
| del pred.surf_vars[f"{name}_density"] |
|
|
| return pred |
|
|
|
|
| class AuroraV1p5(Aurora): |
| """Aurora 1.5 with expanded surface variables, variable lead-time support, and insolation. |
| |
| This variant was trained with an extended set of surface variables (26 total), additional static |
| fields, and prescribed solar insolation as an input channel. It supports variable lead-time |
| embeddings, enabling sub-6-hour prediction steps. Seven surface variables are output-only (not |
| present in the real input data) and are zero-padded during autoregressive rollout. |
| """ |
|
|
| default_checkpoint_repo = "ikwessel/aurora-1.5" |
| default_checkpoint_name = "aurora-0.25-v1.5.ckpt" |
| default_checkpoint_revision = "9751bb56e8e4a0f0a780e3cbe978f4c721e12bc7" |
|
|
| def __init__( |
| self, |
| *, |
| surf_vars: tuple[str, ...] = ( |
| ("2t", "10u", "10v", "msl", "2d", "tcwv", "tcc", "100u", "100v", "sp", "lcc", "mcc") |
| + ("hcc", "skt", "stl1", "swvl1", "ci", "scaled_sd", "i10fg", "blh", "uvb_1h") |
| + ("ssrd_1h", "ttr_1h", "scaled_tp_1h", "scaled_sf_1h", "insolation") |
| ), |
| static_vars: tuple[str, ...] = ( |
| ("lsm", "z", "anor", "isor", "cvh", "cl", "dl", "cvl", "slor", "slt_0", "slt_1") |
| + ("slt_2", "slt_3", "slt_4", "slt_5", "slt_6", "slt_7", "sdfor", "sdor", "tvh_0") |
| + ("tvh_18", "tvh_19", "tvh_3", "tvh_4", "tvh_5", "tvh_6", "tvl_0", "tvl_1", "tvl_10") |
| + ("tvl_11", "tvl_13", "tvl_16", "tvl_17", "tvl_2", "tvl_7", "tvl_9") |
| ), |
| atmos_vars: tuple[str, ...] = ("z", "u", "v", "t", "q"), |
| output_only_surf_vars: tuple[str, ...] = ( |
| ("i10fg", "blh", "uvb_1h", "ssrd_1h", "ttr_1h", "scaled_tp_1h", "scaled_sf_1h") |
| ), |
| rollout_input_clipping: Optional[dict[str, dict[str, Optional[float]]]] = None, |
| variable_lead_time: bool = True, |
| use_updated_lead_time_embedding: bool = True, |
| use_lora: bool = False, |
| use_fp16_safe_attention: bool = True, |
| autocast: bool = True, |
| autocast_dtype: torch.dtype = torch.float16, |
| **kw_args, |
| ) -> None: |
| |
| |
| rollout_input_clipping = rollout_input_clipping or {} |
| if "tcwv" not in rollout_input_clipping: |
| rollout_input_clipping["tcwv"] = {"min": 0.0, "max": None} |
| if "tcc" not in rollout_input_clipping: |
| rollout_input_clipping["tcc"] = {"min": 0.0, "max": 1.0} |
| if "lcc" not in rollout_input_clipping: |
| rollout_input_clipping["lcc"] = {"min": 0.0, "max": 1.0} |
| if "mcc" not in rollout_input_clipping: |
| rollout_input_clipping["mcc"] = {"min": 0.0, "max": 1.0} |
| if "hcc" not in rollout_input_clipping: |
| rollout_input_clipping["hcc"] = {"min": 0.0, "max": 1.0} |
| if "swvl1" not in rollout_input_clipping: |
| rollout_input_clipping["swvl1"] = {"min": 0.0, "max": 70.0} |
| if "ci" not in rollout_input_clipping: |
| rollout_input_clipping["ci"] = {"min": 0.0, "max": 1.0} |
| if "scaled_sd" not in rollout_input_clipping: |
| rollout_input_clipping["scaled_sd"] = {"min": 0.0, "max": 10.0} |
|
|
| super().__init__( |
| surf_vars=surf_vars, |
| static_vars=static_vars, |
| atmos_vars=atmos_vars, |
| output_only_surf_vars=output_only_surf_vars, |
| rollout_input_clipping=rollout_input_clipping, |
| variable_lead_time=variable_lead_time, |
| use_updated_lead_time_embedding=use_updated_lead_time_embedding, |
| use_lora=use_lora, |
| use_fp16_safe_attention=use_fp16_safe_attention, |
| autocast=autocast, |
| autocast_dtype=autocast_dtype, |
| **kw_args, |
| ) |
| self.autocast_encoder = autocast |
| self.autocast_backbone = autocast |
| self.autocast_decoder = autocast |
|
|
| |
| self.log_transformed_surf_vars = tuple(v for v in self.surf_vars if v.startswith("scaled_")) |
|
|
| def _pre_encoder_hook(self, batch: Batch) -> Batch: |
| """Zero-pad output-only variables. |
| |
| Output-only variables are predicted by the model but are not present in real input data. |
| They are added as zero tensors so the encoder receives the correct number of channels. |
| Mutates `batch.surf_vars` / `batch.atmos_vars` in place so that both `batch` and |
| `transformed_batch` in the caller see the new keys. Zero tensors are added post- |
| normalization. |
| """ |
| for var in self.output_only_surf_vars: |
| ref = next(iter(batch.surf_vars.values())) |
| batch.surf_vars[var] = torch.zeros_like(ref) |
| for var in self.output_only_atmos_vars: |
| ref = next(iter(batch.atmos_vars.values())) |
| batch.atmos_vars[var] = torch.zeros_like(ref) |
| return batch |
|
|
| def _pre_norm_hook(self, batch: Batch) -> Batch: |
| """Apply log-transform to scaled surface variables before normalisation.""" |
| return dataclasses.replace( |
| batch, |
| surf_vars={ |
| k: log_transform(v) if k in self.log_transformed_surf_vars else v |
| for k, v in batch.surf_vars.items() |
| }, |
| ) |
|
|
| def _post_unnorm_hook(self, batch: Batch, pred: Batch) -> Batch: |
| """Apply inverse log-transform and recompute prescribed insolation.""" |
| pred = dataclasses.replace( |
| pred, |
| surf_vars={ |
| k: log_untransform(v) if k in self.log_transformed_surf_vars else v |
| for k, v in pred.surf_vars.items() |
| }, |
| ) |
| pred = self._update_insolation(pred) |
| return pred |
|
|
| def _update_insolation(self, pred: Batch) -> Batch: |
| """Recompute prescribed insolation for the prediction's valid time.""" |
| if "insolation" not in pred.surf_vars: |
| return pred |
|
|
| lat_np = pred.metadata.lat.cpu().numpy().astype(np.float32) |
| lon_np = pred.metadata.lon.cpu().numpy().astype(np.float32) |
|
|
| sol_all = [] |
| for t in pred.metadata.time: |
| sol = insolation([t], lat_np, lon_np, enforce_2d=True) |
| sol_all.append(sol[0]) |
| sol_tensor = torch.tensor( |
| np.stack(sol_all, axis=0), |
| dtype=pred.surf_vars["insolation"].dtype, |
| device=pred.surf_vars["insolation"].device, |
| ) |
| |
| sol_tensor = sol_tensor[:, None, :, :] |
|
|
| return dataclasses.replace( |
| pred, |
| surf_vars={ |
| k: (sol_tensor if k == "insolation" else v) for k, v in pred.surf_vars.items() |
| }, |
| ) |
|
|
| def _adapt_checkpoint(self, d: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]: |
| return _adapt_checkpoint_v1p5( |
| self.patch_size, |
| self.surf_vars, |
| self.static_vars, |
| self.atmos_vars, |
| d, |
| ) |
|
|
|
|
| class AuroraV1p5Ensemble(AuroraV1p5): |
| """Aurora 1.5 ensemble version with stochastic noise injection.""" |
|
|
| default_checkpoint_name = "aurora-0.25-v1.5-ensemble.ckpt" |
| default_checkpoint_revision = "9751bb56e8e4a0f0a780e3cbe978f4c721e12bc7" |
|
|
| def __init__( |
| self, |
| *, |
| stochastic: bool = True, |
| **kw_args, |
| ) -> None: |
| super().__init__( |
| stochastic=stochastic, |
| **kw_args, |
| ) |
|
|