sam0310's picture
Upload model.py
951020d verified
Raw History Blame Contribute Delete
2.2 kB
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)