import math import torch import torch.nn as nn class DenseBlock(nn.Module): def __init__(self, nf, gc): super().__init__() self.c1 = nn.Conv2d(nf, gc, 3, 1, 1) self.c2 = nn.Conv2d(nf + gc, gc, 3, 1, 1) self.c3 = nn.Conv2d(nf + 2 * gc, gc, 3, 1, 1) self.c4 = nn.Conv2d(nf + 3 * gc, gc, 3, 1, 1) self.c5 = nn.Conv2d(nf + 4 * gc, nf, 3, 1, 1) self.act = nn.LeakyReLU(0.2, inplace=True) def forward(self, x): x1 = self.act(self.c1(x)) x2 = self.act(self.c2(torch.cat([x, x1], 1))) x3 = self.act(self.c3(torch.cat([x, x1, x2], 1))) x4 = self.act(self.c4(torch.cat([x, x1, x2, x3], 1))) x5 = self.c5(torch.cat([x, x1, x2, x3, x4], 1)) return x + 0.2 * x5 class RRDB(nn.Module): def __init__(self, nf, gc): super().__init__() self.d1 = DenseBlock(nf, gc) self.d2 = DenseBlock(nf, gc) self.d3 = DenseBlock(nf, gc) def forward(self, x): out = self.d1(x) out = self.d2(out) out = self.d3(out) return x + 0.2 * out class DenoiseSRNet(nn.Module): """Noisy LR in -> clean HR out, 4x upsample via pixel-shuffle.""" def __init__(self, nf=64, gc=32, n_blocks=8, scale=4): super().__init__() self.scale = scale self.head = nn.Conv2d(1, nf, 3, 1, 1) self.body = nn.Sequential(*[RRDB(nf, gc) for _ in range(n_blocks)]) self.body_conv = nn.Conv2d(nf, nf, 3, 1, 1) up_layers = [] n_up = int(math.log2(scale)) for _ in range(n_up): up_layers += [ nn.Conv2d(nf, nf * 4, 3, 1, 1), nn.PixelShuffle(2), nn.LeakyReLU(0.2, inplace=True) ] self.upsample = nn.Sequential(*up_layers) self.tail = nn.Sequential( nn.Conv2d(nf, nf, 3, 1, 1), nn.LeakyReLU(0.2, inplace=True), nn.Conv2d(nf, 1, 3, 1, 1) ) def forward(self, x): feat = self.head(x) body_out = self.body_conv(self.body(feat)) feat = feat + body_out feat = self.upsample(feat) out = self.tail(feat) return torch.sigmoid(out)