-
Notifications
You must be signed in to change notification settings - Fork 70
Implementation CRPS decoder #2623
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: develop
Are you sure you want to change the base?
Changes from all commits
f086728
374d882
99c38f2
ab61f4c
361569e
f2ab1a6
62da414
8c48624
4db1f0f
c96c20d
378bdd2
b29a307
9443db0
28d549d
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -344,6 +344,23 @@ def __init__(self, cf: Config, sources_size, targets_num_channels, targets_coord | |
| self.aux_token_idxs = list(range(cf.num_register_tokens + cf.num_class_tokens)) | ||
| self.num_aux_tokens = cf.num_register_tokens + cf.num_class_tokens | ||
|
|
||
| self.ens_latent_perturb = cf.get("decoder_ens_latent_perturbation") | ||
|
|
||
| assert self.cf.ae_local_num_queries == 1, "ae_local_num_queries > 1 is deprecated." | ||
|
|
||
| # Latent-perturbation noise scale (learnable or fixed) | ||
| self.use_latent_perturbation = ( | ||
| self.ens_latent_perturb is not None | ||
| and self.ens_latent_perturb.get("num_members", 0) >= 1 | ||
| ) | ||
|
|
||
| self.latent_perturbation_log_sigma = None | ||
| if self.use_latent_perturbation: | ||
| sigma_learnable = self.ens_latent_perturb.get("sigma_learnable", True) | ||
| self.latent_perturbation_log_sigma = nn.Parameter( | ||
| torch.zeros(1), requires_grad=sigma_learnable | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. So here we initialise this as 0, hence sigma is 1? And then it only takes a sensible value in reset_parameters? Is reset_parameters always called? What happens when we continue or finetune a run? Why don't we just read sigma_init here? |
||
| ) | ||
|
|
||
| def _create_latent_pred_head( | ||
| self, global_cfg, name, loss_cfg, use_class_token, use_patch_token | ||
| ): | ||
|
|
@@ -578,6 +595,26 @@ def create(self) -> "Model": | |
| use_patch_token=False, | ||
| ) | ||
|
|
||
| if self.use_latent_perturbation: | ||
| num_members = self.ens_latent_perturb.get("num_members", 1) | ||
| if num_members > 1: | ||
| # loss_fcts keys that are able to exploit multiple ensemble members | ||
| ensemble_aware_loss_fcts = {"kernel_crps"} # extend as needed | ||
|
|
||
| configured_loss_fcts = { | ||
| fct_name | ||
| for loss_cfg in cf.training_config.losses.values() | ||
| if loss_cfg.get("enabled", True) | ||
| for fct_name in loss_cfg.get("loss_fcts", {}).keys() | ||
| } | ||
|
|
||
| if not (configured_loss_fcts & ensemble_aware_loss_fcts): | ||
| lf = sorted(configured_loss_fcts) | ||
| logger.warning( | ||
| f"decoder_ens_latent_perturbation.num_members={num_members} (>1) is " | ||
| f"configured, but none of the enabled loss_fcts {lf} are ensemble-aware" | ||
| ) | ||
|
|
||
| return self | ||
|
|
||
| def reset_parameters(self): | ||
|
|
@@ -589,6 +626,11 @@ def _reset_params(module): | |
|
|
||
| self.apply(_reset_params) | ||
|
|
||
| if self.latent_perturbation_log_sigma is not None: | ||
| sigma_init = self.ens_latent_perturb.get("sigma_init", 0.01) | ||
| with torch.no_grad(): | ||
| self.latent_perturbation_log_sigma.fill_(math.log(sigma_init)) | ||
|
|
||
| def print_num_parameters(self) -> None: | ||
| """Print number of parameters for entire model and each module used to build the model""" | ||
|
|
||
|
|
@@ -731,6 +773,81 @@ def predict_latent( | |
|
|
||
| return output | ||
|
|
||
| def _gather_neighbor_tokens( | ||
| self, | ||
| model_params: ModelParams, | ||
| tokens_stacked: torch.Tensor, | ||
| ) -> tuple[torch.Tensor, torch.Tensor]: | ||
| """ | ||
| Gather tokens from 1-ring neighborhood of cells (including cell itself) for each | ||
| cell in tokens_stacked | ||
|
|
||
| Output shape is (model_params.hp_nbours x num_hp_cells) x token_dim with | ||
| model_params.hp_nbours = 9 | ||
| """ | ||
|
|
||
| n_hpc = self.num_healpix_cells | ||
| num_stacked, num_tokens, token_dim = tokens_stacked.shape | ||
| num_neighbors = model_params.hp_nbours.shape[-1] | ||
| assert num_tokens == n_hpc | ||
|
|
||
| tokens_flat = tokens_stacked.reshape(num_stacked, n_hpc, 1, token_dim).flatten(0, 1) | ||
|
|
||
| cell_offsets = ( | ||
| torch.arange(num_stacked, device=tokens_stacked.device).view(num_stacked, 1, 1) * n_hpc | ||
| ) | ||
| # indices of neighbors | ||
| idxs = (model_params.hp_nbours.unsqueeze(0) + cell_offsets).flatten(0, 1) | ||
| # collect neighbors for each cell + cell itself | ||
| tokens_nbors = tokens_flat[idxs.flatten()].flatten(0, 1) | ||
|
|
||
| tokens_nbors_lens = tokens_flat.new_zeros(num_stacked * n_hpc + 1, dtype=torch.int32) | ||
| tokens_nbors_lens[1:] = num_neighbors | ||
|
|
||
| return tokens_nbors, tokens_nbors_lens | ||
|
|
||
| def _build_latent_tokens( | ||
| self, | ||
| tokens: torch.Tensor, | ||
| batch: ModelBatch, | ||
| ) -> tuple[torch.Tensor, int]: | ||
| """ | ||
| Strip aux tokens and, if latent perturbation is enabled, build the | ||
| CRPS ensemble by perturbing the assimilated latent tokens with Gaussian | ||
| noise (one perturbed copy per member). | ||
|
|
||
| Ensemble dimension is stacked along batch dimension, i.e. member m, | ||
| batch item b is at index m * batch_size + b along the first dimension of tokens_tiled | ||
| """ | ||
| # remove register and class tokens | ||
| tokens = tokens[:, self.num_aux_tokens :] | ||
|
|
||
| if not self.use_latent_perturbation: | ||
| return tokens, 1 | ||
|
|
||
| batch_size = len(batch) | ||
| token_dim = tokens.shape[-1] | ||
| assert tokens.shape == (batch_size, self.num_healpix_cells, token_dim), ( | ||
| f"unexpected token shape {tokens.shape}" | ||
| ) | ||
|
|
||
| num_members = self.ens_latent_perturb.get("num_members", 1) | ||
| sigma = torch.exp(self.latent_perturbation_log_sigma).to(tokens.dtype) | ||
| # ensemble perturbations | ||
| eps = torch.randn( | ||
| num_members, | ||
| batch_size, | ||
| self.num_healpix_cells * self.cf.ae_local_num_queries, | ||
| token_dim, | ||
| device=tokens.device, | ||
| dtype=tokens.dtype, | ||
| ) | ||
| # apply ensemble perturbations and flatten along ens dimension | ||
| out_shape = (num_members * batch_size, self.num_healpix_cells, token_dim) | ||
| tokens_tiled = (tokens.unsqueeze(0) + sigma * eps).reshape(out_shape) | ||
|
|
||
| return tokens_tiled, num_members | ||
|
|
||
| def predict_decoders( | ||
| self, | ||
| model_params: ModelParams, | ||
|
|
@@ -745,94 +862,119 @@ def predict_decoders( | |
| Predict outputs at the specific target coordinates based on the input weather state and | ||
| pre-training task and projects the latent space representation back to physical space. | ||
|
|
||
| If `latent_perturbation_num_members` > 1 (and `latent_perturbation_log_sigma` is set), | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. This latent_perturbation_num_members no longer exists. It has been replaced by decoder_ens_latent_perturbation.num_members |
||
| this generates a CRPS ensemble by perturbing the assimilated latent tokens with | ||
| Gaussian noise and running the decoder once per member. | ||
|
|
||
| Args: | ||
| model_params : Query and embedding parameters | ||
| fstep : Number of forecast steps | ||
| step : Forecast step | ||
| tokens : Tokens from global assimilation engine | ||
| streams_data : Used to initialize target coordinates tokens and index information | ||
| List of StreamData len(streams_data) == batch_size_per_gpu | ||
| target_coords_idxs : Indices of target coordinates | ||
| batch : Used to initialize target coordinates tokens and index information | ||
| output : Accumulator for predictions | ||
| Returns: | ||
| Prediction output tokens in physical representation for each target_coords. | ||
| """ | ||
| # Empty dicts evaluate to False in python | ||
|
|
||
| # skip if no physical predictions | ||
| if not self.pred_heads: | ||
| return output | ||
|
|
||
| # remove register and class tokens | ||
| tokens = tokens[:, self.num_aux_tokens :] | ||
|
|
||
| # get 1-ring neighborhood for prediction | ||
| # strip aux tokens and, if latent perturbation is enabled, build the ensemble | ||
| tokens_tiled, num_members = self._build_latent_tokens(tokens, batch) | ||
| batch_size = len(batch) | ||
| s = [batch_size, self.num_healpix_cells, self.cf.ae_local_num_queries, tokens.shape[-1]] | ||
| idxs = model_params.hp_nbours.unsqueeze(0).repeat((batch_size, 1, 1)).flatten(0, 1) | ||
| tokens_nbors = tokens.reshape(s).flatten(0, 1)[idxs.flatten()].flatten(0, 1) | ||
| # TODO: precompute in model_params? | ||
| tokens_nbors_lens = torch.full( | ||
| (s[0] * s[1] + 1,), fill_value=9, dtype=torch.int32, device=tokens_nbors.device | ||
| ) | ||
| tokens_nbors_lens[0] = 0 | ||
| token_dim = tokens.shape[-1] | ||
|
|
||
| # pair with tokens from assimilation engine to obtain target tokens | ||
| for stream_name in self.streams.keys(): | ||
| # extract target coords for current stream and fstep and convert to one tensor | ||
| t_coords = [ | ||
| batch.samples[i_b].streams_data[stream_name].target_coords[step] | ||
| for i_b in range(batch_size) | ||
| ] | ||
| t_coords_lens = [len(t) for t in t_coords] | ||
| t_coords = torch.cat(t_coords) | ||
|
|
||
| if len(t_coords) == 0: | ||
| if t_coords.shape[0] == 0: | ||
| continue | ||
|
|
||
| # embed token coords | ||
| tc_embed = self.embed_target_coords[stream_name] | ||
| tc_tokens = checkpoint(tc_embed, t_coords, use_reentrant=False) | ||
|
|
||
| # skip when coordinate embeddings yields nan (i.e. the coord embedding network diverged) | ||
| if torch.isnan(tc_tokens).any(): | ||
| logger.warning( | ||
| ( | ||
| f"Skipping prediction for {stream_name} because", | ||
| f" of {torch.isnan(tc_tokens).sum()} NaN in tc_tokens.", | ||
| ) | ||
| f"Skipping prediction for {stream_name} because" | ||
| f" of {torch.isnan(tc_tokens).sum()} NaN in tc_tokens." | ||
| ) | ||
| pred = torch.tensor([], device=tc_tokens.device) | ||
|
|
||
| # skip empty lengths | ||
| elif tc_tokens.shape[0] == 0: | ||
| pred = torch.tensor([], device=tc_tokens.device) | ||
|
|
||
| else: | ||
| # lens for varlen attention | ||
| tcls = torch.cat( | ||
| [ | ||
| sample.streams_data[stream_name].target_coords_lens[step] | ||
| for sample in batch.samples | ||
| ] | ||
| ) | ||
| tcs_lens = torch.cat([torch.zeros(1, dtype=torch.int32, device=tcls.device), tcls]) | ||
| num_groups = tcs_lens.shape[0] - 1 | ||
| assert batch_size * self.num_healpix_cells == num_groups, ( | ||
| f"expected {batch_size * self.num_healpix_cells} query groups, got {num_groups}" | ||
| ) | ||
|
|
||
| if self.cf.decoder_type == "Linear": | ||
| # repeat target-coord tokens once per ensemble member, | ||
| # matching tokens_tiled's ordering | ||
| tc_tokens_in = tc_tokens.repeat(num_members, 1) | ||
| tcs_lens_in = torch.cat( | ||
| [ | ||
| torch.zeros(1, dtype=torch.int32, device=tcs_lens.device), | ||
| tcs_lens[1:].repeat(num_members), | ||
| ] | ||
| ) | ||
| pred = self.target_token_engines[stream_name]( | ||
| tc_tokens, | ||
| tokens.reshape(-1, s[-1]), # collapse the batch and token dimensions | ||
| tcs_lens, | ||
| ).unsqueeze(0) # add ensemble dim: shape is then [1, preds_per_coord, channels] | ||
| tc_tokens_in, | ||
| tokens_tiled.reshape(-1, token_dim), | ||
| tcs_lens_in, | ||
| ) | ||
| pred = pred.reshape(num_members, t_coords.shape[0], pred.shape[-1]) | ||
|
|
||
| else: | ||
| tc_tokens = self.target_token_engines[stream_name]( | ||
| latent=tokens_nbors, | ||
| output=tc_tokens, | ||
| latent_lens=tokens_nbors_lens, | ||
| output_lens=tcs_lens, | ||
| coordinates=t_coords, | ||
| # Run members sequentially to stay under CUDA's max grid size | ||
| # for large ensembles or large hpl. Gather HEALPix neighbor | ||
| # tokens inside the loop so memory scales with one member | ||
| # instead of num_members. | ||
| tc_tokens_outs = [] | ||
| for m in range(num_members): | ||
| member_tokens = tokens_tiled[m * batch_size : (m + 1) * batch_size] | ||
| tokens_nbors, tokens_nbors_lens = self._gather_neighbor_tokens( | ||
| model_params, | ||
| member_tokens, | ||
| ) | ||
| tc_tokens_out_m = self.target_token_engines[stream_name]( | ||
| latent=tokens_nbors, | ||
| output=tc_tokens, | ||
| latent_lens=tokens_nbors_lens, | ||
| output_lens=tcs_lens, | ||
| coordinates=t_coords, | ||
| ) | ||
| tc_tokens_outs.append(tc_tokens_out_m) | ||
| tc_tokens_out = torch.cat(tc_tokens_outs, dim=0) | ||
|
|
||
| pred = self.pred_heads[stream_name](tc_tokens_out) | ||
| ens_head_size, total_target_points, output_channels = pred.shape | ||
| assert total_target_points == num_members * t_coords.shape[0], ( | ||
| f"expected {num_members * t_coords.shape[0]} pts in pred, " | ||
| f"got {total_target_points}" | ||
| ) | ||
| n_pts = t_coords.shape[0] | ||
| split_shape = (ens_head_size, num_members, n_pts, output_channels) | ||
| merged_shape = (num_members * ens_head_size, n_pts, output_channels) | ||
| pred = pred.reshape(split_shape).permute(1, 0, 2, 3).reshape(merged_shape) | ||
|
|
||
| # final prediction head to map back to physical space | ||
| pred = self.pred_heads[stream_name](tc_tokens) | ||
| assert pred.shape[1] == t_coords.shape[0], ( | ||
| f"expected {t_coords.shape[0]} points in dim=1, got {pred.shape[1]}" | ||
| ) | ||
|
|
||
| # recover batch dimension (ragged, so as list) | ||
| # recover batch dimension | ||
| pred = torch.split(pred, t_coords_lens, dim=1) | ||
| output.add_physical_prediction(step, stream_name, pred) | ||
|
|
||
|
|
||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Since this is nn.Parameter, is this also subject to weight decay on all parameters? Do we want this?