class ResidualConvBlock(nn.Module):
def __init__(self, channels):
super().__init__()
self.conv1 = nn.Conv2d(channels, channels, kernel_size=3, padding=1)
self.conv2 = nn.Conv2d(channels, channels, kernel_size=3, padding=1)
self.relu = nn.ReLU()
nn.init.kaiming_normal_(self.conv1.weight, nonlinearity="relu")
nn.init.kaiming_normal_(self.conv2.weight, nonlinearity="relu")
def forward(self, x):
out = self.relu(self.conv1(x))
out = self.conv2(out)
return self.relu(x + out) # connexion résiduelle


