Download model/firecubenet.py from OneScience-Group/FireCubeNet: direct link, hf CLI and curl.
- Browser
- Download file 2.47 kB
-
https://huggingface.co/OneScience-Group/FireCubeNet/resolve/main/model/firecubenet.py
- Command line
-
hf download hf://OneScience-Group/FireCubeNet/model/firecubenet.py
-
curl -L -o firecubenet.py https://huggingface.co/OneScience-Group/FireCubeNet/resolve/main/model/firecubenet.py
2.47 kB
| """ConvLSTM model for center-pixel next-day wildfire danger.""" | |
| from __future__ import annotations | |
| import torch | |
| from torch import nn | |
| class ConvLSTMCell(nn.Module): | |
| """Standard ConvLSTM cell with input, forget, output, and candidate gates.""" | |
| def __init__(self, input_channels: int, hidden_channels: int, kernel_size: int = 3): | |
| super().__init__() | |
| self.hidden_channels = int(hidden_channels) | |
| padding = kernel_size // 2 | |
| self.gates = nn.Conv2d( | |
| input_channels + hidden_channels, 4 * hidden_channels, | |
| kernel_size=kernel_size, padding=padding, | |
| ) | |
| def forward(self, inputs: torch.Tensor, state: tuple[torch.Tensor, torch.Tensor]): | |
| hidden, cell = state | |
| input_gate, forget_gate, output_gate, candidate = self.gates( | |
| torch.cat((inputs, hidden), dim=1) | |
| ).chunk(4, dim=1) | |
| input_gate = torch.sigmoid(input_gate) | |
| forget_gate = torch.sigmoid(forget_gate) | |
| output_gate = torch.sigmoid(output_gate) | |
| candidate = torch.tanh(candidate) | |
| next_cell = forget_gate * cell + input_gate * candidate | |
| next_hidden = output_gate * torch.tanh(next_cell) | |
| return next_hidden, next_cell | |
| class FireCubeNet(nn.Module): | |
| """Propagate ConvLSTM state over ten days and classify the center pixel.""" | |
| def __init__(self, input_channels: int = 25, hidden_channels: int = 4, | |
| kernel_size: int = 3, dropout: float = 0.1): | |
| super().__init__() | |
| self.input_channels = int(input_channels) | |
| self.hidden_channels = int(hidden_channels) | |
| self.cell = ConvLSTMCell(input_channels, hidden_channels, kernel_size) | |
| self.head = nn.Sequential(nn.Dropout(dropout), nn.Linear(hidden_channels, 1)) | |
| def forward(self, inputs: torch.Tensor) -> torch.Tensor: | |
| if inputs.ndim != 5 or inputs.shape[2] != self.input_channels: | |
| raise ValueError( | |
| f"expected BTCHW with C={self.input_channels}, got {tuple(inputs.shape)}" | |
| ) | |
| batch, _, _, height, width = inputs.shape | |
| hidden = inputs.new_zeros(batch, self.hidden_channels, height, width) | |
| cell = inputs.new_zeros(batch, self.hidden_channels, height, width) | |
| for time_index in range(inputs.shape[1]): | |
| hidden, cell = self.cell(inputs[:, time_index], (hidden, cell)) | |
| center_features = hidden[:, :, height // 2, width // 2] | |
| return self.head(center_features) | |