# coding=utf-8 # # SPDX-License-Identifier: MIT # # Copyright (c) 2021 Open Climate Fix # # This module is a thin configuration wrapper around the DGMR (Deep # Generative Model of Radar) architecture from Ravuri et al. (2021, # "Skilful Precipitation Nowcasting using Deep Generative Models of Radar", # Nature 597), as re-implemented in PyTorch by Open Climate Fix # (``openclimatefix/skillful_nowcasting``, MIT License). The network modules # are vendored verbatim (minus HuggingFace hub mixins) under # ``dgmr_official/``; only config plumbing and the plain ``forward`` are # added here for YAML-driven usage. import torch import torch.nn as nn from model.dgmr_official.common import ContextConditioningStack, LatentConditioningStack from model.dgmr_official.discriminators import Discriminator from model.dgmr_official.generators import Generator, Sampler from model.dgmr_official.losses import GridCellLoss def weight_fn(y, precip_weight_cap=24.0): """ Weight function for the grid cell loss: w(y) = max(y + 1, cap). """ return torch.max(y + 1, torch.tensor(precip_weight_cap, device=y.device)) class DGMR(nn.Module): """ Config-driven DGMR wrapper (generator + discriminator). The generator is a conditional GAN generator that takes ``num_context`` observed radar frames of shape [B, T, C, H, W] and produces ``forecast_steps`` future frames of the same spatial size. The discriminator scores full sequences (context + forecast) spatially and temporally; during GAN training the hinge losses plus the grid-cell regularizer are applied (see ``dgmr_official/losses.py`` and the paper). Args: forecast_steps: Number of frames to predict in the future (paper: 18). num_context: Number of input/context frames (paper: 4). input_channels: Number of channels per frame (paper: 1, radar). output_shape: Spatial size of the frames; must be divisible by 32 (paper: 256). Discriminators additionally need >= 128 px. conv_type: Convolution flavour used by the conditioning stack, one of "standard" / "coord" / "3d". latent_channels / context_channels: DGMR architecture sizes (paper: 768 / 384). generation_steps: Number of Monte-Carlo generator samples used when computing the grid-cell regularizer during training (paper: 6). grid_lambda: Weight of the grid-cell regularizer (paper: 20). precip_weight_cap: Ceiling for the grid-cell weight function (paper: 24). """ def __init__( self, forecast_steps: int = 18, num_context: int = 4, input_channels: int = 1, output_shape: int = 256, conv_type: str = "standard", latent_channels: int = 768, context_channels: int = 384, generation_steps: int = 6, grid_lambda: float = 20.0, precip_weight_cap: float = 24.0, ): super().__init__() self.forecast_steps = int(forecast_steps) self.num_context = int(num_context) self.input_channels = int(input_channels) self.output_shape = int(output_shape) self.conv_type = conv_type self.latent_channels = int(latent_channels) self.context_channels = int(context_channels) self.generation_steps = int(generation_steps) self.grid_lambda = float(grid_lambda) self.precip_weight_cap = float(precip_weight_cap) self.conditioning_stack = ContextConditioningStack( input_channels=self.input_channels, conv_type=self.conv_type, output_channels=self.context_channels, ) self.latent_stack = LatentConditioningStack( shape=( 8 * self.input_channels, self.output_shape // 32, self.output_shape // 32, ), output_channels=self.latent_channels, ) self.sampler = Sampler( forecast_steps=self.forecast_steps, latent_channels=self.latent_channels, context_channels=self.context_channels, ) self.generator = Generator(self.conditioning_stack, self.latent_stack, self.sampler) self.discriminator = Discriminator(self.input_channels) self.grid_regularizer = GridCellLoss( weight_fn=weight_fn, precip_weight_cap=self.precip_weight_cap ) def forward(self, x: torch.Tensor) -> torch.Tensor: """ Args: x: Observed radar frames, shape [batch, num_context, C, H, W]. Returns: Forecast frames, shape [batch, forecast_steps, C, H, W]. """ return self.generator(x)