Download model/climatebench.py from OneScience-Group/ClimateBench: direct link, hf CLI and curl.
- Browser
- Download file 3.94 kB
-
https://huggingface.co/OneScience-Group/ClimateBench/resolve/main/model/climatebench.py
- Command line
-
hf download hf://OneScience-Group/ClimateBench/model/climatebench.py
-
curl -L -o climatebench.py https://huggingface.co/OneScience-Group/ClimateBench/resolve/main/model/climatebench.py
3.94 kB
| """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()} | |