RainNet / model /rainnet.py
Niulh's picture
Publish RainNet model package
9f29df6 verified
Raw
History Blame Contribute Delete
3.64 kB
"""PyTorch implementation of RainNet v1.0."""
import torch
from torch import nn
class RainNet(nn.Module):
"""The 20-convolution encoder-decoder from Ayzel et al. (2020)."""
def __init__(self, in_channels=4, out_channels=1, dropout=0.5, mode="regression"):
super().__init__()
if mode != "regression":
raise ValueError("This package implements RainNet regression mode only")
self.in_channels = in_channels
self.pool = nn.MaxPool2d(2)
self.upsample = nn.Upsample(scale_factor=2, mode="nearest")
self.relu = nn.ReLU(inplace=True)
self.conv1f = nn.Conv2d(in_channels, 64, 3, padding=1)
self.conv1s = nn.Conv2d(64, 64, 3, padding=1)
self.conv2f = nn.Conv2d(64, 128, 3, padding=1)
self.conv2s = nn.Conv2d(128, 128, 3, padding=1)
self.conv3f = nn.Conv2d(128, 256, 3, padding=1)
self.conv3s = nn.Conv2d(256, 256, 3, padding=1)
self.conv4f = nn.Conv2d(256, 512, 3, padding=1)
self.conv4s = nn.Conv2d(512, 512, 3, padding=1)
self.drop4 = nn.Dropout(dropout)
self.conv5f = nn.Conv2d(512, 1024, 3, padding=1)
self.conv5s = nn.Conv2d(1024, 1024, 3, padding=1)
self.drop5 = nn.Dropout(dropout)
self.conv6f = nn.Conv2d(1536, 512, 3, padding=1)
self.conv6s = nn.Conv2d(512, 512, 3, padding=1)
self.conv7f = nn.Conv2d(768, 256, 3, padding=1)
self.conv7s = nn.Conv2d(256, 256, 3, padding=1)
self.conv8f = nn.Conv2d(384, 128, 3, padding=1)
self.conv8s = nn.Conv2d(128, 128, 3, padding=1)
self.conv9f = nn.Conv2d(192, 64, 3, padding=1)
self.conv9s = nn.Conv2d(64, 64, 3, padding=1)
self.conv9out = nn.Conv2d(64, 2, 3, padding=1)
self.output = nn.Conv2d(2, out_channels, 1)
self.apply(self._initialize)
@staticmethod
def _initialize(module):
if isinstance(module, nn.Conv2d):
nn.init.kaiming_normal_(module.weight, mode="fan_in", nonlinearity="relu")
if module.bias is not None:
nn.init.zeros_(module.bias)
def forward(self, x):
if x.ndim != 4:
raise ValueError(f"RainNet expects BCHW input, got shape {tuple(x.shape)}")
if x.shape[1] != self.in_channels:
raise ValueError(f"RainNet expects {self.in_channels} channels, got {x.shape[1]}")
if x.shape[-2] % 16 or x.shape[-1] % 16:
raise ValueError(
f"Spatial dimensions must be divisible by 16, got {tuple(x.shape[-2:])}"
)
c1 = self.relu(self.conv1s(self.relu(self.conv1f(x))))
c2 = self.relu(self.conv2s(self.relu(self.conv2f(self.pool(c1)))))
c3 = self.relu(self.conv3s(self.relu(self.conv3f(self.pool(c2)))))
c4 = self.relu(self.conv4s(self.relu(self.conv4f(self.pool(c3)))))
c5 = self.relu(self.conv5s(self.relu(self.conv5f(self.pool(self.drop4(c4))))))
c5 = self.drop5(c5)
c6 = torch.cat((self.upsample(c5), c4), dim=1)
c6 = self.relu(self.conv6s(self.relu(self.conv6f(c6))))
c7 = torch.cat((self.upsample(c6), c3), dim=1)
c7 = self.relu(self.conv7s(self.relu(self.conv7f(c7))))
c8 = torch.cat((self.upsample(c7), c2), dim=1)
c8 = self.relu(self.conv8s(self.relu(self.conv8f(c8))))
c9 = torch.cat((self.upsample(c8), c1), dim=1)
c9 = self.relu(self.conv9s(self.relu(self.conv9f(c9))))
return self.output(self.relu(self.conv9out(c9)))
def build_rainnet(in_channels=4, out_channels=1, dropout=0.5, mode="regression"):
return RainNet(in_channels, out_channels, dropout, mode)