Capricorn / Baselines /Models /HiCSR.py
chnftq's picture
Upload folder using huggingface_hub
e3814d7 verified
Raw
History Blame Contribute Delete
5.53 kB
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
# --------------------------------------------------------
# References:
# HiCSR: https://github.com/PSI-Lab/HiCSR
# --------------------------------------------------------
import torch
from torch import nn
import torch.nn.functional as F
import math
class ResidualBlock(nn.Module):
def __init__(self, channels):
super(ResidualBlock, self).__init__()
self.conv1 = nn.Conv2d(channels, channels, kernel_size=3, padding=1)
self.bn1 = nn.BatchNorm2d(channels)
self.relu = nn.ReLU()
self.conv2 = nn.Conv2d(channels, channels, kernel_size=3, padding=1)
self.bn2 = nn.BatchNorm2d(channels)
def forward(self, x):
res = self.conv1(x)
res = self.bn1(res)
res = self.relu(res)
res = self.conv2(res)
res = self.bn2(res)
return x + res
class Generator(nn.Module):
def __init__(self, num_res_blocks=15, input_channels=1, out_channels=None):
super(Generator, self).__init__()
if not out_channels:
out_channels = input_channels
self.pre_res_block = nn.Sequential(
nn.Conv2d(input_channels, 64, kernel_size=3),
nn.ReLU(),
)
res_blocks = [ResidualBlock(64) for _ in range(num_res_blocks)]
self.res_blocks = nn.Sequential(*res_blocks)
self.post_res_block = nn.Sequential(
nn.Conv2d(64, 64, kernel_size=3, padding=1),
nn.BatchNorm2d(64)
)
self.final_block = nn.Sequential(
nn.Conv2d(64, 128, kernel_size=3),
nn.Conv2d(128, 128, kernel_size=3),
nn.Conv2d(128, 256, kernel_size=3),
nn.Conv2d(256, 256, kernel_size=3),
nn.Conv2d(256, out_channels, kernel_size=3),
)
def forward(self, x):
first_block = self.pre_res_block(x)
res_blocks = self.res_blocks(first_block)
post_res_block = self.post_res_block(res_blocks)
final_block = self.final_block(first_block + post_res_block)
return torch.tanh(final_block)
class Discriminator(nn.Module):
def __init__(self, input_channels=1):
super(Discriminator, self).__init__()
self.conv = nn.Sequential(
nn.Conv2d(input_channels, 64, kernel_size=4, stride=2, padding=1, bias=False),
nn.LeakyReLU(0.2, inplace=True),
nn.Conv2d(64, 128, kernel_size=4, stride=2, padding=1, bias=False),
nn.BatchNorm2d(128),
nn.LeakyReLU(0.2, inplace=True),
nn.Conv2d(128, 256, kernel_size=4, stride=2, padding=1, bias=False),
nn.BatchNorm2d(256),
nn.LeakyReLU(0.2, inplace=True),
nn.Conv2d(256, 512, kernel_size=4, stride=2, padding=1, bias=False),
nn.BatchNorm2d(512),
nn.LeakyReLU(0.2, inplace=True),
nn.Conv2d(512, input_channels, kernel_size=1, stride=1, padding=0, bias=False),
nn.Sigmoid())# temporarily add the sigmoid
self.init_params()
def forward(self, x):
x = self.conv(x)
return x
def init_params(self):
for module in self.modules():
if isinstance(module, nn.Conv2d):
nn.init.normal_(module.weight.data, 0.0, 0.02)
elif isinstance(module, nn.BatchNorm2d):
nn.init.normal_(module.weight.data, 1.0, 0.02)
nn.init.constant_(module.bias.data, 0)
class DAE(nn.Module):
def __init__(self, num_layers=5, num_features=64, input_channels=1):
super(DAE, self).__init__()
self.num_layers = num_layers
conv_layers = []
deconv_layers = []
conv_layers.append(nn.Sequential(nn.Conv2d(input_channels, num_features, kernel_size=3, stride=2, padding=1),
nn.ReLU(inplace=True)))
for i in range(num_layers - 1):
conv_layers.append(nn.Sequential(nn.Conv2d(num_features, num_features, kernel_size=3, padding=1),
nn.ReLU(inplace=True)))
for i in range(num_layers - 1):
deconv_layers.append(nn.Sequential(nn.ConvTranspose2d(num_features, num_features, kernel_size=3, padding=1),
nn.ReLU(inplace=True)))
deconv_layers.append(nn.ConvTranspose2d(num_features, input_channels, kernel_size=3, stride=2, padding=1, output_padding=1))
self.conv_layers = nn.Sequential(*conv_layers)
self.deconv_layers = nn.Sequential(*deconv_layers)
self.relu = nn.ReLU(inplace=True)
def forward(self, x):
residual = x
conv_feats = []
for i in range(self.num_layers):
x = self.conv_layers[i](x)
if (i + 1) % 2 == 0 and len(conv_feats) < math.ceil(self.num_layers / 2) - 1:
conv_feats.append(x)
conv_feats_idx = 0
for i in range(self.num_layers):
x = self.deconv_layers[i](x)
if (i + 1 + self.num_layers) % 2 == 0 and conv_feats_idx < len(conv_feats):
conv_feat = conv_feats[-(conv_feats_idx + 1)]
conv_feats_idx += 1
x = x + conv_feat
x = self.relu(x)
x += residual
x = torch.tanh(x)
return x