ConvLSTM / model /convlstm.py
zhangrenchao's picture
Add engineering reproduction package
9be39c5 verified
Raw
History Blame Contribute Delete
4.45 kB
"""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)