| """Peephole ConvLSTM encoder-forecaster for precipitation nowcasting.""" |
|
|
| import torch |
| from torch import nn |
| from torch.nn import functional as F |
|
|
|
|
| def patchify(sequence, patch_size): |
| batch, steps, channels, height, width = sequence.shape |
| flattened = sequence.flatten(0, 1) |
| patched = F.pixel_unshuffle(flattened, patch_size) |
| return patched.unflatten(0, (batch, steps)) |
|
|
|
|
| def unpatchify(sequence, patch_size): |
| batch, steps = sequence.shape[:2] |
| images = F.pixel_shuffle(sequence.flatten(0, 1), patch_size) |
| return images.unflatten(0, (batch, steps)) |
|
|
|
|
| class ConvLSTMCell(nn.Module): |
| def __init__(self, input_channels, hidden_channels, kernel_size): |
| super().__init__() |
| padding = kernel_size // 2 |
| self.hidden_channels = hidden_channels |
| self.input_conv = None if input_channels == 0 else nn.Conv2d( |
| input_channels, 4 * hidden_channels, kernel_size, padding=padding |
| ) |
| self.hidden_conv = nn.Conv2d(hidden_channels, 4 * hidden_channels, kernel_size, |
| padding=padding, bias=False) |
| self.peephole_input = nn.Parameter(torch.zeros(1, hidden_channels, 1, 1)) |
| self.peephole_forget = nn.Parameter(torch.zeros(1, hidden_channels, 1, 1)) |
| self.peephole_output = nn.Parameter(torch.zeros(1, hidden_channels, 1, 1)) |
| self.bias = nn.Parameter(torch.zeros(1, 4 * hidden_channels, 1, 1)) |
|
|
| def forward(self, values, state): |
| hidden, cell = state |
| gates = self.hidden_conv(hidden) + self.bias |
| if values is not None: |
| if self.input_conv is None: |
| raise ValueError("this ConvLSTM cell has no external input projection") |
| gates = gates + self.input_conv(values) |
| input_gate, forget_gate, candidate, output_gate = gates.chunk(4, dim=1) |
| input_gate = torch.sigmoid(input_gate + self.peephole_input * cell) |
| forget_gate = torch.sigmoid(forget_gate + self.peephole_forget * cell) |
| cell = forget_gate * cell + input_gate * torch.tanh(candidate) |
| output_gate = torch.sigmoid(output_gate + self.peephole_output * cell) |
| hidden = output_gate * torch.tanh(cell) |
| return hidden, cell |
|
|
|
|
| class ConvLSTM(nn.Module): |
| def __init__(self, config): |
| super().__init__() |
| self.patch_size = int(config["patch_size"]) |
| self.output_frames = int(config["output_frames"]) |
| patch_channels = int(config["input_channels"]) * self.patch_size ** 2 |
| hidden = [int(value) for value in config["hidden_channels"]] |
| kernel = int(config["kernel_size"]) |
| self.encoder = nn.ModuleList([ |
| ConvLSTMCell(patch_channels, hidden[0], kernel), |
| ConvLSTMCell(hidden[0], hidden[1], kernel), |
| ]) |
| self.forecaster = nn.ModuleList([ |
| ConvLSTMCell(0, hidden[0], kernel), |
| ConvLSTMCell(hidden[0], hidden[1], kernel), |
| ]) |
| self.output = nn.Conv2d(sum(hidden), patch_channels, 1) |
|
|
| @staticmethod |
| def _zero_state(batch, channels, height, width, reference): |
| zeros = reference.new_zeros(batch, channels, height, width) |
| return zeros, zeros.clone() |
|
|
| def forward(self, sequence, return_states=False): |
| patched = patchify(sequence, self.patch_size) |
| batch, _, _, height, width = patched.shape |
| states = [self._zero_state(batch, cell.hidden_channels, height, width, sequence) |
| for cell in self.encoder] |
| for step in range(patched.shape[1]): |
| values = patched[:, step] |
| for index, cell in enumerate(self.encoder): |
| states[index] = cell(values, states[index]) |
| values = states[index][0] |
| forecast_states = [(hidden.clone(), cell.clone()) for hidden, cell in states] |
| predictions, traces = [], [] |
| for _ in range(self.output_frames): |
| forecast_states[0] = self.forecaster[0](None, forecast_states[0]) |
| forecast_states[1] = self.forecaster[1](forecast_states[0][0], forecast_states[1]) |
| hidden = torch.cat((forecast_states[0][0], forecast_states[1][0]), dim=1) |
| predictions.append(self.output(hidden)) |
| traces.append([state[0] for state in forecast_states]) |
| logits = torch.stack(predictions, dim=1) |
| images = unpatchify(logits.sigmoid(), self.patch_size) |
| return (images, logits, traces) if return_states else (images, logits) |
|
|