File size: 2,471 Bytes
ef2ae28
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""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)