File size: 4,699 Bytes
5a5d1a8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
# 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)