diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..69ca6cc --- /dev/null +++ b/.gitignore @@ -0,0 +1,135 @@ +# Byte-compiled / optimized / DLL files +__pycache__/ +*.py[cod] +*$py.class + +# C extensions +*.so + +# Distribution / packaging +.Python +build/ +develop-eggs/ +dist/ +downloads/ +eggs/ +.eggs/ +lib/ +lib64/ +parts/ +sdist/ +var/ +wheels/ +pip-wheel-metadata/ +share/python-wheels/ +*.egg-info/ +.installed.cfg +*.egg +MANIFEST + +# PyInstaller +# Usually these files are written by a python script from a template +# before PyInstaller builds the exe, so as to inject date/other infos into it. +*.manifest +*.spec + +# Installer logs +pip-log.txt +pip-delete-this-directory.txt + +# Unit test / coverage reports +htmlcov/ +.tox/ +.nox/ +.coverage +.coverage.* +.cache +nosetests.xml +coverage.xml +*.cover +*.py,cover +.hypothesis/ +.pytest_cache/ + +# Translations +*.mo +*.pot + +# Django stuff: +*.log +local_settings.py +db.sqlite3 +db.sqlite3-journal + +# Flask stuff: +instance/ +.webassets-cache + +# Scrapy stuff: +.scrapy + +# Sphinx documentation +docs/_build/ + +# PyBuilder +target/ + +# Jupyter Notebook +.ipynb_checkpoints + +# IPython +profile_default/ +ipython_config.py + +# pyenv +.python-version + +# pipenv +# According to pypa/pipenv#598, it is recommended to include Pipfile.lock in version control. +# However, in case of collaboration, if having platform-specific dependencies or dependencies +# having no cross-platform support, pipenv may install dependencies that don't work, or not +# install all needed dependencies. +#Pipfile.lock + +# PEP 582; used by e.g. github.com/David-OConnor/pyflow +__pypackages__/ + +# Celery stuff +celerybeat-schedule +celerybeat.pid + +# SageMath parsed files +*.sage.py + +# Environments +.env +.venv +env/ +venv/ +ENV/ +env.bak/ +venv.bak/ + +# Spyder project settings +.spyderproject +.spyproject + +# Rope project settings +.ropeproject + +# mkdocs documentation +/site + +# mypy +.mypy_cache/ +.dmypy.json +dmypy.json + +# Pyre type checker +.pyre/ + +# custom +/data +src/samples/ +src/models +.vscode \ No newline at end of file diff --git a/Makefile b/Makefile new file mode 100644 index 0000000..0cff594 --- /dev/null +++ b/Makefile @@ -0,0 +1,22 @@ +SHELL := /bin/bash +CONDA_ACTIVATE=source $$(conda info --base)/etc/profile.d/conda.sh ; conda activate ; conda activate + +.PHONY: help lint local build pull run + +.DEFAULT: help + +help: + @echo "make run" + @echo " run python main.py" + +ENV_NAME := super-res +visdom: + @cd src;\ + $(CONDA_ACTIVATE) $(ENV_NAME);\ + python -m visdom.server;\ + +run: + @cd src;\ + $(CONDA_ACTIVATE) $(ENV_NAME);\ + CUDA_VISIBLE_DEVICES=2,3 nohup python main.py > main.log & + cd ..; diff --git a/README.md b/README.md new file mode 100644 index 0000000..a7bfce2 --- /dev/null +++ b/README.md @@ -0,0 +1,53 @@ +# Project 7 +- Topic - Image Super Resolution +- Paper - Learning a Single Convolutional Super-Resolution Network for +Multiple Degradations +- Paper Link - https://arxiv.org/pdf/1712.06116v2.pdf + + +# Run Instructions +1. `make run` + +# Usage +## Requirements +- `python=3.8` +- `pytorch=1.7.1` +- `cudatoolkit=10.2` + +## Code structure +``` +. +├── docs +| ├── TeamKota-MidEvalReport.pdf +| ├── Project Proposal +│ ├── Method.pdf +│ └── TeamKota-ProjectProposal.pdf +├── README.md +├── Makefile +├── data + ├── BSDS300 + ├── DIV2K_train_HR + ├── waterloo +├── environment.yml +└── src + ├── dataset.py + ├── kernels.py + ├── logger.py + ├── main.py + ├── model.py + ├── train.py + └── utils.py +``` +## Setup env +- `conda env create --file environment.yaml` + +## Data + +- use `Name` in main.py + +| Id | Name |Info | Link | +| --- | --- | ----------- | ----------- | +| 1 | DIV2K_train_HR | HR images | http://data.vision.ee.ethz.ch/cvl/DIV2K/DIV2K_train_HR.zip | +| 2 | DIV2K_train_LR_bicubic_X2 | Bicubic x2 downscaling | http://data.vision.ee.ethz.ch/cvl/DIV2K/DIV2K_train_LR_bicubic_X2.zip | +| 3 | Waterloo Exploration Database | Pristine Natural Images | http://ivc.uwaterloo.ca/database/WaterlooExploration/exploration_database_and_code.rar | +| 4 | Berkeley Segementation Dataset and Benchmark (BSDS300) | Natural Images with greyscale and color segementation | https://www2.eecs.berkeley.edu/Research/Projects/CS/vision/bsds/BSDS300-images.tgz | \ No newline at end of file diff --git a/docs/Project-Proposal/Method.pdf b/docs/Project-Proposal/Method.pdf new file mode 100644 index 0000000..3949ce8 Binary files /dev/null and b/docs/Project-Proposal/Method.pdf differ diff --git a/docs/Project-Proposal/TeamKota-ProjectProposal.pdf b/docs/Project-Proposal/TeamKota-ProjectProposal.pdf new file mode 100644 index 0000000..66fdf85 Binary files /dev/null and b/docs/Project-Proposal/TeamKota-ProjectProposal.pdf differ diff --git a/docs/Team Kota - Final Eval Report.pdf b/docs/Team Kota - Final Eval Report.pdf new file mode 100644 index 0000000..35badea Binary files /dev/null and b/docs/Team Kota - Final Eval Report.pdf differ diff --git a/docs/TeamKota-MidEvalReport.pdf b/docs/TeamKota-MidEvalReport.pdf new file mode 100644 index 0000000..6b200f8 Binary files /dev/null and b/docs/TeamKota-MidEvalReport.pdf differ diff --git a/environment.yml b/environment.yml new file mode 100644 index 0000000..ed89f3d --- /dev/null +++ b/environment.yml @@ -0,0 +1,13 @@ +name: super-res +channels: + - defaults +dependencies: + - python=3.8 + - cudatoolkit=10.2 + - torchaudio + - pytorch + - torchvision + - scipy + - visdom + - scikit-learn + - jsonpatch diff --git a/src/dataset.py b/src/dataset.py new file mode 100644 index 0000000..e92ea4d --- /dev/null +++ b/src/dataset.py @@ -0,0 +1,84 @@ +import torch +from torch.utils import data +from PIL import Image +import numpy as np +from kernels import Kernels +from torchvision import transforms +import math +import random +from torchvision.utils import save_image +import glob + +def Scaling(image): + return np.array(image) / 255.0 + +class AddGaussianNoise(object): + def __init__(self, mean=0.): + self.mean = mean + + def __call__(self, vec): + vec = np.asarray(vec) + h, w, n = vec.shape + self.std = float(random.randint(0, 76)) + return vec + np.random.rand(h,w,n) * self.std + self.mean + + def __repr__(self): + return self.__class__.__name__ + '(mean={0}, std={1})'.format(self.mean, self.std) + +class DIV2K_train(data.Dataset): + def __init__(self, config=None): + + self.image_paths = [] + num_images = 800 + for i in range(1, num_images+1): + name = '0000' + str(i) + name = name[-4:] + Y_path = config.y_path + name + '.png' + self.image_paths.append(Y_path) + self.image_paths += glob.glob(config.y_path2) + self.image_paths += glob.glob(config.y_path3) + self.scale_factor = config.scale_factor + self.image_size = config.image_size + + self.kernels = Kernels(self.scale_factor) + + + def __getitem__(self, index): + Y_path = self.image_paths[index] + + Y_image = Image.open(Y_path).convert('RGB') # hr image + # Y_image.save("ogyimage"+str(index)+".jpg") + X_imageact,X_image, Y_image = self.transformlr(Y_image,index) + + return X_imageact.to(torch.float64), X_image.to(torch.float64), Y_image.to(torch.float64) + + def __len__(self): + return len(self.image_paths) + + def transformlr(self, Y_image,index): + transform = transforms.RandomCrop(self.image_size * self.scale_factor) + hr_image = transform(Y_image) #image + # print(hr_image) + # hr_image.save("transyimage"+str(index)+".jpg") + + kernel, degradinfo = random.choice(self.kernels.allkernels) + + transform = transforms.Compose([ + transforms.Lambda(lambda x: self.kernels.Blur(x,kernel)), + transforms.Resize((self.image_size, self.image_size), interpolation=3), + AddGaussianNoise() + ]) + + lr_image = np.asarray(transform(hr_image)) #numpy + transform = transforms.ToTensor() + lr_imageact= transform(lr_image) + # print("here",lr_image.shape) + # temp = Image.fromarray(lr_image.astype(np.uint8)) + # temp.save("transximage"+str(index)+".jpg") + + transform = transforms.Compose([transforms.Lambda(lambda x: self.kernels.ConcatDegraInfo(x,degradinfo))]) + lr_image = transform(lr_image) + + transform = transforms.ToTensor() + lr_image, hr_image = transform(lr_image), transform(hr_image) + return lr_imageact,lr_image, hr_image diff --git a/src/kernels.py b/src/kernels.py new file mode 100644 index 0000000..b17b4cb --- /dev/null +++ b/src/kernels.py @@ -0,0 +1,118 @@ +import cv2 +import math +import torch +import random +import numpy as np +from PIL import Image +from scipy import signal +from scipy.ndimage import convolve +from sklearn.decomposition import PCA +from scipy.stats import multivariate_normal +import random + + +class Kernels(object): + def __init__(self, scaleFactor): + # big-d: the class has other values initialsed, do we not need them? + self.allkernels = [] + self.scaleFactor = scaleFactor + + self.allkernels = np.zeros((10000,15,15)) + + for count in range(10000): + + theta = random.random() * np.pi + l1 = 0.5 + random.random() * ((self.scaleFactor * 2 ) + 1.5) + l2 = 0.5 + random.random() * (l1-0.5) + + ker = self.getKernel(theta,l1,l2) + self.allkernels[count,:,:] = ker + + self.degradation = self.PCA() + temp = [] + for index in range(len(self.allkernels)): + temp.append([self.allkernels[index,:,:],self.degradation[index]]) + + self.allkernels = temp + + # def __init__(self, scaleFactor): + # # big-d: the class has other values initialsed, do we not need them? + # self.allkernels = [] + # self.scaleFactor = scaleFactor + # + # # sai: Add anisotropic kernels + # widths = [x/10 for x in range(2, 10*self.scaleFactor + 1)] + # + # self.allkernels = np.zeros((len(widths),15,15)) + # for index, width in enumerate(widths): + # # yeet=random.randint(0, 1) + # # if yeet==0: + # ker = self.isogkern(15,width) + # # else: + # # ker = self.anisogkern(15,width,random.uniform(0.2,width)) + # self.allkernels[index,:,:] = ker + # + # self.degradation = self.PCA() + # temp = [] + # for index in range(len(self.allkernels)): + # temp.append([self.allkernels[index,:,:],self.degradation[index]]) + # + # self.allkernels = temp + + + def Blur(self, image, kernel): + image = np.asarray(image) + dst = cv2.filter2D(image,-1,kernel) + return Image.fromarray(dst) + + def ConcatDegraInfo(self, image, degradation): + h, w = list(image.shape[0:2]) + n = 15 # taking n=15 PCA components + maps = np.ones((h, w, n)) + for i in range(15): + maps[:, :, i] = degradation[i] * maps[:, :, i] + image = np.concatenate((image, maps), axis=-1) + return image + + def PCA(self, k=15): + + data = self.allkernels.reshape(-1,225) + pca = PCA(n_components=k) + new_data = pca.fit_transform(data) + return torch.from_numpy(new_data) + + + def isogkern(self, kernlen, std): + gkern1d = signal.gaussian(kernlen, std=std).reshape(kernlen, 1) + gkern2d = np.outer(gkern1d, gkern1d) + gkern2d = gkern2d/np.sum(gkern2d) + return gkern2d + + + def anisogkern(self, kernlen, std1, std2): + # big-d: angle NOT used + gkern1d_1 = signal.gaussian(kernlen, std=std1).reshape(kernlen, 1) + gkern1d_2 = signal.gaussian(kernlen, std=std2).reshape(kernlen, 1) + gkern2d = np.outer(gkern1d_1, gkern1d_2) + gkern2d = gkern2d/np.sum(gkern2d) + return gkern2d + + + def getKernel(self,theta,l1,l2): + + v = np.dot([[math.cos(theta), -math.sin(theta)],[math.sin(theta), math.cos(theta)]],[[1],[0]]) + + V = [[v[0][0], v[1][0]],[v[1][0], -v[0][0]]] + D = [[l1, 0],[0, l2]] + + Sigma = np.dot(np.dot(V,D),np.linalg.inv(V)) + rv = multivariate_normal([7,7], Sigma) + + ker = np.zeros((15,15)) + + for i in range(15): + for j in range(15): + ker[i][j] = rv.pdf([i,j]) + + ker /= np.sum(ker) + return ker diff --git a/src/kernels/SRMDNFx2.mat b/src/kernels/SRMDNFx2.mat new file mode 100644 index 0000000..765c023 Binary files /dev/null and b/src/kernels/SRMDNFx2.mat differ diff --git a/src/kernels/SRMDNFx3.mat b/src/kernels/SRMDNFx3.mat new file mode 100644 index 0000000..1d07ee6 Binary files /dev/null and b/src/kernels/SRMDNFx3.mat differ diff --git a/src/kernels/SRMDNFx4.mat b/src/kernels/SRMDNFx4.mat new file mode 100644 index 0000000..dd06ccf Binary files /dev/null and b/src/kernels/SRMDNFx4.mat differ diff --git a/src/kernels/SRMDx1_color.mat b/src/kernels/SRMDx1_color.mat new file mode 100644 index 0000000..81221ab Binary files /dev/null and b/src/kernels/SRMDx1_color.mat differ diff --git a/src/kernels/SRMDx1_gray.mat b/src/kernels/SRMDx1_gray.mat new file mode 100644 index 0000000..533218c Binary files /dev/null and b/src/kernels/SRMDx1_gray.mat differ diff --git a/src/kernels/SRMDx2.mat b/src/kernels/SRMDx2.mat new file mode 100644 index 0000000..43b56dc Binary files /dev/null and b/src/kernels/SRMDx2.mat differ diff --git a/src/kernels/SRMDx3.mat b/src/kernels/SRMDx3.mat new file mode 100644 index 0000000..59fc886 Binary files /dev/null and b/src/kernels/SRMDx3.mat differ diff --git a/src/kernels/SRMDx4.mat b/src/kernels/SRMDx4.mat new file mode 100644 index 0000000..c974d1d Binary files /dev/null and b/src/kernels/SRMDx4.mat differ diff --git a/src/logger.py b/src/logger.py new file mode 100644 index 0000000..b4c4d49 --- /dev/null +++ b/src/logger.py @@ -0,0 +1,25 @@ +import visdom + + +class Logger(object): + def __init__(self, log_dir): + self.last = None + self.viz = visdom.Visdom() + + def scalar_summary(self, tag, value, step, scope=None): + if self.last and self.last['step'] != step: + self.last = None + + if self.last is None: + self.last = {'step': step, 'iter': step, 'epoch': 1} + self.last[tag] = value + + def images_summary(self, tag, images, step): + """Log a list of images.""" + self.viz.images( + images, + opts=dict(title='%s/%d' % (tag, step), caption='%s/%d' % (tag, step)), + ) + + def histo_summary(self, tag, values, step, bins=1000): + pass diff --git a/src/main.py b/src/main.py new file mode 100644 index 0000000..5d37ce4 --- /dev/null +++ b/src/main.py @@ -0,0 +1,68 @@ +import argparse + +import torch +from torch.utils import data +from train import Train +from utils import Utils +from dataset import DIV2K_train + +if __name__=='__main__': + parser = argparse.ArgumentParser() + + # data path + parser.add_argument('--x_path', type=str, default='../data/DIV2K_train_LR_bicubic/X2/') + parser.add_argument('--y_path', type=str, default='../data/DIV2K_train_HR/') + parser.add_argument('--y_path2', type=str, default='../data/waterloo/*.bmp') + parser.add_argument('--y_path3', type=str, default='../data/BSDS300/images/train/*.jpg') + parser.add_argument('--model_save_path', type=str, default='./models') + + # training settings + parser.add_argument('--image_size', type=int, default=40) + parser.add_argument('--total_step', type=int, default=200000) + parser.add_argument('--batch_size', type=int, default=2) + parser.add_argument('--num_workers', type=int, default=2) + parser.add_argument('--num_blocks', type=int, default=11) + parser.add_argument('--num_channels', type=int, default=18) + parser.add_argument('--conv_dim', type=int, default=128) + parser.add_argument('--scale_factor', type=int, default=3) + parser.add_argument('--lr', type=float, default=0.001) + parser.add_argument('--beta1', type=float, default=0.5) + parser.add_argument('--beta2', type=float, default=0.999) + parser.add_argument('--trained_model', type=int, default=None) + parser.add_argument('--device', type=str, default='cpu') + parser.add_argument('--testflag', type=int, default=0) + #misc + parser.add_argument('--log_step', type=int, default=10) + parser.add_argument('--sample_step', type=int, default=100) # todo: 100 + parser.add_argument('--model_save_step', type=int, default=1000) + parser.add_argument('--use_tensorboard', type=bool, default=True) + + parser.add_argument('--log_path', type=str, default='./logs') + parser.add_argument('--result_path', type=str, default='./results') + + # config + config = parser.parse_args() + + # device + config.device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') + # print(config) + print("GPU available = ", torch.cuda.is_available()) + print("# GPU = ", torch.cuda.device_count()) + + # download data + # utils = Utils() + # utils.download('DIV2K_train_HR') # download the HR images + # utils.download('DIV2K_train_LR_bicubic_X2') # download the LR images + # print("[CUSTOM LOG]: data download done") + + # data loader` + dataset = DIV2K_train(config=config) + data_loader = data.DataLoader(dataset=dataset, + batch_size=config.batch_size, + shuffle=True, + num_workers=config.num_workers) + print("Data Size:", len(data_loader.dataset)) + training = Train(data_loader, config) + training.train() + + print("DONE BOOM BOOM!!") diff --git a/src/model.py b/src/model.py new file mode 100644 index 0000000..40a1f6b --- /dev/null +++ b/src/model.py @@ -0,0 +1,34 @@ +import torch.nn as nn +import torch.nn.functional as F + +class SRMD(nn.Module): + # number of conv layers = 12; given in paper + def __init__(self, conv_layers=12, channels=18, conv_dim=128, scale_factor=2): + super(SRMD, self).__init__() + + self.nonlinear_mapping = self.create_layers(conv_layers, channels, conv_dim) + self.conv_layer = nn.Sequential( + nn.Conv2d(conv_dim, 3*scale_factor**2, kernel_size=3, padding=1), + nn.PixelShuffle(scale_factor), + nn.Sigmoid() + ) + + def forward(self, x): + x = self.nonlinear_mapping(x) + x = self.conv_layer(x) + return x + + + def create_layers(self, conv_layers, channels, conv_dim): + layers = [] + in_dim = channels + for i in range(conv_layers): + conv2d = nn.Conv2d(in_dim, conv_dim, kernel_size=3, padding=1) + bn = nn.BatchNorm2d(conv_dim) + relu = nn.ReLU() + + layers += [conv2d, bn, relu] + + in_dim = conv_dim + + return nn.Sequential(*layers) \ No newline at end of file diff --git a/src/train.py b/src/train.py new file mode 100644 index 0000000..ad15b2a --- /dev/null +++ b/src/train.py @@ -0,0 +1,245 @@ +import torch +import torch.nn as nn +import os +from torchvision.utils import save_image, make_grid +from model import SRMD +import numpy as np +import math +import cv2 +import sys +from piqa import SSIM +torch.set_default_tensor_type(torch.DoubleTensor) + +class Train(object): + def __init__(self, data_loader, config): + # Data loader + self.data_loader = data_loader + + # Model hyper-parameters + self.num_blocks = config.num_blocks + self.num_channels = config.num_channels + self.conv_dim = config.conv_dim + self.scale_factor = config.scale_factor + + # Training settings + self.total_step = config.total_step + self.lr = config.lr + self.beta1 = config.beta1 + self.beta2 = config.beta2 + self.trained_model = config.trained_model + self.use_tensorboard = config.use_tensorboard + + # Path and step size + self.log_path = config.log_path + self.result_path = config.result_path + self.model_save_path = config.model_save_path + self.log_step = config.log_step + self.sample_step = config.sample_step + self.model_save_step = config.model_save_step + + # Device configuration + self.device = config.device + + self.build_model() + if self.use_tensorboard: + self.build_tensorboard() + + # Start with trained model + if self.trained_model: + self.load_trained_model() + + def calculate_psnr(self, img1, img2): + # img1 and img2 have range [0, 255] + img1 = img1.astype(np.float64) + + img2 = img2.astype(np.float64) + mse = np.mean((img1 - img2)**2) + if mse == 0: + return float('inf') + return 20 * math.log10(255.0 / math.sqrt(mse)) + + def ssim(self, img1, img2): + C1 = (0.01 * 255)**2 + C2 = (0.03 * 255)**2 + + img1 = img1.astype(np.float64) + img2 = img2.astype(np.float64) + kernel = cv2.getGaussianKernel(11, 1.5) + window = np.outer(kernel, kernel.transpose()) + + mu1 = cv2.filter2D(img1, -1, window)[5:-5, 5:-5] # valid + mu2 = cv2.filter2D(img2, -1, window)[5:-5, 5:-5] + mu1_sq = mu1**2 + mu2_sq = mu2**2 + mu1_mu2 = mu1 * mu2 + sigma1_sq = cv2.filter2D(img1**2, -1, window)[5:-5, 5:-5] - mu1_sq + sigma2_sq = cv2.filter2D(img2**2, -1, window)[5:-5, 5:-5] - mu2_sq + sigma12 = cv2.filter2D(img1 * img2, -1, window)[5:-5, 5:-5] - mu1_mu2 + + ssim_map = ((2 * mu1_mu2 + C1) * (2 * sigma12 + C2)) / ((mu1_sq + mu2_sq + C1) * + (sigma1_sq + sigma2_sq + C2)) + return ssim_map.mean() + + + def calculate_ssim(self, img1, img2): + '''calculate SSIM + the same outputs as MATLAB's + img1, img2: [0, 255] + ''' + if not img1.shape == img2.shape: + raise ValueError('Input images must have the same dimensions.') + if img1.ndim == 2: + return self.ssim(img1, img2) + elif img1.ndim == 3: + if img1.shape[2] == 3: + ssims = [] + for i in range(3): + ssims.append(self.ssim(img1, img2)) + return np.array(ssims).mean() + elif img1.shape[2] == 1: + return self.ssim(np.squeeze(img1), np.squeeze(img2)) + else: + print("Image 1:", img1.ndim) + print("Image 2:", img2.ndim) + raise ValueError('Wrong input image dimensions.') + + def build_model(self): + # model and optimizer + self.model = SRMD(self.num_blocks, self.num_channels, self.conv_dim, self.scale_factor) + self.optimizer = torch.optim.Adam(self.model.parameters(), self.lr, [self.beta1, self.beta2]) + + self.model.to(self.device) + + def load_trained_model(self): + self.load(os.path.join( + self.model_save_path, '{}.pth'.format(self.trained_model))) + print('loaded trained models (step: {})..!'.format(self.trained_model)) + + def load(self, filename): + S = torch.load(filename) + self.model.load_state_dict(S['SR']) + + def build_tensorboard(self): + from logger import Logger + self.logger = Logger(self.log_path) + + def update_lr(self, lr): + for param_group in self.optimizer.param_groups: + param_group['lr'] = lr + + def reset_grad(self): + self.optimizer.zero_grad() + + def detach(self, x): + return x.data + + def train(self): + self.model.train() + + # Reconst loss + reconst_loss = nn.MSELoss() + + # Data iter + data_iter = iter(self.data_loader) + iter_per_epoch = len(self.data_loader) + + # Start with trained model + if self.trained_model: + start = self.trained_model + 1 + else: + start = 0 + + for step in range(start, self.total_step): + # Reset data_iter for each epoch + if (step+1) % iter_per_epoch == 0: + data_iter = iter(self.data_loader) + + actx, x, y = next(data_iter) + actx, x, y = actx.to(self.device), x.to(self.device), y.to(self.device) + y = y.to(torch.float64) + + out = self.model(x) + loss = reconst_loss(out, y) + + self.reset_grad() + + # For decoder + loss.backward(retain_graph=True) + + self.optimizer.step() + + # Print out log info + if (step+1) % self.log_step == 0: + print("[{}/{}] loss: {:.4f}".format(step+1, self.total_step, loss.item())) + + class SSIMLoss(SSIM): + def forward(self, x, y): + print("SSIM Shape GPU:", x.shape, y.shape) + print("SSIM type:", x.dtype, y.dtype) + x, y = x.to('cpu').unsqueeze(0).type(torch.FloatTensor), y.to('cpu').unsqueeze(0).type(torch.FloatTensor) + print("SSIM type double:", x.dtype, y.dtype) + print("SSIM Shape CPU:", x.shape, y.shape) + return 1. - super().forward(x, y) + + criterion_ssim = SSIMLoss() + + # Sample images + if (step+1) % self.sample_step == 0: + self.model.eval() + reconst = self.model(x) + + def to_np(x): + return x.data.cpu().numpy() + + tmp = nn.Upsample(scale_factor=self.scale_factor)(actx.data[:,:,:]) + pairs = torch.cat((tmp.data[0:2,:], reconst.data[0:2,:], y.data[0:2,:]), dim=3) + psnrscore1= self.calculate_psnr(to_np(reconst.data[0,:]),to_np(y.data[0,:])) + print("Shape: ",tmp.data.shape, reconst.data.shape, y.data.shape) + ssimscore1= criterion_ssim(reconst.data[0,:],y.data[0,:]) + psnrscore2= self.calculate_psnr(to_np(reconst.data[1,:]),to_np(y.data[1,:])) + ssimscore2= criterion_ssim(reconst.data[1,:],y.data[1,:]) + with open('score.txt', 'a') as f: + print('test_{}.jpg PSNR1:{} SSIM1:{} PSNR2:{} SSIM2:{}'.format(step + 1,psnrscore1,ssimscore1,psnrscore2,ssimscore2), file=f) + f.close() + pairs = pairs.to('cpu') + grid = make_grid(pairs, 2) + from PIL import Image + tmp = tmp.to('cpu') + tmp = np.squeeze(grid.numpy().transpose((1, 2, 0))) + # tmp = torch.from_numpy(tmp) + tmp = (255 * tmp).astype(np.uint8) + Image.fromarray(tmp).save('./samples/test_%d.jpg' % (step + 1)) + + # Save check points + if (step+1) % self.model_save_step == 0: + self.save(os.path.join(self.model_save_path, '{}.pth'.format(self.trained_model))) + def test(self): + self.model.eval() + reconst_loss = nn.MSELoss() + avg_psnr=0 + avg_ssim=0 + class SSIMLoss(SSIM): + def forward(self, x, y): + return 1. - super().forward(x, y) + + criterion_ssim = SSIMLoss() + data_iter = iter(self.data_loader) + num_iters = len(self.data_loader) + for step in range(num_iters): + actx, x, y = next(data_iter) + actx, x, y = actx.to(self.device), x.to(self.device), y.to(self.device) + y = y.to(torch.float64) + out = self.model(x) + for i in range(self.batch_size): + ssim = criterion_ssim(out.data[i,:], y.data[i,:]) + avg_ssim += ssim + mse = reconst_loss(out, y) + psnr = 10 * math.log10(1 / mse.item()) + avg_psnr += psnr + with open('testscore.txt', 'a') as f: + print(f'PSNR : {avg_psnr/num_iters} SSIM : {avg_ssim/num_iters}',file=f) + f.close() + + def save(self, filename): + model = self.model.state_dict() + torch.save({'SR': model}, filename) diff --git a/src/utils.py b/src/utils.py new file mode 100644 index 0000000..f78c476 --- /dev/null +++ b/src/utils.py @@ -0,0 +1,31 @@ +import os +import shutil + +class Utils: + + DATA_BASE_PATH = '../data/' + URL = { + 'DIV2K_train_LR_bicubic_X2': 'http://data.vision.ee.ethz.ch/cvl/DIV2K/DIV2K_train_LR_bicubic_X2.zip', + 'DIV2K_train_HR': 'http://data.vision.ee.ethz.ch/cvl/DIV2K/DIV2K_train_HR.zip' + } + + def download(self, dataset_name, remove_zip=False): + + # make data folder if it doesn't exist + try: + os.mkdir(self.DATA_BASE_PATH) + except OSError as err: + print(err, "[IGNORE]") + + # wget: save as .zip in ../data + ZIP_FILE_PATH = self.DATA_BASE_PATH + dataset_name + ".zip" + wget_cmd = "wget --continue " + self.URL[dataset_name] + " --output-document " + ZIP_FILE_PATH + + os.system(wget_cmd) + + # unpack dataset + shutil.unpack_archive(ZIP_FILE_PATH, self.DATA_BASE_PATH) + + # remove zip file + if remove_zip: + os.system("rm " + ZIP_FILE_PATH)