File size: 4,449 Bytes
9be39c5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
"""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)