Download model/convlstm.py from OneScience-Group/ConvLSTM: direct link, hf CLI and curl.
- Browser
- Download file 4.45 kB
-
https://huggingface.co/OneScience-Group/ConvLSTM/resolve/main/model/convlstm.py
- Command line
-
hf download hf://OneScience-Group/ConvLSTM/model/convlstm.py
-
curl -L -o convlstm.py https://huggingface.co/OneScience-Group/ConvLSTM/resolve/main/model/convlstm.py
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) | |
| 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) | |