FireCubeNet / model /firecubenet.py
zhangrenchao's picture
Publish FireCubeNet engineering reproduction
ef2ae28 verified
Raw
History Blame Contribute Delete
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)