File size: 3,941 Bytes
2e913c2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Paper-faithful PyTorch CNN-LSTM baseline for ClimateBench."""

from __future__ import annotations

import torch
from torch import nn
from torch.nn import functional as F


TARGETS = ("tas", "dtr", "pr", "pr90")


class ReLULSTM(nn.Module):
    """Keras-compatible LSTM using sigmoid gates and ReLU activation."""

    def __init__(self, input_size: int = 20, hidden_size: int = 25):
        super().__init__()
        self.input_size = int(input_size)
        self.hidden_size = int(hidden_size)
        self.kernel = nn.Parameter(torch.empty(input_size, 4 * hidden_size))
        self.recurrent_kernel = nn.Parameter(torch.empty(hidden_size, 4 * hidden_size))
        self.bias = nn.Parameter(torch.zeros(4 * hidden_size))
        self.reset_parameters()

    def reset_parameters(self) -> None:
        nn.init.xavier_uniform_(self.kernel)
        nn.init.orthogonal_(self.recurrent_kernel)
        nn.init.zeros_(self.bias)
        with torch.no_grad():
            self.bias[self.hidden_size:2 * self.hidden_size].fill_(1.0)

    def forward(self, sequence: torch.Tensor) -> torch.Tensor:
        if sequence.ndim != 3 or sequence.shape[-1] != self.input_size:
            raise ValueError(f"expected [B,T,{self.input_size}], got {tuple(sequence.shape)}")
        batch = sequence.shape[0]
        hidden = sequence.new_zeros(batch, self.hidden_size)
        cell = sequence.new_zeros(batch, self.hidden_size)
        for step in range(sequence.shape[1]):
            gates = sequence[:, step] @ self.kernel + hidden @ self.recurrent_kernel + self.bias
            input_gate, forget_gate, candidate, output_gate = gates.chunk(4, dim=-1)
            input_gate = torch.sigmoid(input_gate)
            forget_gate = torch.sigmoid(forget_gate)
            candidate = F.relu(candidate)
            output_gate = torch.sigmoid(output_gate)
            cell = forget_gate * cell + input_gate * candidate
            hidden = output_gate * F.relu(cell)
        return hidden


class ClimateBenchBranch(nn.Module):
    """One 364,764-parameter paper CNN-LSTM target emulator."""

    PAPER_PARAMETER_COUNT = 364_764

    def __init__(self, height: int = 96, width: int = 144):
        super().__init__()
        self.height = int(height)
        self.width = int(width)
        self.conv = nn.Conv2d(4, 20, kernel_size=3, padding="same")
        self.pool = nn.AvgPool2d(kernel_size=2)
        self.lstm = ReLULSTM(20, 25)
        self.dense = nn.Linear(25, height * width)
        count = sum(parameter.numel() for parameter in self.parameters())
        if count != self.PAPER_PARAMETER_COUNT:
            raise RuntimeError(f"paper branch must have 364764 parameters, got {count}")

    def forward(self, inputs: torch.Tensor) -> torch.Tensor:
        if inputs.ndim != 5 or inputs.shape[2:] != (4, self.height, self.width):
            raise ValueError(f"expected [B,T,4,{self.height},{self.width}], got {tuple(inputs.shape)}")
        batch, time = inputs.shape[:2]
        features = F.relu(self.conv(inputs.reshape(batch * time, 4, self.height, self.width)))
        features = self.pool(features).mean(dim=(-2, -1)).reshape(batch, time, 20)
        return self.dense(self.lstm(features)).reshape(batch, 1, self.height, self.width)


class ClimateBench(nn.Module):
    """Four independent paper branches ordered as tas, dtr, pr and pr90."""

    def __init__(self, height: int = 96, width: int = 144):
        super().__init__()
        self.height = int(height)
        self.width = int(width)
        self.branches = nn.ModuleDict({name: ClimateBenchBranch(height, width) for name in TARGETS})

    def forward(self, inputs: torch.Tensor) -> torch.Tensor:
        return torch.cat([self.branches[name](inputs) for name in TARGETS], dim=1)

    def parameter_counts(self) -> dict[str, int]:
        return {name: sum(parameter.numel() for parameter in branch.parameters())
                for name, branch in self.branches.items()}