"""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)