diff --git a/README.md b/README.md deleted file mode 100644 index be9ba98..0000000 --- a/README.md +++ /dev/null @@ -1,214 +0,0 @@ -

- -

- -## Improving the Stability of Diffusion Models for Content Consistent Super-Resolution - - - - - -[Lingchen Sun](https://scholar.google.com/citations?hl=zh-CN&tzom=-480&user=ZCDjTn8AAAAJ)1,2 -| [Rongyuan Wu](https://scholar.google.com/citations?user=A-U8zE8AAAAJ&hl=zh-CN)1,2 | -[Zhengqiang Zhang](https://scholar.google.com/citations?hl=zh-CN&user=UX26wSMAAAAJ&view_op=list_works&sortby=pubdate)1,2 | -[Hongwei Yong](https://scholar.google.com.hk/citations?user=Xii74qQAAAAJ&hl=zh-CN)1 | -[Lei Zhang](https://www4.comp.polyu.edu.hk/~cslzhang)1,2 - -1The Hong Kong Polytechnic University, 2OPPO Research Institute - - -## ⏰ Update -- **2024.1.4**: Code and the model for real-world SR are released. -- **2024.1.3**: Paper is released. -- **2023.12.23**: Repo is released. - - -:star: If CCSR is helpful to your images or projects, please help star this repo. Thanks! :hugs: - -## 🌟 Overview Framework -![ccsr](figs/framework.png) - -## 😍 Visual Results -### Demo on Real-World SR -[](https://imgsli.com/MjMxMzA0) [](https://imgsli.com/MjMxMzEx) [](https://imgsli.com/MjMxMzE1) [](https://imgsli.com/MjMxMzI3) -[](https://imgsli.com/MjMxMzEy) [](https://imgsli.com/MjMxMzE5) - -### Comparisons on Real-World SR -For the diffusion model-based method, two restored images that have the best and worst PSNR values over 10 runs are shown for a more comprehensive and fair comparison. - -![ccsr](figs/realworld.png) - -### Comparisons on Bicubic SR -![ccsr](figs/bicubic.png) -For more comparisons, please refer to our paper for details. - -## 📝 Quantitative comparisons -We propose new stability metrics, namely global standard deviation (G-STD) and local standard deviation (L-STD), to respectively measure the image-level and pixel-level variations of the SR results of diffusion-based methods. - -More details about G-STD and L-STD can be found in our paper. - -![ccsr](figs/table.png) -## ⚙ Dependencies and Installation -```shell -## git clone this repository -git clone https://github.com/csslc/CCSR.git -cd CCSR - -# create an environment with python >= 3.9 -conda create -n ccsr python=3.9 -conda activate ccsr -pip install -r requirements.txt -pip install -e git+https://github.com/CompVis/taming-transformers.git@master#egg=taming-transformers -``` -## 🍭 Quick Inference -#### Step 1: Download the pretrained models -- Download the CCSR models from: - -| Model Name | Description | GoogleDrive | OneDive | -|:---------------------|:---------------------------------------------|:--------------------------------------------------------------------------------------|:------------------------------------------------------------------------------| -| real-world_ccsr.ckpt | CCSR model for real-world image restoration. | [download](https://drive.google.com/drive/folders/1jM1mxDryPk9CTuFTvYcraP2XIVzbPiw_?usp=drive_link) | download | -| bicubic_ccsr.ckpt | CCSR model for bicubic image restoration. | download | download | - - -#### Step 2: Prepare testing data -You can put the testing images in the `preset/test_datasets`. - -#### Step 3: Running testing command -``` -python inference_ccsr.py \ ---input preset/test_datasets \ ---config configs/model/ccsr_stage2.yaml \ ---ckpt weights/real-world_ccsr.ckpt \ ---steps 45 \ ---sr_scale 4 \ ---t_max 0.6667 \ ---t_min 0.3333 \ ---color_fix_type adain \ ---output experiments/test \ ---device cuda \ ---repeat_times 1 -``` -You can obtain `N` different SR results by setting `repeat_time` as `N` to test the stability of CCSR. The data folder should be like this: - -``` - experiments/test - ├── sample0 # the first group of SR results - └── sample1 # the second group of SR results - ... - └── sampleN # the N-th group of SR results -``` - -## 📏 Evaluation -1. Calculate the Image Quality Assessment for each restored group. - - Fill in the required information in [cal_iqa.py](cal_iqa/cal_iqa.py) and run, then you can obtain the evaluation results in the folder like this: - ``` - log_path - ├── log_name_npy # save the IQA values of each restored group as the npy files - └── log_name.log # log recode - ``` - -2. Calculate the G-STD value for the diffusion-based SR method. - - Fill in the required information in [iqa_G-STD.py](cal_iqa/iqa_G-STD.py) and run, then you can obtain the mean IQA values of N restored groups and G-STD value. - -3. Calculate the L-STD value for the diffusion-based SR method. - - Fill in the required information in [iqa_L-STD.py](cal_iqa/iqa_L-STD.py) and run, then you can obtain the L-STD value. - - -## 🚋 Train - -#### Step1: Prepare training data - -1. Generate file list of training set and validation set. - - ```shell - python scripts/make_file_list.py \ - --img_folder [hq_dir_path] \ - --val_size [validation_set_size] \ - --save_folder [save_dir_path] \ - --follow_links - ``` - - This script will collect all image files in `img_folder` and split them into training set and validation set automatically. You will get two file lists in `save_folder`, each line in a file list contains an absolute path of an image file: - - ``` - save_dir_path - ├── train.list # training file list - └── val.list # validation file list - ``` - -2. Configure training set and validation set. - - For real-world image restoration, fill in the following configuration files with appropriate values. - - - [training set](configs/dataset/general_deg_stablesr_realesrgan_train.yaml) and [validation set](configs/dataset/general_deg_stablesr_realesrgan_val.yaml) for **Real-ESRGAN** degradation. - -#### Step2: Train Stage1 Model -1. Download pretrained [Stable Diffusion v2.1](https://huggingface.co/stabilityai/stable-diffusion-2-1-base) to provide generative capabilities. - - ```shell - wget https://huggingface.co/stabilityai/stable-diffusion-2-1-base/resolve/main/v2-1_512-ema-pruned.ckpt --no-check-certificate - ``` - -2. Create the initial model weights. - - ```shell - python scripts/make_stage2_init_weight.py \ - --cldm_config configs/model/ccsr_stage1.yaml \ - --sd_weight [sd_v2.1_ckpt_path] \ - --output weights/init_weight_ccsr.ckpt - ``` - -3. Configure training-related information. - - Fill in the configuration file of [training of stage1](configs/train_ccsr_stage1.yaml) with appropriate settings. - -4. Start training. - - ```shell - python train.py --config configs/train_ccsr_stage1.yaml - ``` - -#### Step3: Train Stage2 Model -1. Configure training-related information. - - Fill in the configuration file of [training of stage2](configs/train_ccsr_stage2.yaml) with appropriate settings. - -2. Start training. - ```shell - python train.py --config configs/train_ccsr_stage2.yaml - ``` - -### Citations -If our code helps your research or work, please consider citing our paper. -The following are BibTeX references: - -``` -@article{sun2023ccsr, - title={Improving the Stability of Diffusion Models for Content Consistent Super-Resolution}, - author={Sun, Lingchen and Wu, Rongyuan and Zhang, Zhengqiang and Yong, Hongwei and Zhang, Lei}, - journal={arXiv preprint arXiv:2401.00877}, - year={2024} -} -``` - -### License -This project is released under the [Apache 2.0 license](LICENSE). - -### Acknowledgement -This project is based on [ControlNet](https://github.com/lllyasviel/ControlNet), [BasicSR](https://github.com/XPixelGroup/BasicSR) and [DiffBIR](https://github.com/XPixelGroup/DiffBIR). Some codes are brought from [StableSR](https://github.com/IceClear/StableSR). Thanks for their awesome works. - -### Contact -If you have any questions, please contact: ling-chen.sun@connect.polyu.hk - - -
-statistics - -![visitors](https://visitor-badge.laobi.icu/badge?page_id=csslc/CCSR) - -
- - diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..5109219 --- /dev/null +++ b/__init__.py @@ -0,0 +1,4 @@ +from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS + +WEB_DIRECTORY = "./web" +__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS", "WEB_DIRECTORY"] \ No newline at end of file diff --git a/__pycache__/__init__.cpython-310.pyc b/__pycache__/__init__.cpython-310.pyc new file mode 100644 index 0000000..137e2dc Binary files /dev/null and b/__pycache__/__init__.cpython-310.pyc differ diff --git a/__pycache__/nodes.cpython-310.pyc b/__pycache__/nodes.cpython-310.pyc new file mode 100644 index 0000000..64d24a3 Binary files /dev/null and b/__pycache__/nodes.cpython-310.pyc differ diff --git a/cal_iqa/cal_iqa.py b/cal_iqa/cal_iqa.py deleted file mode 100644 index bc6d6b4..0000000 --- a/cal_iqa/cal_iqa.py +++ /dev/null @@ -1,236 +0,0 @@ -# evaluate the restored images with IQA -# PSNR, SSIM, LPIPS are given as example, you can add more IQA in this file - -import cv2 -import argparse, os, sys, glob -import logging -from datetime import datetime - -import pyiqa - -from torch.utils import data as data -import glob -import numpy as np -import math -import random -import torch - -def get_timestamp(): - return datetime.now().strftime('%y%m%d-%H%M%S') - -def setup_logger(logger_name, root, phase, level=logging.INFO, screen=False, tofile=False): - '''set up logger''' - lg = logging.getLogger(logger_name) - formatter = logging.Formatter('%(asctime)s.%(msecs)03d - %(levelname)s: %(message)s', - datefmt='%y-%m-%d %H:%M:%S') - lg.setLevel(level) - if tofile: - log_file = os.path.join(root, phase + '_{}.log'.format(get_timestamp())) - fh = logging.FileHandler(log_file, mode='w') - fh.setFormatter(formatter) - lg.addHandler(fh) - if screen: - sh = logging.StreamHandler() - sh.setFormatter(formatter) - lg.addHandler(sh) - -def dict2str(opt, indent_l=1): - '''dict to string for logger''' - msg = '' - for k in opt: - if isinstance(v, dict): - msg += ' ' * (indent_l * 2) + k + ':[\n' - msg += dict2str(v, indent_l + 1) - msg += ' ' * (indent_l * 2) + ']\n' - else: - msg += ' ' * (indent_l * 2) + k + ': ' + str(v) + '\n' - return msg - -def img2tensor(imgs, bgr2rgb=True, float32=True): - """from BasicSR - Numpy array to tensor. - - Args: - imgs (list[ndarray] | ndarray): Input images. - bgr2rgb (bool): Whether to change bgr to rgb. - float32 (bool): Whether to change to float32. - - Returns: - list[tensor] | tensor: Tensor images. If returned results only have - one element, just return tensor. - """ - - def _totensor(img, bgr2rgb, float32): - if img.shape[2] == 3 and bgr2rgb: - if img.dtype == 'float64': - img = img.astype('float32') - img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) - img = torch.from_numpy(img.transpose(2, 0, 1)) - if float32: - img = img.float() - return img - - if isinstance(imgs, list): - return [_totensor(img, bgr2rgb, float32) for img in imgs] - else: - return _totensor(imgs, bgr2rgb, float32) - - -def main(): - parser = argparse.ArgumentParser() - - parser.add_argument( - "--init-imgs", - nargs="+", - help="path to the input image", - default=['****/sample0', - '****/sample1', - '****/sample2', - '****/sample3', - '****/sample4', - '****/sample5', - '****/sample6', - '****/sample7', - '****/sample8', - '****/sample9', - ], - ) - - parser.add_argument( - "--init-imgs-names", - nargs="+", - help="name of the input image", - default=['****-0', '****-1', '****-2', '****-3', '****-4', '****-5', '****-6', '****-7', '****-8', '****-9' - ], - ) - - parser.add_argument( - "--gt-imgs", - nargs="+", - help="path to the gt image, you need to add the paths of gt folders corresponding to init-imgs", - default=['****', '****', '****', '****', '****','****', - '****', '****', '****', '****', '****','****', - ], - - ) - - parser.add_argument( - "--log", - type=str, - nargs="?", - help="path to the log", - default='/home/notebook/data/group/SunLingchen/code/CCSR/CCSR-main/experiments') - - parser.add_argument( - "--log-name", - type=str, - nargs="?", - help="name of your log", - default='test', - ) - - parser.add_argument( - "--num_img", - type=int, - nargs="?", - help="the number of images evaluated in the folder; 0: all the images are evaludated.", - default=0, - ) - - opt = parser.parse_args() - device = torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu") - - os.makedirs(opt.log, exist_ok=True) - # init logger - setup_logger('base', opt.log, 'test_' + opt.log_name, level=logging.INFO, - screen=True, tofile=True) - logger = logging.getLogger('base') - logger.info(opt) - - # init metrics: you can add more metrics here - - iqa_ssim = pyiqa.create_metric('ssim', test_y_channel=True, color_space='ycbcr').to(device) - iqa_psnr = pyiqa.create_metric('psnr', test_y_channel=True, color_space='ycbcr').to(device) - iqa_lpips = pyiqa.create_metric('lpips', device=device) - - - - for dir_idx in range(len(opt.init_imgs)): - - gt_dir = opt.gt_imgs[dir_idx] - img_gt_list = sorted(glob.glob(os.path.join(gt_dir, '*.png'))) - - img_sr_dir = opt.init_imgs[dir_idx] - img_sr_list = sorted(glob.glob(os.path.join(img_sr_dir, '*.png'))) - print(opt.init_imgs_names[dir_idx]) - print(f'GT IMAGES LEN IS {len(img_gt_list)}, SR IMAGES LEN IS {len(img_sr_list)}') - - assert len(img_gt_list) == len(img_sr_list) - - - for dir_idx in range(len(opt.init_imgs)): - - # record metrics - metrics = {} - metrics['psnr'], metrics['ssim'], metrics['lpips'] = \ - [], [], [] - - gt_dir = opt.gt_imgs[dir_idx] - img_gt_list = sorted(glob.glob(os.path.join(gt_dir, '*.png'))) - - img_sr_dir = opt.init_imgs[dir_idx] - img_sr_list = sorted(glob.glob(os.path.join(img_sr_dir, '*.png'))) - - if opt.num_img != 0: - img_gt_list = img_gt_list[0:opt.num_img] - img_sr_list = img_sr_list[0:opt.num_img] - - PSNR_all, SSIM_all, lpips_all = 0.0, 0.0, 0.0 - logger.info('\nTesting [{:s}]...'.format(opt.init_imgs_names[dir_idx])) - - for img_idx in range(len(img_sr_list)): - img_sr_name = os.path.basename(img_sr_list[img_idx]) - print(f'Processing {img_sr_name} ...') - print(img_sr_list[img_idx]) - - input_sr_img = cv2.imread(img_sr_list[img_idx], cv2.IMREAD_COLOR) - sr = img2tensor(input_sr_img, bgr2rgb=True, float32=True).unsqueeze(0).cuda().contiguous() - - input_gt_img = cv2.imread(img_gt_list[img_idx], cv2.IMREAD_COLOR) - hr = img2tensor(input_gt_img, bgr2rgb=True, float32=True).unsqueeze(0).cuda().contiguous() - - if sr.shape != hr.shape: - continue - - # PSNR: convert the ycbcr to calculate - hr = hr[..., 4:-4, 4:-4]/255. - sr = sr[..., 4:-4, 4:-4]/255. - - PSNR_now = iqa_psnr(sr, hr).item() - PSNR_all += PSNR_now - metrics['psnr'].append(PSNR_now) - - # SSIM - ssim_now = iqa_ssim(sr, hr).item() - SSIM_all += ssim_now - metrics['ssim'].append(ssim_now) - - # lpips - lpips_now = iqa_lpips(sr, hr).item() - lpips_all += lpips_now - metrics['lpips'].append(lpips_now) - - logger.info('{:20s}_{} - PSNR: {:.6f} dB; SSIM: {:.6f}; LPIPS: {:.6f}'.format(opt.init_imgs_names[dir_idx], img_sr_name, PSNR_now, ssim_now, lpips_now)) - - PSNR_all = round(PSNR_all/len(img_sr_list) , 4) - SSIM_all = round(SSIM_all/len(img_sr_list) , 4) - lpips_all = round(lpips_all/len(img_sr_list) , 4) - - logger.info('{:20s}_all - PSNR: {:.6f} dB; SSIM: {:.6f}; LPIPS: {:.6f}'.format(opt.init_imgs_names[dir_idx], PSNR_all, SSIM_all, lpips_all)) - - # save metrics - npy_path = os.path.join(opt.log, 'test_' + opt.log_name + '_npy') - os.makedirs(npy_path, exist_ok=True) - np.save(npy_path + '/' + opt.init_imgs_names[dir_idx]+'.npy', metrics) - -main() \ No newline at end of file diff --git a/cal_iqa/iqa_G-STD.py b/cal_iqa/iqa_G-STD.py deleted file mode 100644 index bf3fda8..0000000 --- a/cal_iqa/iqa_G-STD.py +++ /dev/null @@ -1,111 +0,0 @@ -# evaluate the G-STD of restored N images - - -import os -import numpy as np -import glob -from datetime import datetime -import logging - -def get_timestamp(): - return datetime.now().strftime('%y%m%d-%H%M%S') - -def setup_logger(logger_name, root, phase, level=logging.INFO, screen=False, tofile=False): - '''set up logger''' - lg = logging.getLogger(logger_name) - formatter = logging.Formatter('%(asctime)s.%(msecs)03d - %(levelname)s: %(message)s', - datefmt='%y-%m-%d %H:%M:%S') - lg.setLevel(level) - if tofile: - log_file = os.path.join(root, phase + '_{}.log'.format(get_timestamp())) - fh = logging.FileHandler(log_file, mode='w') - fh.setFormatter(formatter) - lg.addHandler(fh) - if screen: - sh = logging.StreamHandler() - sh.setFormatter(formatter) - lg.addHandler(sh) - -def img2tensor(imgs, bgr2rgb=True, float32=True): - """from BasicSR - Numpy array to tensor. - - Args: - imgs (list[ndarray] | ndarray): Input images. - bgr2rgb (bool): Whether to change bgr to rgb. - float32 (bool): Whether to change to float32. - - Returns: - list[tensor] | tensor: Tensor images. If returned results only have - one element, just return tensor. - """ - - def _totensor(img, bgr2rgb, float32): - if img.shape[2] == 3 and bgr2rgb: - if img.dtype == 'float64': - img = img.astype('float32') - img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) - img = torch.from_numpy(img.transpose(2, 0, 1)) - if float32: - img = img.float() - return img - - if isinstance(imgs, list): - return [_totensor(img, bgr2rgb, float32) for img in imgs] - else: - return _totensor(imgs, bgr2rgb, float32) - -# log_name -name_log = 'DIV2K-valid' -path_log = '****/G-STD' -os.makedirs(path_log, exist_ok=True) - -# load metrics -npy_file_path = '****' -npy_file_lists = sorted(glob.glob(os.path.join(npy_file_path, '*npy'))) -l_sample = len(npy_file_lists) -a = np.load(npy_file_lists[0], allow_pickle=True).item() -l_file = len(a['psnr']) - -# init logger -setup_logger('base', path_log, 'test_' + name_log, level=logging.INFO, screen=True, tofile=True) -logger = logging.getLogger('base') -logger.info(name_log) - -# init the metrics: you can add other metrics here -metric_psnr = np.zeros([l_file, l_sample]) -metric_ssim = np.zeros([l_file, l_sample]) -metric_lpips = np.zeros([l_file, l_sample]) - -i = 0 -for npy_file in npy_file_lists: - # npy_file_list = sorted(glob.glob(os.path.join(npy_files, '*'))) - a = np.load(npy_file, allow_pickle=True).item() - metric_psnr[:, i] = np.array(a['psnr']) - metric_ssim[:, i] = np.array(a['ssim']) - metric_lpips[:, i] = np.array(a['lpips']) - i = i + 1 - -# calculate the mean of the metrics -mean_psnr = np.mean(metric_psnr, axis=1) -mean_ssim = np.mean(metric_ssim, axis=1) -mean_lpips = np.mean(metric_lpips, axis=1) - -mean_mean_psnr = np.mean(mean_psnr) -mean_mean_ssim = np.mean(mean_ssim) -mean_mean_lpips = np.mean(mean_lpips) - - -# calculate the std of the metrics -std_psnr = np.std(metric_psnr, axis=1) -std_ssim = np.std(metric_ssim, axis=1) -std_lpips = np.std(metric_lpips, axis=1) - -mean_std_psnr = np.mean(std_psnr) -mean_std_ssim = np.mean(std_ssim) -mean_std_lpips = np.mean(std_lpips) - - -logger.info('mean_mean - PSNR: {:.6f} dB; SSIM: {:.6f}; LPIPS: {:.6f}'.format(mean_mean_psnr, mean_mean_ssim, mean_mean_lpips)) - -logger.info('G_STD - PSNR: {:.6f} dB; SSIM: {:.6f}; LPIPS: {:.6f}'.format(mean_std_psnr, mean_std_ssim, mean_std_lpips)) \ No newline at end of file diff --git a/cal_iqa/iqa_L-STD.py b/cal_iqa/iqa_L-STD.py deleted file mode 100644 index ab56658..0000000 --- a/cal_iqa/iqa_L-STD.py +++ /dev/null @@ -1,190 +0,0 @@ -# evaluate the L-STD of restored N images - - -import cv2 -import argparse, os, sys, glob -import logging -from datetime import datetime - -from torch.utils import data as data -import glob -import numpy as np -import math -import random -import torch - - -def img2tensor(imgs, bgr2rgb=True, float32=True): - """from BasicSR - Numpy array to tensor. - - Args: - imgs (list[ndarray] | ndarray): Input images. - bgr2rgb (bool): Whether to change bgr to rgb. - float32 (bool): Whether to change to float32. - - Returns: - list[tensor] | tensor: Tensor images. If returned results only have - one element, just return tensor. - """ - - def _totensor(img, bgr2rgb, float32): - if img.shape[2] == 3 and bgr2rgb: - if img.dtype == 'float64': - img = img.astype('float32') - img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) - img = torch.from_numpy(img.transpose(2, 0, 1)) - if float32: - img = img.float() - return img - - if isinstance(imgs, list): - return [_totensor(img, bgr2rgb, float32) for img in imgs] - else: - return _totensor(imgs, bgr2rgb, float32) - - -def get_timestamp(): - return datetime.now().strftime('%y%m%d-%H%M%S') - -def setup_logger(logger_name, root, phase, level=logging.INFO, screen=False, tofile=False): - '''set up logger''' - lg = logging.getLogger(logger_name) - formatter = logging.Formatter('%(asctime)s.%(msecs)03d - %(levelname)s: %(message)s', - datefmt='%y-%m-%d %H:%M:%S') - lg.setLevel(level) - if tofile: - log_file = os.path.join(root, phase + '_{}.log'.format(get_timestamp())) - fh = logging.FileHandler(log_file, mode='w') - fh.setFormatter(formatter) - lg.addHandler(fh) - if screen: - sh = logging.StreamHandler() - sh.setFormatter(formatter) - lg.addHandler(sh) - -def dict2str(opt, indent_l=1): - '''dict to string for logger''' - msg = '' - for k in opt: - if isinstance(v, dict): - msg += ' ' * (indent_l * 2) + k + ':[\n' - msg += dict2str(v, indent_l + 1) - msg += ' ' * (indent_l * 2) + ']\n' - else: - msg += ' ' * (indent_l * 2) + k + ': ' + str(v) + '\n' - return msg - -def calc_psnr(sr, hr): - sr, hr = sr.double(), hr.double() - diff = (sr - hr) / 255.00 - mse = diff.pow(2).mean() - psnr = -10 * math.log10(mse) - return float(psnr) - -def rgb_to_ycbcr(image: torch.Tensor) -> torch.Tensor: - r"""Convert an RGB image to YCbCr. - - Args: - image (torch.Tensor): RGB Image to be converted to YCbCr. - - Returns: - torch.Tensor: YCbCr version of the image. - """ - - if not torch.is_tensor(image): - raise TypeError("Input type is not a torch.Tensor. Got {}".format(type(image))) - - if len(image.shape) < 3 or image.shape[-3] != 3: - raise ValueError("Input size must have a shape of (*, 3, H, W). Got {}".format(image.shape)) - - image = image / 255. ## image in range (0, 1) - r: torch.Tensor = image[..., 0, :, :] - g: torch.Tensor = image[..., 1, :, :] - b: torch.Tensor = image[..., 2, :, :] - - y: torch.Tensor = 65.481 * r + 128.553 * g + 24.966 * b + 16.0 - cb: torch.Tensor = -37.797 * r + -74.203 * g + 112.0 * b + 128.0 - cr: torch.Tensor = 112.0 * r + -93.786 * g + -18.214 * b + 128.0 - - return torch.stack((y, cb, cr), -3) - -def main(): - parser = argparse.ArgumentParser() - - parser.add_argument( - "--init-imgs", - nargs="+", - help="path to the input image", - default=['****/sample0', - '****/sample1', - '****/sample2', - '****/sample3', - '****/sample4', - '****/sample5', - '****/sample6', - '****/sample7', - '****/sample8', - '****/sample9', - ] - ) - - - - opt = parser.parse_args() - device = torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu") - - name_log = 'ours' - path_log = '****/L-STD' - os.makedirs(path_log, exist_ok=True) - - # init logger - setup_logger('base', path_log, 'test_' + name_log, level=logging.INFO, screen=True, tofile=True) - logger = logging.getLogger('base') - logger.info(name_log) - - num_samples = len(opt.init_imgs) - - # combine all the sampled data - - img_sr_dirs = [] - for dir_idx in range(num_samples): - img_sr_dir = opt.init_imgs[dir_idx] - img_sr_dirs.append(img_sr_dir) - - img_sr_list = sorted(glob.glob(os.path.join(img_sr_dir, '*.png'))) - num_imgs = len(img_sr_list) - input_sr_img = cv2.imread(img_sr_list[0], cv2.IMREAD_COLOR) - sr = img2tensor(input_sr_img, bgr2rgb=True, float32=True).unsqueeze(0).cuda().contiguous() - - sr = sr[..., 4:-4, 4:-4]/255. #[0,1] - B, C, H, W = sr.shape - sr = sr[0] - - std_alls = [] - - for img_idx in range(len(img_sr_list)): - d = torch.zeros([num_samples,C,H,W]) - for dir_idx in range(num_samples): - img_sr_dir = opt.init_imgs[dir_idx] - img_sr_list = sorted(glob.glob(os.path.join(img_sr_dir, '*.png'))) - - input_sr_img = cv2.imread(img_sr_list[img_idx], cv2.IMREAD_COLOR) - img_sr_name = os.path.basename(img_sr_list[img_idx]) - - sr = img2tensor(input_sr_img, bgr2rgb=True, float32=True).unsqueeze(0).cuda().contiguous() - - sr = sr[..., 4:-4, 4:-4]/255. #[0,1] - sr = sr[0] - d[dir_idx,...] = sr - - d = np.array(d) - stds = np.std(d, axis=0) - stds = np.mean(stds) - logger.info('_{} - L-STD: {:.6f}'.format(img_sr_name, stds)) - std_alls.append(stds) - std_all = np.mean(std_alls) - print(std_all) - logger.info('_all - L-STD: {:.6f}; '.format(std_all)) - -main() \ No newline at end of file diff --git a/configs/dataset/general_deg_stablesr_realesrgan_train.yaml b/configs/dataset/general_deg_stablesr_realesrgan_train.yaml deleted file mode 100644 index ba0a4e3..0000000 --- a/configs/dataset/general_deg_stablesr_realesrgan_train.yaml +++ /dev/null @@ -1,63 +0,0 @@ -dataset: - target: dataset.realesrgan.RealESRGANDataset - params: - # Path to the file list. - file_list: preset/train_datasets/train.list - out_size: 512 - crop_type: center - - use_hflip: true - use_rot: false - - blur_kernel_size: 21 - kernel_list: ['iso', 'aniso', 'generalized_iso', 'generalized_aniso', 'plateau_iso', 'plateau_aniso'] - kernel_prob: [0.45, 0.25, 0.12, 0.03, 0.12, 0.03] - sinc_prob: 0.1 - blur_sigma: [0.2, 1.5] - betag_range: [0.5, 2.0] - betap_range: [1, 1.5] - - blur_kernel_size2: 11 - kernel_list2: ['iso', 'aniso', 'generalized_iso', 'generalized_aniso', 'plateau_iso', 'plateau_aniso'] - kernel_prob2: [0.45, 0.25, 0.12, 0.03, 0.12, 0.03] - sinc_prob2: 0.1 - blur_sigma2: [0.2, 1.0] - betag_range2: [0.5, 2.0] - betap_range2: [1, 1.5] - - final_sinc_prob: 0.8 - -data_loader: - batch_size: 16 - shuffle: true - num_workers: 16 - prefetch_factor: 2 - drop_last: true - -batch_transform: - target: dataset.batch_transform.RealESRGANBatchTransform - params: - use_sharpener: false - resize_hq: false - # Queue size of training pool, this should be multiples of batch_size. - queue_size: 192 - - # the first degradation process - resize_prob: [0.2, 0.7, 0.1] # up, down, keep - resize_range: [0.3, 1.5] - gaussian_noise_prob: 0.5 - noise_range: [1, 15] - poisson_scale_range: [0.05, 2.0] - gray_noise_prob: 0.4 - jpeg_range: [60, 95] - - # the second degradation process - stage2_scale: 4 - second_blur_prob: 0.5 - resize_prob2: [0.3, 0.4, 0.3] # up, down, keep - resize_range2: [0.6, 1.2] - gaussian_noise_prob2: 0.5 - noise_range2: [1, 12] - poisson_scale_range2: [0.05, 1.0] - gray_noise_prob2: 0.4 - jpeg_range2: [60, 100] diff --git a/configs/dataset/general_deg_stablesr_realesrgan_val.yaml b/configs/dataset/general_deg_stablesr_realesrgan_val.yaml deleted file mode 100644 index 92ae1f8..0000000 --- a/configs/dataset/general_deg_stablesr_realesrgan_val.yaml +++ /dev/null @@ -1,64 +0,0 @@ -dataset: - target: dataset.realesrgan.RealESRGANDataset - params: - # Path to the file list. - file_list: preset/train_datasets/val.list - - out_size: 512 - crop_type: center - - use_hflip: true - use_rot: false - - blur_kernel_size: 21 - kernel_list: ['iso', 'aniso', 'generalized_iso', 'generalized_aniso', 'plateau_iso', 'plateau_aniso'] - kernel_prob: [0.45, 0.25, 0.12, 0.03, 0.12, 0.03] - sinc_prob: 0.1 - blur_sigma: [0.2, 1.5] - betag_range: [0.5, 2.0] - betap_range: [1, 1.5] - - blur_kernel_size2: 11 - kernel_list2: ['iso', 'aniso', 'generalized_iso', 'generalized_aniso', 'plateau_iso', 'plateau_aniso'] - kernel_prob2: [0.45, 0.25, 0.12, 0.03, 0.12, 0.03] - sinc_prob2: 0.1 - blur_sigma2: [0.2, 1.0] - betag_range2: [0.5, 2.0] - betap_range2: [1, 1.5] - - final_sinc_prob: 0.8 - -data_loader: - batch_size: 1 - shuffle: false - num_workers: 16 - prefetch_factor: 2 - drop_last: true - -batch_transform: - target: dataset.batch_transform.RealESRGANBatchTransform - params: - use_sharpener: false - resize_hq: false - # Queue size of training pool, this should be multiples of batch_size. - queue_size: 256 - - # the first degradation process - resize_prob: [0.2, 0.7, 0.1] # up, down, keep - resize_range: [0.3, 1.5] - gaussian_noise_prob: 0.5 - noise_range: [1, 15] - poisson_scale_range: [0.05, 2.0] - gray_noise_prob: 0.4 - jpeg_range: [60, 95] - - # the second degradation process - stage2_scale: 4 - second_blur_prob: 0.5 - resize_prob2: [0.3, 0.4, 0.3] # up, down, keep - resize_range2: [0.6, 1.2] - gaussian_noise_prob2: 0.5 - noise_range2: [1, 12] - poisson_scale_range2: [0.05, 1.0] - gray_noise_prob2: 0.4 - jpeg_range2: [60, 100] \ No newline at end of file diff --git a/configs/model/ccsr_stage2.yaml b/configs/model/ccsr_stage2.yaml index 51a2923..fc68fc8 100644 --- a/configs/model/ccsr_stage2.yaml +++ b/configs/model/ccsr_stage2.yaml @@ -1,4 +1,4 @@ -target: model.ccsr_stage2.ControlLDM +target: ComfyUI-CCSR.model.ccsr_stage2.ControlLDM params: linear_start: 0.00085 linear_end: 0.0120 @@ -24,7 +24,7 @@ params: learning_rate: 5e-6 control_stage_config: - target: model.ccsr_stage2.ControlNet + target: ComfyUI-CCSR.model.ccsr_stage2.ControlNet params: use_checkpoint: True image_size: 32 # unused @@ -42,7 +42,7 @@ params: legacy: False unet_config: - target: model.ccsr_stage2.ControlledUnetModel + target: ComfyUI-CCSR.model.ccsr_stage2.ControlledUnetModel params: use_checkpoint: True image_size: 32 # unused @@ -60,7 +60,7 @@ params: legacy: False first_stage_config: - target: ldm.models.autoencoder.AutoencoderKL + target: ComfyUI-CCSR.ldm.models.autoencoder.AutoencoderKL params: embed_dim: 4 monitor: val/rec_loss @@ -84,13 +84,13 @@ params: target: torch.nn.Identity cond_stage_config: - target: ldm.modules.encoders.modules.FrozenOpenCLIPEmbedder + target: ComfyUI-CCSR.ldm.modules.encoders.modules.FrozenOpenCLIPEmbedder params: freeze: True layer: "penultimate" lossconfig: - target: ldm.modules.losses.LPIPSWithDiscriminator + target: ComfyUI-CCSR.ldm.modules.losses.LPIPSWithDiscriminator params: disc_start: 1.0 kl_weight: 0 diff --git a/configs/train_ccsr_stage1.yaml b/configs/train_ccsr_stage1.yaml deleted file mode 100644 index 4987365..0000000 --- a/configs/train_ccsr_stage1.yaml +++ /dev/null @@ -1,47 +0,0 @@ -data: - target: dataset.data_module.BIRDataModule - params: - # Path to training set configuration file. - train_config: configs/dataset/general_deg_stablesr_realesrgan_train.yaml - # Path to validation set configuration file. - val_config: configs/dataset/general_deg_stablesr_realesrgan_val.yaml - -model: - # You can set learning rate in the following configuration file. - config: configs/model/ccsr_stage1.yaml - # Path to the checkpoints or weights you want to resume. At the begining, - # this should be set to the initial weights created by scripts/make_stage2_init_weight.py. - resume: weights/init_weight_ccsr.ckpt - -lightning: - seed: 231 - - trainer: - accelerator: ddp - precision: 32 - # Indices of GPUs used for training. - gpus: [0,1,2,3,] - # Path to save logs and checkpoints. - default_root_dir: experiments/test_ccsr_stage1 - # Max number of training steps (batches). - max_steps: 25001 - # Validation frequency in terms of training steps. - val_check_interval: 390 - log_every_n_steps: 50 - # Accumulate gradients from multiple batches so as to increase batch size. - accumulate_grad_batches: 3 - - callbacks: - - target: model.callbacks.ImageLogger - params: - # Log frequency of image logger. - log_every_n_steps: 150000 - max_images_each_step: 4 - log_images_kwargs: ~ - - - target: model.callbacks.ModelCheckpoint - params: - # Frequency of saving checkpoints. - every_n_train_steps: 500 - save_top_k: -1 - filename: "{step}" diff --git a/configs/train_ccsr_stage2.yaml b/configs/train_ccsr_stage2.yaml deleted file mode 100644 index 9671fa1..0000000 --- a/configs/train_ccsr_stage2.yaml +++ /dev/null @@ -1,49 +0,0 @@ -data: - target: dataset.data_module.BIRDataModule - params: - # Path to training set configuration file. - train_config: configs/dataset/general_deg_stablesr_realesrgan_train.yaml - # Path to validation set configuration file. - val_config: configs/dataset/general_deg_stablesr_realesrgan_val.yaml - -model: - # You can set learning rate in the following configuration file. - config: configs/model/ccsr_stage2.yaml - # Path to the checkpoints or weights you want to resume. At the begining, - # this should be set to the weights obtained from ccsr_stage1 - resume: /ccsr_stage1.ckpt - - -lightning: - seed: 231 - - trainer: - accelerator: ddp - precision: 32 - # Indices of GPUs used for training. - gpus: [0,1,2,3,] - # Path to save logs and checkpoints. - default_root_dir: experiments/test_ccsr_stage2 - # Max number of training steps (batches). - max_steps: 60 - # Validation frequency in terms of training steps. - val_check_interval: 390 - log_every_n_steps: 50 - # 50 - # Accumulate gradients from multiple batches so as to increase batch size. - accumulate_grad_batches: 8 - - callbacks: - - target: model.callbacks.ImageLogger - params: - # Log frequency of image logger. - log_every_n_steps: 150000 - max_images_each_step: 4 - log_images_kwargs: ~ - - - target: model.callbacks.ModelCheckpoint - params: - # Frequency of saving checkpoints. - every_n_train_steps: 20 - save_top_k: -1 - filename: "{step}" diff --git a/dataset/__pycache__/batch_transform.cpython-310.pyc b/dataset/__pycache__/batch_transform.cpython-310.pyc deleted file mode 100644 index 6bfee36..0000000 Binary files a/dataset/__pycache__/batch_transform.cpython-310.pyc and /dev/null differ diff --git a/dataset/__pycache__/batch_transform.cpython-37.pyc b/dataset/__pycache__/batch_transform.cpython-37.pyc deleted file mode 100644 index c2ebb17..0000000 Binary files a/dataset/__pycache__/batch_transform.cpython-37.pyc and /dev/null differ diff --git a/dataset/__pycache__/bicubic_torchvision.cpython-310.pyc b/dataset/__pycache__/bicubic_torchvision.cpython-310.pyc deleted file mode 100644 index 535fa31..0000000 Binary files a/dataset/__pycache__/bicubic_torchvision.cpython-310.pyc and /dev/null differ diff --git a/dataset/__pycache__/data_module.cpython-310.pyc b/dataset/__pycache__/data_module.cpython-310.pyc deleted file mode 100644 index 9858ce9..0000000 Binary files a/dataset/__pycache__/data_module.cpython-310.pyc and /dev/null differ diff --git a/dataset/__pycache__/data_module.cpython-37.pyc b/dataset/__pycache__/data_module.cpython-37.pyc deleted file mode 100644 index 6b67885..0000000 Binary files a/dataset/__pycache__/data_module.cpython-37.pyc and /dev/null differ diff --git a/dataset/__pycache__/projectdata.cpython-310.pyc b/dataset/__pycache__/projectdata.cpython-310.pyc deleted file mode 100644 index c553acd..0000000 Binary files a/dataset/__pycache__/projectdata.cpython-310.pyc and /dev/null differ diff --git a/dataset/__pycache__/projectdata.cpython-37.pyc b/dataset/__pycache__/projectdata.cpython-37.pyc deleted file mode 100644 index 08ed181..0000000 Binary files a/dataset/__pycache__/projectdata.cpython-37.pyc and /dev/null differ diff --git a/dataset/__pycache__/realesrgan.cpython-310.pyc b/dataset/__pycache__/realesrgan.cpython-310.pyc deleted file mode 100644 index b556ade..0000000 Binary files a/dataset/__pycache__/realesrgan.cpython-310.pyc and /dev/null differ diff --git a/dataset/__pycache__/realesrgan.cpython-37.pyc b/dataset/__pycache__/realesrgan.cpython-37.pyc deleted file mode 100644 index bb85401..0000000 Binary files a/dataset/__pycache__/realesrgan.cpython-37.pyc and /dev/null differ diff --git a/dataset/batch_transform.py b/dataset/batch_transform.py deleted file mode 100644 index ef1a70e..0000000 --- a/dataset/batch_transform.py +++ /dev/null @@ -1,374 +0,0 @@ -from typing import Any, overload, Dict, Union, List, Sequence -import random - -import torch -from torch.nn import functional as F -import numpy as np - -from utils.image import USMSharp, DiffJPEG, filter2D -from utils.degradation import ( - random_add_gaussian_noise_pt, random_add_poisson_noise_pt -) - - -class BatchTransform: - - @overload - def __call__(self, batch: Any) -> Any: - ... - - -class IdentityBatchTransform(BatchTransform): - - def __call__(self, batch: Any) -> Any: - return batch - - -class RealESRGANBatchTransform(BatchTransform): - """ - It's too slow to process a batch of images under RealESRGAN degradation - model on CPU (by dataloader), which may cost 0.2 ~ 1 second per image. - So we execute the degradation process on GPU after loading a batch of images - and kernels from dataloader. - """ - def __init__( - self, - use_sharpener: bool, - resize_hq: bool, - queue_size: int, - resize_prob: Sequence[float], - resize_range: Sequence[float], - gray_noise_prob: float, - gaussian_noise_prob: float, - noise_range: Sequence[float], - poisson_scale_range: Sequence[float], - jpeg_range: Sequence[int], - second_blur_prob: float, - stage2_scale: Union[float, Sequence[Union[float, int]]], - resize_prob2: Sequence[float], - resize_range2: Sequence[float], - gray_noise_prob2: float, - gaussian_noise_prob2: float, - noise_range2: Sequence[float], - poisson_scale_range2: Sequence[float], - jpeg_range2: Sequence[int] - ) -> "RealESRGANBatchTransform": - super().__init__() - # resize settings for the first degradation process - self.resize_prob = resize_prob - self.resize_range = resize_range - - # noise settings for the first degradation process - self.gray_noise_prob = gray_noise_prob - self.gaussian_noise_prob = gaussian_noise_prob - self.noise_range = noise_range - self.poisson_scale_range = poisson_scale_range - self.jpeg_range = jpeg_range - - self.second_blur_prob = second_blur_prob - self.stage2_scale = stage2_scale - assert ( - isinstance(stage2_scale, (float, int)) or ( - isinstance(stage2_scale, Sequence) and len(stage2_scale) == 2 and - all(isinstance(x, (float, int)) for x in stage2_scale) - ) - ), f"stage2_scale can not be {type(stage2_scale)}" - - # resize settings for the second degradation process - self.resize_prob2 = resize_prob2 - self.resize_range2 = resize_range2 - - # noise settings for the second degradation process - self.gray_noise_prob2 = gray_noise_prob2 - self.gaussian_noise_prob2 = gaussian_noise_prob2 - self.noise_range2 = noise_range2 - self.poisson_scale_range2 = poisson_scale_range2 - self.jpeg_range2 = jpeg_range2 - - self.use_sharpener = use_sharpener - if self.use_sharpener: - self.usm_sharpener = USMSharp() - else: - self.usm_sharpener = None - self.resize_hq = resize_hq - self.queue_size = queue_size - self.jpeger = DiffJPEG(differentiable=False) - - @torch.no_grad() - def _dequeue_and_enqueue(self): - """It is the training pair pool for increasing the diversity in a batch. - - Batch processing limits the diversity of synthetic degradations in a batch. For example, samples in a - batch could not have different resize scaling factors. Therefore, we employ this training pair pool - to increase the degradation diversity in a batch. - """ - # initialize - b, c, h, w = self.lq.size() - if not hasattr(self, "queue_lr"): - # TODO: Being multiple of batch_size seems not necessary for queue_size - assert self.queue_size % b == 0, f"queue size {self.queue_size} should be divisible by batch size {b}" - self.queue_lr = torch.zeros(self.queue_size, c, h, w).to(self.lq) - _, c, h, w = self.gt.size() - self.queue_gt = torch.zeros(self.queue_size, c, h, w).to(self.lq) - self.queue_ptr = 0 - if self.queue_ptr == self.queue_size: # the pool is full - # do dequeue and enqueue - # shuffle - idx = torch.randperm(self.queue_size) - self.queue_lr = self.queue_lr[idx] - self.queue_gt = self.queue_gt[idx] - # get first b samples - lq_dequeue = self.queue_lr[0:b, :, :, :].clone() - gt_dequeue = self.queue_gt[0:b, :, :, :].clone() - # update the queue - self.queue_lr[0:b, :, :, :] = self.lq.clone() - self.queue_gt[0:b, :, :, :] = self.gt.clone() - - self.lq = lq_dequeue - self.gt = gt_dequeue - else: - # only do enqueue - self.queue_lr[self.queue_ptr:self.queue_ptr + b, :, :, :] = self.lq.clone() - self.queue_gt[self.queue_ptr:self.queue_ptr + b, :, :, :] = self.gt.clone() - self.queue_ptr = self.queue_ptr + b - - @torch.no_grad() - def __call__(self, batch: Dict[str, Union[torch.Tensor, str]]) -> Dict[str, Union[torch.Tensor, List[str]]]: - # training data synthesis - hq = batch["hq"] - if self.use_sharpener: - self.usm_sharpener.to(hq) - hq = self.usm_sharpener(hq) - self.jpeger.to(hq) - - kernel1 = batch["kernel1"] - kernel2 = batch["kernel2"] - sinc_kernel = batch["sinc_kernel"] - - ori_h, ori_w = hq.size()[2:4] - - # ----------------------- The first degradation process ----------------------- # - # blur - out = filter2D(hq, kernel1) - # random resize - updown_type = random.choices(["up", "down", "keep"], self.resize_prob)[0] - if updown_type == "up": - scale = np.random.uniform(1, self.resize_range[1]) - elif updown_type == "down": - scale = np.random.uniform(self.resize_range[0], 1) - else: - scale = 1 - mode = random.choice(["area", "bilinear", "bicubic"]) - out = F.interpolate(out, scale_factor=scale, mode=mode) - # add noise - if np.random.uniform() < self.gaussian_noise_prob: - out = random_add_gaussian_noise_pt( - out, sigma_range=self.noise_range, clip=True, - rounds=False, gray_prob=self.gray_noise_prob - ) - else: - out = random_add_poisson_noise_pt( - out, - scale_range=self.poisson_scale_range, - gray_prob=self.gray_noise_prob, - clip=True, - rounds=False - ) - # JPEG compression - jpeg_p = out.new_zeros(out.size(0)).uniform_(*self.jpeg_range) - # clamp to [0, 1], otherwise JPEGer will result in unpleasant artifacts - out = torch.clamp(out, 0, 1) - out = self.jpeger(out, quality=jpeg_p) - - # ----------------------- The second degradation process ----------------------- # - # blur - if np.random.uniform() < self.second_blur_prob: - out = filter2D(out, kernel2) - - # select scale of second degradation stage - if isinstance(self.stage2_scale, Sequence): - min_scale, max_scale = self.stage2_scale - stage2_scale = np.random.uniform(min_scale, max_scale) - else: - stage2_scale = self.stage2_scale - stage2_h, stage2_w = int(ori_h / stage2_scale), int(ori_w / stage2_scale) - # print(f"stage2 scale = {stage2_scale}") - - # random resize - updown_type = random.choices(["up", "down", "keep"], self.resize_prob2)[0] - if updown_type == "up": - scale = np.random.uniform(1, self.resize_range2[1]) - elif updown_type == "down": - scale = np.random.uniform(self.resize_range2[0], 1) - else: - scale = 1 - mode = random.choice(["area", "bilinear", "bicubic"]) - out = F.interpolate( - out, size=(int(stage2_h * scale), int(stage2_w * scale)), mode=mode - ) - # add noise - if np.random.uniform() < self.gaussian_noise_prob2: - out = random_add_gaussian_noise_pt( - out, sigma_range=self.noise_range2, clip=True, - rounds=False, gray_prob=self.gray_noise_prob2 - ) - else: - out = random_add_poisson_noise_pt( - out, - scale_range=self.poisson_scale_range2, - gray_prob=self.gray_noise_prob2, - clip=True, - rounds=False - ) - - # JPEG compression + the final sinc filter - # We also need to resize images to desired sizes. We group [resize back + sinc filter] together - # as one operation. - # We consider two orders: - # 1. [resize back + sinc filter] + JPEG compression - # 2. JPEG compression + [resize back + sinc filter] - # Empirically, we find other combinations (sinc + JPEG + Resize) will introduce twisted lines. - if np.random.uniform() < 0.5: - # resize back + the final sinc filter - mode = random.choice(["area", "bilinear", "bicubic"]) - out = F.interpolate(out, size=(stage2_h, stage2_w), mode=mode) - out = filter2D(out, sinc_kernel) - # JPEG compression - jpeg_p = out.new_zeros(out.size(0)).uniform_(*self.jpeg_range2) - out = torch.clamp(out, 0, 1) - out = self.jpeger(out, quality=jpeg_p) - else: - # JPEG compression - jpeg_p = out.new_zeros(out.size(0)).uniform_(*self.jpeg_range2) - out = torch.clamp(out, 0, 1) - out = self.jpeger(out, quality=jpeg_p) - # resize back + the final sinc filter - mode = random.choice(["area", "bilinear", "bicubic"]) - out = F.interpolate(out, size=(stage2_h, stage2_w), mode=mode) - out = filter2D(out, sinc_kernel) - - # resize back to gt_size since We are doing restoration task - if stage2_scale != 1: - out = F.interpolate(out, size=(ori_h, ori_w), mode="bicubic") - # clamp and round - lq = torch.clamp((out * 255.0).round(), 0, 255) / 255. - - if self.resize_hq and stage2_scale != 1: - # resize hq - hq = F.interpolate(hq, size=(stage2_h, stage2_w), mode="bicubic", antialias=True) - hq = F.interpolate(hq, size=(ori_h, ori_w), mode="bicubic", antialias=True) - - self.gt = hq - self.lq = lq - self._dequeue_and_enqueue() - - # [0, 1], float32, rgb, nhwc - lq = self.lq.float().permute(0, 2, 3, 1).contiguous() - # [-1, 1], float32, rgb, nhwc - hq = (self.gt * 2 - 1).float().permute(0, 2, 3, 1).contiguous() - - return dict(jpg=hq, hint=lq, txt=batch["txt"]) - - -class BicubicBatchTransform(BatchTransform): - """ - It's too slow to process a batch of images under RealESRGAN degradation - model on CPU (by dataloader), which may cost 0.2 ~ 1 second per image. - So we execute the degradation process on GPU after loading a batch of images - and kernels from dataloader. - """ - def __init__( - self, - use_sharpener: bool, - resize_hq: bool, - queue_size: int, - scale: Union[float, Sequence[Union[float, int]]], - ) -> "BicubicBatchTransform": - super().__init__() - - self.scale = scale - assert ( - isinstance(scale, (float, int)) or ( - isinstance(scale, Sequence) and len(scale) == 2 and - all(isinstance(x, (float, int)) for x in scale) - ) - ), f"scale can not be {type(scale)}" - - self.use_sharpener = use_sharpener - if self.use_sharpener: - self.usm_sharpener = USMSharp() - else: - self.usm_sharpener = None - self.resize_hq = resize_hq - self.queue_size = queue_size - - - @torch.no_grad() - def _dequeue_and_enqueue(self): - """It is the training pair pool for increasing the diversity in a batch. - - Batch processing limits the diversity of synthetic degradations in a batch. For example, samples in a - batch could not have different resize scaling factors. Therefore, we employ this training pair pool - to increase the degradation diversity in a batch. - """ - # initialize - b, c, h, w = self.lq.size() - if not hasattr(self, "queue_lr"): - # TODO: Being multiple of batch_size seems not necessary for queue_size - assert self.queue_size % b == 0, f"queue size {self.queue_size} should be divisible by batch size {b}" - self.queue_lr = torch.zeros(self.queue_size, c, h, w).to(self.lq) - _, c, h, w = self.gt.size() - self.queue_gt = torch.zeros(self.queue_size, c, h, w).to(self.lq) - self.queue_ptr = 0 - if self.queue_ptr == self.queue_size: # the pool is full - # do dequeue and enqueue - # shuffle - idx = torch.randperm(self.queue_size) - self.queue_lr = self.queue_lr[idx] - self.queue_gt = self.queue_gt[idx] - # get first b samples - lq_dequeue = self.queue_lr[0:b, :, :, :].clone() - gt_dequeue = self.queue_gt[0:b, :, :, :].clone() - # update the queue - self.queue_lr[0:b, :, :, :] = self.lq.clone() - self.queue_gt[0:b, :, :, :] = self.gt.clone() - - self.lq = lq_dequeue - self.gt = gt_dequeue - else: - # only do enqueue - self.queue_lr[self.queue_ptr:self.queue_ptr + b, :, :, :] = self.lq.clone() - self.queue_gt[self.queue_ptr:self.queue_ptr + b, :, :, :] = self.gt.clone() - self.queue_ptr = self.queue_ptr + b - - @torch.no_grad() - def __call__(self, batch: Dict[str, Union[torch.Tensor, str]]) -> Dict[str, Union[torch.Tensor, List[str]]]: - # training data synthesis - hq = batch["hq"] - if self.use_sharpener: - self.usm_sharpener.to(hq) - hq = self.usm_sharpener(hq) - - ori_h, ori_w = hq.size()[2:4] - h, w = int(ori_h / self.scale), int(ori_w / self.scale) - - # if self.resize_hq and self.scale != 1: - # # resize hq - # lq = F.interpolate(hq, size=(h, w), mode="bicubic", antialias=True) - - out = F.interpolate(hq, size=(h, w), mode="bicubic", antialias=True) - out = F.interpolate(out, size=(ori_h, ori_w), mode="bicubic", antialias=True) - - # clamp and round - lq = torch.clamp((out * 255.0).round(), 0, 255) / 255. - - self.gt = hq - self.lq = lq - self._dequeue_and_enqueue() - - # [0, 1], float32, rgb, nhwc - lq = self.lq.float().permute(0, 2, 3, 1).contiguous() - # [-1, 1], float32, rgb, nhwc - hq = (self.gt * 2 - 1).float().permute(0, 2, 3, 1).contiguous() - - return dict(jpg=hq, hint=lq, txt=batch["txt"]) \ No newline at end of file diff --git a/dataset/bicubic_torchvision.py b/dataset/bicubic_torchvision.py deleted file mode 100644 index c54f513..0000000 --- a/dataset/bicubic_torchvision.py +++ /dev/null @@ -1,107 +0,0 @@ -from typing import Dict, Sequence -import math -import random -import time - -import numpy as np -import torch -from torch.utils import data -from PIL import Image - -from utils.degradation import circular_lowpass_kernel, random_mixed_kernels -from utils.image import augment, random_crop_arr, center_crop_arr -from utils.file import load_file_list - - -class BicubicDataset(data.Dataset): - """ - # TODO: add comment - """ - - def __init__( - self, - file_list: str, - out_size: int, - crop_type: str, - use_hflip: bool, - use_rot: bool - ) -> "BicubicDataset": - super(BicubicDataset, self).__init__() - self.paths = load_file_list(file_list) - self.out_size = out_size - self.crop_type = crop_type - assert self.crop_type in ["center", "random", "none"], f"invalid crop type: {self.crop_type}" - - # self.blur_kernel_size = blur_kernel_size - # self.kernel_list = kernel_list - # # a list for each kernel probability - # self.kernel_prob = kernel_prob - # self.blur_sigma = blur_sigma - # # betag used in generalized Gaussian blur kernels - # self.betag_range = betag_range - # # betap used in plateau blur kernels - # self.betap_range = betap_range - # # the probability for sinc filters - # self.sinc_prob = sinc_prob - - # self.blur_kernel_size2 = blur_kernel_size2 - # self.kernel_list2 = kernel_list2 - # self.kernel_prob2 = kernel_prob2 - # self.blur_sigma2 = blur_sigma2 - # self.betag_range2 = betag_range2 - # self.betap_range2 = betap_range2 - # self.sinc_prob2 = sinc_prob2 - - # # a final sinc filter - # self.final_sinc_prob = final_sinc_prob - - self.use_hflip = use_hflip - self.use_rot = use_rot - - # kernel size ranges from 7 to 21 - # self.kernel_range = [2 * v + 1 for v in range(3, 11)] - # # TODO: kernel range is now hard-coded, should be in the configure file - # # convolving with pulse tensor brings no blurry effect - # self.pulse_tensor = torch.zeros(21, 21).float() - # self.pulse_tensor[10, 10] = 1 - - @torch.no_grad() - def __getitem__(self, index: int) -> Dict[str, torch.Tensor]: - # -------------------------------- Load hq images -------------------------------- # - hq_path = self.paths[index] - success = False - for _ in range(3): - try: - pil_img = Image.open(hq_path).convert("RGB") - success = True - break - except: - time.sleep(1) - assert success, f"failed to load image {hq_path}" - - if self.crop_type == "random": - pil_img = random_crop_arr(pil_img, self.out_size) - elif self.crop_type == "center": - pil_img = center_crop_arr(pil_img, self.out_size) - # self.crop_type is "none" - else: - pil_img = np.array(pil_img) - assert pil_img.shape[:2] == (self.out_size, self.out_size) - # hwc, rgb to bgr, [0, 255] to [0, 1], float32 - img_hq = (pil_img[..., ::-1] / 255.0).astype(np.float32) - - # -------------------- Do augmentation for training: flip, rotation -------------------- # - img_hq = augment(img_hq, self.use_hflip, self.use_rot) - - # [0, 1], BGR to RGB, HWC to CHW - img_hq = torch.from_numpy( - img_hq[..., ::-1].transpose(2, 0, 1).copy() - ).float() - - return { - "hq": img_hq, - 'txt': "" - } - - def __len__(self) -> int: - return len(self.paths) diff --git a/dataset/codeformer.py b/dataset/codeformer.py deleted file mode 100644 index f002d86..0000000 --- a/dataset/codeformer.py +++ /dev/null @@ -1,109 +0,0 @@ -from typing import Sequence, Dict, Union -import math -import time - -import numpy as np -import cv2 -from PIL import Image -import torch.utils.data as data - -from utils.file import load_file_list -from utils.image import center_crop_arr, augment, random_crop_arr -from utils.degradation import ( - random_mixed_kernels, random_add_gaussian_noise, random_add_jpg_compression -) - - -class CodeformerDataset(data.Dataset): - - def __init__( - self, - file_list: str, - out_size: int, - crop_type: str, - use_hflip: bool, - blur_kernel_size: int, - kernel_list: Sequence[str], - kernel_prob: Sequence[float], - blur_sigma: Sequence[float], - downsample_range: Sequence[float], - noise_range: Sequence[float], - jpeg_range: Sequence[int] - ) -> "CodeformerDataset": - super(CodeformerDataset, self).__init__() - self.file_list = file_list - self.paths = load_file_list(file_list) - self.out_size = out_size - self.crop_type = crop_type - assert self.crop_type in ["none", "center", "random"] - self.use_hflip = use_hflip - # degradation configurations - self.blur_kernel_size = blur_kernel_size - self.kernel_list = kernel_list - self.kernel_prob = kernel_prob - self.blur_sigma = blur_sigma - self.downsample_range = downsample_range - self.noise_range = noise_range - self.jpeg_range = jpeg_range - - def __getitem__(self, index: int) -> Dict[str, Union[np.ndarray, str]]: - # load gt image - # Shape: (h, w, c); channel order: BGR; image range: [0, 1], float32. - gt_path = self.paths[index] - success = False - for _ in range(3): - try: - pil_img = Image.open(gt_path).convert("RGB") - success = True - break - except: - time.sleep(1) - assert success, f"failed to load image {gt_path}" - - if self.crop_type == "center": - pil_img_gt = center_crop_arr(pil_img, self.out_size) - elif self.crop_type == "random": - pil_img_gt = random_crop_arr(pil_img, self.out_size) - else: - pil_img_gt = np.array(pil_img) - assert pil_img_gt.shape[:2] == (self.out_size, self.out_size) - img_gt = (pil_img_gt[..., ::-1] / 255.0).astype(np.float32) - - # random horizontal flip - img_gt = augment(img_gt, hflip=self.use_hflip, rotation=False, return_status=False) - h, w, _ = img_gt.shape - - # ------------------------ generate lq image ------------------------ # - # blur - kernel = random_mixed_kernels( - self.kernel_list, - self.kernel_prob, - self.blur_kernel_size, - self.blur_sigma, - self.blur_sigma, - [-math.pi, math.pi], - noise_range=None - ) - img_lq = cv2.filter2D(img_gt, -1, kernel) - # downsample - scale = np.random.uniform(self.downsample_range[0], self.downsample_range[1]) - img_lq = cv2.resize(img_lq, (int(w // scale), int(h // scale)), interpolation=cv2.INTER_LINEAR) - # noise - if self.noise_range is not None: - img_lq = random_add_gaussian_noise(img_lq, self.noise_range) - # jpeg compression - if self.jpeg_range is not None: - img_lq = random_add_jpg_compression(img_lq, self.jpeg_range) - - # resize to original size - img_lq = cv2.resize(img_lq, (w, h), interpolation=cv2.INTER_LINEAR) - - # BGR to RGB, [-1, 1] - target = (img_gt[..., ::-1] * 2 - 1).astype(np.float32) - # BGR to RGB, [0, 1] - source = img_lq[..., ::-1].astype(np.float32) - - return dict(jpg=target, txt="", hint=source) - - def __len__(self) -> int: - return len(self.paths) diff --git a/dataset/data_module.py b/dataset/data_module.py deleted file mode 100644 index d4dcde3..0000000 --- a/dataset/data_module.py +++ /dev/null @@ -1,68 +0,0 @@ -from typing import Any, Tuple, Mapping - -from pytorch_lightning.utilities.types import EVAL_DATALOADERS, TRAIN_DATALOADERS -import pytorch_lightning as pl -from torch.utils.data import DataLoader, Dataset -from omegaconf import OmegaConf - -from utils.common import instantiate_from_config -from dataset.batch_transform import BatchTransform, IdentityBatchTransform - - -class BIRDataModule(pl.LightningDataModule): - - def __init__( - self, - train_config: str, - val_config: str=None - ) -> "BIRDataModule": - super().__init__() - self.train_config = OmegaConf.load(train_config) - self.val_config = OmegaConf.load(val_config) if val_config else None - - def load_dataset(self, config: Mapping[str, Any]) -> Tuple[Dataset, BatchTransform]: - dataset = instantiate_from_config(config["dataset"]) - batch_transform = ( - instantiate_from_config(config["batch_transform"]) - if config.get("batch_transform") else IdentityBatchTransform() - ) - return dataset, batch_transform - - def setup(self, stage: str) -> None: - if stage == "fit": - self.train_dataset, self.train_batch_transform = self.load_dataset(self.train_config) - if self.val_config: - self.val_dataset, self.val_batch_transform = self.load_dataset(self.val_config) - else: - self.val_dataset, self.val_batch_transform = None, None - else: - raise NotImplementedError(stage) - - def train_dataloader(self) -> TRAIN_DATALOADERS: - return DataLoader( - dataset=self.train_dataset, **self.train_config["data_loader"] - ) - - def val_dataloader(self) -> EVAL_DATALOADERS: - if self.val_dataset is None: - return None - return DataLoader( - dataset=self.val_dataset, **self.val_config["data_loader"] - ) - - def on_after_batch_transfer(self, batch: Any, dataloader_idx: int) -> Any: - self.trainer: pl.Trainer - - if self.trainer.training: - return self.train_batch_transform(batch) - elif self.trainer.validating or self.trainer.sanity_checking: - return self.val_batch_transform(batch) - else: - raise RuntimeError( - "Trainer state: \n" - f"training: {self.trainer.training}\n" - f"validating: {self.trainer.validating}\n" - f"testing: {self.trainer.testing}\n" - f"predicting: {self.trainer.predicting}\n" - f"sanity_checking: {self.trainer.sanity_checking}" - ) diff --git a/dataset/projectdata.py b/dataset/projectdata.py deleted file mode 100644 index 3e139be..0000000 --- a/dataset/projectdata.py +++ /dev/null @@ -1,121 +0,0 @@ -from typing import Dict, Sequence -import math -import random -import time -import glob -import os -import cv2 - -import numpy as np -import torch -from torch.utils import data -from PIL import Image - -from utils.degradation import circular_lowpass_kernel, random_mixed_kernels -from utils.image import augment, random_crop_arr, center_crop_arr -from utils.file import load_file_list - - -class PROJECTDataset(data.Dataset): - """ - # TODO: add comment - """ - - def __init__(self, - file_path: list, - out_size: int, - scale: int, - hr_pattern: str, - blur_sigma: float, - blur_kernel: int, - crop_type: str, - use_hflip: bool, - use_rot: bool, - ): - super().__init__() - - self.patch_size = out_size // 4 - self.scale = scale - - lq_files = [] - hr_files = [] - for path in file_path: - lq_files += glob.glob(os.path.join(path, "*phone.*")) - hr_files += glob.glob(os.path.join(path, "*screen.*")) - - self.lq_files = lq_files - self.hr_files = hr_files - self.lq_files.sort() - self.hr_files.sort() - - self.hr_pattern = hr_pattern - self.blur_sigma = blur_sigma - self.blur_kernel = blur_kernel - self.noise_scale1 = 0.002 - self.noise_scale2 = 0.003 - self.blur_kernal_list = [ i for i in range(1, self.blur_kernel * 2, 2)] - self.blur_sigma_list = [ i / 10 for i in range(1, int(2 * self.blur_sigma * 10), 1)] - - self.edge_pixel = 0 - - @torch.no_grad() - def __len__(self): - return len(self.lq_files) - - def __getitem__(self, index): - lowRes_file = self.lq_files[index] - highRes_file = self.hr_files[index] - - hightRes_img = Image.open(highRes_file).convert("RGB") - hightRes_img = np.array(hightRes_img) - hightRes_img = (hightRes_img[..., ::-1] / 255.0).astype(np.float32) - lowRes_img = Image.open(lowRes_file).convert("RGB") - lowRes_img = np.array(lowRes_img) - lowRes_img = (lowRes_img[..., ::-1] / 255.0).astype(np.float32) - - low_h = lowRes_img.shape[0] - high_h = hightRes_img.shape[0] - if low_h == high_h: - lowRes_img = cv2.resize(lowRes_img, None, fx=1 / self.scale, fy=1 / self.scale) - - # scale为4数据集用作2x训练 - if high_h // low_h == 4 and self.scale == 2: - hightRes_img = cv2.resize(hightRes_img, None, fx=1 / 2, fy=1 / 2) - - # scale为2数据集用作4x训练 - if "20231021_plant_lightroom_Crop320" in lowRes_file and self.scale == 4: - lowRes_img = cv2.resize(lowRes_img, None, fx=1 / 2, fy=1 / 2) - - - img_size_ori_h, img_size_ori_w = lowRes_img.shape[:2] - i = random.randint(self.edge_pixel, img_size_ori_h - self.patch_size - self.edge_pixel) - j = random.randint(self.edge_pixel, img_size_ori_w - self.patch_size - self.edge_pixel) - hightRes_img = hightRes_img[i * self.scale : i * self.scale + self.patch_size * self.scale, j * self.scale : j * self.scale + self.patch_size * self.scale] - lowRes_img = lowRes_img[i : i + self.patch_size, j : j + self.patch_size] - - if random.uniform(0, 1) < 0.5: - hightRes_img = hightRes_img[::-1] - lowRes_img = lowRes_img[::-1] - if random.uniform(0, 1) < 0.5: - hightRes_img = hightRes_img[:,::-1] - lowRes_img = lowRes_img[:,::-1] - if 1: - sigma_x_offset = random.choice(self.blur_sigma_list) - sigma_y_offset = random.choice(self.blur_sigma_list) - kernal_x_offset = random.choice(self.blur_kernal_list) - kernal_y_offset = random.choice(self.blur_kernal_list) - lowRes_img = cv2.GaussianBlur(lowRes_img, (kernal_x_offset, kernal_y_offset), sigmaX=sigma_x_offset, sigmaY=sigma_y_offset) - # lowRes_img = lowRes_img * np.random.normal(loc=1, scale=self.noise_scale1, size=lowRes_img.shape) + np.random.normal(loc=0, scale=self.noise_scale2, size=lowRes_img.shape) - - hightRes_img = np.ascontiguousarray(hightRes_img) - lowRes_img = np.ascontiguousarray(lowRes_img) - - # lowRes_img = lowRes_img[:,:,None] - # hightRes_img = hightRes_img[:,:,None] - lowRes_img = lowRes_img.clip(0, 1) - hightRes_img = hightRes_img.clip(0, 1) - - lowRes_img = torch.from_numpy(np.ascontiguousarray(np.transpose(np.stack(lowRes_img, axis=0), (2, 0, 1)))).float() - hightRes_img = torch.from_numpy(np.ascontiguousarray(np.transpose(np.stack(hightRes_img, axis=0), (2, 0, 1)))).float() - - return {'lq': lowRes_img, 'hq': hightRes_img, 'txt': '', 'file_name': lowRes_file} \ No newline at end of file diff --git a/dataset/realesrgan.py b/dataset/realesrgan.py deleted file mode 100644 index 2a4a812..0000000 --- a/dataset/realesrgan.py +++ /dev/null @@ -1,183 +0,0 @@ -from typing import Dict, Sequence -import math -import random -import time - -import numpy as np -import torch -from torch.utils import data -from PIL import Image - -from utils.degradation import circular_lowpass_kernel, random_mixed_kernels -from utils.image import augment, random_crop_arr, center_crop_arr -from utils.file import load_file_list - - -class RealESRGANDataset(data.Dataset): - """ - # TODO: add comment - """ - - def __init__( - self, - file_list: str, - out_size: int, - crop_type: str, - use_hflip: bool, - use_rot: bool, - # blur kernel settings of the first degradation stage - blur_kernel_size: int, - kernel_list: Sequence[str], - kernel_prob: Sequence[float], - blur_sigma: Sequence[float], - betag_range: Sequence[float], - betap_range: Sequence[float], - sinc_prob: float, - # blur kernel settings of the second degradation stage - blur_kernel_size2: int, - kernel_list2: Sequence[str], - kernel_prob2: Sequence[float], - blur_sigma2: Sequence[float], - betag_range2: Sequence[float], - betap_range2: Sequence[float], - sinc_prob2: float, - final_sinc_prob: float - ) -> "RealESRGANDataset": - super(RealESRGANDataset, self).__init__() - self.paths = load_file_list(file_list) - self.out_size = out_size - self.crop_type = crop_type - assert self.crop_type in ["center", "random", "none"], f"invalid crop type: {self.crop_type}" - - self.blur_kernel_size = blur_kernel_size - self.kernel_list = kernel_list - # a list for each kernel probability - self.kernel_prob = kernel_prob - self.blur_sigma = blur_sigma - # betag used in generalized Gaussian blur kernels - self.betag_range = betag_range - # betap used in plateau blur kernels - self.betap_range = betap_range - # the probability for sinc filters - self.sinc_prob = sinc_prob - - self.blur_kernel_size2 = blur_kernel_size2 - self.kernel_list2 = kernel_list2 - self.kernel_prob2 = kernel_prob2 - self.blur_sigma2 = blur_sigma2 - self.betag_range2 = betag_range2 - self.betap_range2 = betap_range2 - self.sinc_prob2 = sinc_prob2 - - # a final sinc filter - self.final_sinc_prob = final_sinc_prob - - self.use_hflip = use_hflip - self.use_rot = use_rot - - # kernel size ranges from 7 to 21 - self.kernel_range = [2 * v + 1 for v in range(3, 11)] - # TODO: kernel range is now hard-coded, should be in the configure file - # convolving with pulse tensor brings no blurry effect - self.pulse_tensor = torch.zeros(21, 21).float() - self.pulse_tensor[10, 10] = 1 - - @torch.no_grad() - def __getitem__(self, index: int) -> Dict[str, torch.Tensor]: - # -------------------------------- Load hq images -------------------------------- # - hq_path = self.paths[index] - success = False - for _ in range(3): - try: - pil_img = Image.open(hq_path).convert("RGB") - success = True - break - except: - time.sleep(1) - assert success, f"failed to load image {hq_path}" - - if self.crop_type == "random": - pil_img = random_crop_arr(pil_img, self.out_size) - elif self.crop_type == "center": - pil_img = center_crop_arr(pil_img, self.out_size) - # self.crop_type is "none" - else: - pil_img = np.array(pil_img) - assert pil_img.shape[:2] == (self.out_size, self.out_size) - # hwc, rgb to bgr, [0, 255] to [0, 1], float32 - img_hq = (pil_img[..., ::-1] / 255.0).astype(np.float32) - - # -------------------- Do augmentation for training: flip, rotation -------------------- # - img_hq = augment(img_hq, self.use_hflip, self.use_rot) - - # ------------------------ Generate kernels (used in the first degradation) ------------------------ # - kernel_size = random.choice(self.kernel_range) - if np.random.uniform() < self.sinc_prob: - # this sinc filter setting is for kernels ranging from [7, 21] - if kernel_size < 13: - omega_c = np.random.uniform(np.pi / 3, np.pi) - else: - omega_c = np.random.uniform(np.pi / 5, np.pi) - kernel = circular_lowpass_kernel(omega_c, kernel_size, pad_to=False) - else: - kernel = random_mixed_kernels( - self.kernel_list, - self.kernel_prob, - kernel_size, - self.blur_sigma, - self.blur_sigma, [-math.pi, math.pi], - self.betag_range, - self.betap_range, - noise_range=None - ) - # pad kernel - pad_size = (21 - kernel_size) // 2 - kernel = np.pad(kernel, ((pad_size, pad_size), (pad_size, pad_size))) - - # ------------------------ Generate kernels (used in the second degradation) ------------------------ # - kernel_size = random.choice(self.kernel_range) - if np.random.uniform() < self.sinc_prob2: - if kernel_size < 13: - omega_c = np.random.uniform(np.pi / 3, np.pi) - else: - omega_c = np.random.uniform(np.pi / 5, np.pi) - kernel2 = circular_lowpass_kernel(omega_c, kernel_size, pad_to=False) - else: - kernel2 = random_mixed_kernels( - self.kernel_list2, - self.kernel_prob2, - kernel_size, - self.blur_sigma2, - self.blur_sigma2, [-math.pi, math.pi], - self.betag_range2, - self.betap_range2, - noise_range=None - ) - - # pad kernel - pad_size = (21 - kernel_size) // 2 - kernel2 = np.pad(kernel2, ((pad_size, pad_size), (pad_size, pad_size))) - - # ------------------------------------- the final sinc kernel ------------------------------------- # - if np.random.uniform() < self.final_sinc_prob: - kernel_size = random.choice(self.kernel_range) - omega_c = np.random.uniform(np.pi / 3, np.pi) - sinc_kernel = circular_lowpass_kernel(omega_c, kernel_size, pad_to=21) - sinc_kernel = torch.FloatTensor(sinc_kernel) - else: - sinc_kernel = self.pulse_tensor - - # [0, 1], BGR to RGB, HWC to CHW - img_hq = torch.from_numpy( - img_hq[..., ::-1].transpose(2, 0, 1).copy() - ).float() - kernel = torch.FloatTensor(kernel) - kernel2 = torch.FloatTensor(kernel2) - - return { - "hq": img_hq, "kernel1": kernel, "kernel2": kernel2, - "sinc_kernel": sinc_kernel, "txt": "" - } - - def __len__(self) -> int: - return len(self.paths) diff --git a/figs/bicubic.png b/figs/bicubic.png deleted file mode 100644 index 4db36f7..0000000 Binary files a/figs/bicubic.png and /dev/null differ diff --git a/figs/compare_1.png b/figs/compare_1.png deleted file mode 100644 index 95d1406..0000000 Binary files a/figs/compare_1.png and /dev/null differ diff --git a/figs/compare_2.png b/figs/compare_2.png deleted file mode 100644 index 67e89f7..0000000 Binary files a/figs/compare_2.png and /dev/null differ diff --git a/figs/compare_3.png b/figs/compare_3.png deleted file mode 100644 index 5ae0143..0000000 Binary files a/figs/compare_3.png and /dev/null differ diff --git a/figs/compare_4.png b/figs/compare_4.png deleted file mode 100644 index a7920f4..0000000 Binary files a/figs/compare_4.png and /dev/null differ diff --git a/figs/compare_5.png b/figs/compare_5.png deleted file mode 100644 index a64b71e..0000000 Binary files a/figs/compare_5.png and /dev/null differ diff --git a/figs/compare_6.png b/figs/compare_6.png deleted file mode 100644 index bcf806e..0000000 Binary files a/figs/compare_6.png and /dev/null differ diff --git a/figs/framework.png b/figs/framework.png deleted file mode 100644 index f85905e..0000000 Binary files a/figs/framework.png and /dev/null differ diff --git a/figs/logo.png b/figs/logo.png deleted file mode 100644 index aca7662..0000000 Binary files a/figs/logo.png and /dev/null differ diff --git a/figs/realworld.png b/figs/realworld.png deleted file mode 100644 index d6fd24b..0000000 Binary files a/figs/realworld.png and /dev/null differ diff --git a/figs/table.png b/figs/table.png deleted file mode 100644 index 220fd4c..0000000 Binary files a/figs/table.png and /dev/null differ diff --git a/inference_ccsr.py b/inference_ccsr.py deleted file mode 100644 index 05c0cdd..0000000 --- a/inference_ccsr.py +++ /dev/null @@ -1,230 +0,0 @@ -from typing import List, Tuple, Optional -import os -import math -from argparse import ArgumentParser, Namespace - -import numpy as np -import torch -import einops -from torch.nn import functional as F -import pytorch_lightning as pl -from PIL import Image -from omegaconf import OmegaConf - -from ldm.xformers_state import disable_xformers -from model.q_sampler import SpacedSampler -from model.ccsr_stage1 import ControlLDM -from model.cond_fn import MSEGuidance -from utils.image import auto_resize, pad -from utils.common import instantiate_from_config, load_state_dict -from utils.file import list_image_files, get_file_name_parts - - -@torch.no_grad() -def process( - model: ControlLDM, - control_imgs: List[np.ndarray], - steps: int, - t_max: float, - t_min: float, - strength: float, - color_fix_type: str, - cond_fn: Optional[MSEGuidance], - tiled: bool, - tile_size: int, - tile_stride: int -) -> Tuple[List[np.ndarray], List[np.ndarray]]: - """ - Apply CCSR model on a list of low-quality images. - - Args: - model (ControlLDM): Model. - control_imgs (List[np.ndarray]): A list of low-quality images (HWC, RGB, range in [0, 255]). - steps (int): Sampling steps. - t_max (float): - t_min (float): - strength (float): Control strength. Set to 1.0 during training. - color_fix_type (str): Type of color correction for samples. - cond_fn (Guidance | None): Guidance function that returns gradient to guide the predicted x_0. - tiled (bool): If specified, a patch-based sampling strategy will be used for sampling. - tile_size (int): Size of patch. - tile_stride (int): Stride of sliding patch. - - Returns: - preds (List[np.ndarray]): Restoration results (HWC, RGB, range in [0, 255]). - """ - n_samples = len(control_imgs) - sampler = SpacedSampler(model, var_type="fixed_small") - control = torch.tensor(np.stack(control_imgs) / 255.0, dtype=torch.float32, device=model.device).clamp_(0, 1) - control = einops.rearrange(control, "n h w c -> n c h w").contiguous() - - model.control_scales = [strength] * 13 - - if cond_fn is not None: - cond_fn.load_target(2 * control - 1) - - height, width = control.size(-2), control.size(-1) - shape = (n_samples, 4, height // 8, width // 8) - x_T = torch.randn(shape, device=model.device, dtype=torch.float32) - if not tiled: - # samples = sampler.sample_ccsr_stage1( - # steps=steps, t_max=t_max, shape=shape, cond_img=control, - # positive_prompt="", negative_prompt="", x_T=x_T, - # cfg_scale=1.0, cond_fn=cond_fn, - # color_fix_type=color_fix_type - # ) - samples = sampler.sample_ccsr( - steps=steps, t_max=t_max, t_min=t_min, shape=shape, cond_img=control, - positive_prompt="", negative_prompt="", x_T=x_T, - cfg_scale=1.0, cond_fn=cond_fn, - color_fix_type=color_fix_type - ) - else: - samples = sampler.sample_with_mixdiff_ccsr( - tile_size=tile_size, tile_stride=tile_stride, - steps=steps, t_max=t_max, t_min=t_min, shape=shape, cond_img=control, - positive_prompt="", negative_prompt="", x_T=x_T, - cfg_scale=1.0, cond_fn=cond_fn, - color_fix_type=color_fix_type - ) - - x_samples = samples.clamp(0, 1) - x_samples = (einops.rearrange(x_samples, "b c h w -> b h w c") * 255).cpu().numpy().clip(0, 255).astype(np.uint8) - - preds = [x_samples[i] for i in range(n_samples)] - - return preds - - -def parse_args() -> Namespace: - parser = ArgumentParser() - - parser.add_argument("--ckpt", type=str, help="full checkpoint path", - default='weights/real-world_ccsr.ckpt') - parser.add_argument("--config", type=str, help="model config path", default='configs/model/ccsr_stage2.yaml') - - parser.add_argument("--input", type=str, default='preset/test_datasets') - parser.add_argument("--steps", type=int, default=45) - parser.add_argument("--sr_scale", type=float, default=4) - parser.add_argument("--repeat_times", type=int, default=1) - - # patch-based sampling (tiling settings) - parser.add_argument("--tiled", action="store_true") - parser.add_argument("--tile_size", type=int, default=512) # image size - parser.add_argument("--tile_stride", type=int, default=256) # image size - - parser.add_argument("--color_fix_type", type=str, default="adain", choices=["wavelet", "adain", "none"]) - parser.add_argument("--output", type=str, default="experiments/test") - parser.add_argument("--t_max", type=float, default=0.6667) - parser.add_argument("--t_min", type=float, default=0.3333) - parser.add_argument("--show_lq", action="store_true") - parser.add_argument("--skip_if_exist", action="store_true") - - parser.add_argument("--seed", type=int, default=233) - parser.add_argument("--device", type=str, default="cuda", choices=["cpu", "cuda", "mps"]) - - return parser.parse_args() - - -def check_device(device): - if device == "cuda": - # check if CUDA is available - if not torch.cuda.is_available(): - print("CUDA not available because the current PyTorch install was not " - "built with CUDA enabled.") - device = "cpu" - else: - # xformers only support CUDA. Disable xformers when using cpu or mps. - disable_xformers() - if device == "mps": - # check if MPS is available - if not torch.backends.mps.is_available(): - if not torch.backends.mps.is_built(): - print("MPS not available because the current PyTorch install was not " - "built with MPS enabled.") - device = "cpu" - else: - print("MPS not available because the current MacOS version is not 12.3+ " - "and/or you do not have an MPS-enabled device on this machine.") - device = "cpu" - print(f'using device {device}') - return device - - -def main() -> None: - args = parse_args() - pl.seed_everything(args.seed) - - args.device = check_device(args.device) - - model: ControlLDM = instantiate_from_config(OmegaConf.load(args.config)) - load_state_dict(model, torch.load(args.ckpt, map_location="cpu"), strict=True) - # reload preprocess model if specified - - model.freeze() - model.to(args.device) - - assert os.path.isdir(args.input) - args.input_list = [args.input] - for file_path in list_image_files(args.input_list, follow_links=True): - lq = Image.open(file_path).convert("RGB") - if args.sr_scale != 1: - lq = lq.resize( - tuple(math.ceil(x * args.sr_scale) for x in lq.size), - Image.BICUBIC - ) - if not args.tiled: - lq_resized = auto_resize(lq, 512) - else: - lq_resized = auto_resize(lq, args.tile_size) - - x = lq_resized.resize( - tuple(s // 64 * 64 for s in lq_resized.size), Image.LANCZOS - ) - x = np.array(x) - # x = pad(np.array(lq_resized), scale=64) - - for i in range(args.repeat_times): - save_path = os.path.join(args.output, os.path.relpath(file_path, args.input)) - parent_path, stem, _ = get_file_name_parts(save_path) - save_path_now = os.path.join(parent_path, 'sample' + str(i)) - - save_path = os.path.join(save_path_now, f"{stem}.png") - if os.path.exists(save_path): - if args.skip_if_exist: - print(f"skip {save_path}") - continue - else: - raise RuntimeError(f"{save_path} already exist") - - os.makedirs(save_path_now, exist_ok=True) - - # initialize latent image guidance - cond_fn = None - - preds = process( - model, [x], steps=args.steps, - t_max=args.t_max, t_min=args.t_min, - strength=1, - color_fix_type=args.color_fix_type, - cond_fn=cond_fn, - tiled=args.tiled, tile_size=args.tile_size, tile_stride=args.tile_stride - ) - pred = preds[0] - # remove padding - # pred = pred[:lq_resized.height, :lq_resized.width, :] - - if args.show_lq: - pred = np.array(Image.fromarray(pred).resize(lq.size, Image.LANCZOS)) - stage1_pred = np.array(Image.fromarray(stage1_pred).resize(lq.size, Image.LANCZOS)) - lq = np.array(lq) - images = [lq, pred] - Image.fromarray(np.concatenate(images, axis=1)).save(save_path) - else: - Image.fromarray(pred).resize(lq.size, Image.LANCZOS).save(save_path) - # pred.save(save_path) - print(f"save to {save_path}") - - -if __name__ == "__main__": - main() diff --git a/ldm/__pycache__/util.cpython-310.pyc b/ldm/__pycache__/util.cpython-310.pyc index dd9a8cd..8fad02b 100644 Binary files a/ldm/__pycache__/util.cpython-310.pyc and b/ldm/__pycache__/util.cpython-310.pyc differ diff --git a/ldm/__pycache__/xformers_state.cpython-310.pyc b/ldm/__pycache__/xformers_state.cpython-310.pyc index 2a2b79b..ae435e8 100644 Binary files a/ldm/__pycache__/xformers_state.cpython-310.pyc and b/ldm/__pycache__/xformers_state.cpython-310.pyc differ diff --git a/ldm/models/__pycache__/autoencoder.cpython-310.pyc b/ldm/models/__pycache__/autoencoder.cpython-310.pyc index 8045f7e..191aaa9 100644 Binary files a/ldm/models/__pycache__/autoencoder.cpython-310.pyc and b/ldm/models/__pycache__/autoencoder.cpython-310.pyc differ diff --git a/ldm/models/diffusion/__pycache__/__init__.cpython-310.pyc b/ldm/models/diffusion/__pycache__/__init__.cpython-310.pyc index cef85ef..2a84304 100644 Binary files a/ldm/models/diffusion/__pycache__/__init__.cpython-310.pyc and b/ldm/models/diffusion/__pycache__/__init__.cpython-310.pyc differ diff --git a/ldm/models/diffusion/__pycache__/ddim.cpython-310.pyc b/ldm/models/diffusion/__pycache__/ddim.cpython-310.pyc index b5d2fce..0539b8b 100644 Binary files a/ldm/models/diffusion/__pycache__/ddim.cpython-310.pyc and b/ldm/models/diffusion/__pycache__/ddim.cpython-310.pyc differ diff --git a/ldm/models/diffusion/__pycache__/ddpm_ccsr_stage1.cpython-310.pyc b/ldm/models/diffusion/__pycache__/ddpm_ccsr_stage1.cpython-310.pyc index 8b02144..cd32ec5 100644 Binary files a/ldm/models/diffusion/__pycache__/ddpm_ccsr_stage1.cpython-310.pyc and b/ldm/models/diffusion/__pycache__/ddpm_ccsr_stage1.cpython-310.pyc differ diff --git a/ldm/models/diffusion/__pycache__/ddpm_ccsr_stage2.cpython-310.pyc b/ldm/models/diffusion/__pycache__/ddpm_ccsr_stage2.cpython-310.pyc index bddc9fb..a0cd72b 100644 Binary files a/ldm/models/diffusion/__pycache__/ddpm_ccsr_stage2.cpython-310.pyc and b/ldm/models/diffusion/__pycache__/ddpm_ccsr_stage2.cpython-310.pyc differ diff --git a/ldm/models/diffusion/ddpm_ccsr_stage1.py b/ldm/models/diffusion/ddpm_ccsr_stage1.py index 597faf3..6b30caa 100644 --- a/ldm/models/diffusion/ddpm_ccsr_stage1.py +++ b/ldm/models/diffusion/ddpm_ccsr_stage1.py @@ -17,16 +17,16 @@ from functools import partial import itertools from tqdm import tqdm from torchvision.utils import make_grid -from pytorch_lightning.utilities.distributed import rank_zero_only +from pytorch_lightning.utilities.rank_zero import rank_zero_only from omegaconf import ListConfig -from ldm.util import log_txt_as_img, exists, default, ismap, isimage, mean_flat, count_params, instantiate_from_config -from ldm.modules.ema import LitEma -from ldm.modules.distributions.distributions import normal_kl, DiagonalGaussianDistribution -from ldm.models.autoencoder import IdentityFirstStage, AutoencoderKL -from ldm.modules.diffusionmodules.util import make_beta_schedule, extract_into_tensor, noise_like -from ldm.models.diffusion.ddim import DDIMSampler -from model.mixins import ImageLoggerMixin +from ....ldm.util import log_txt_as_img, exists, default, ismap, isimage, mean_flat, count_params, instantiate_from_config +from ....ldm.modules.ema import LitEma +from ....ldm.modules.distributions.distributions import normal_kl, DiagonalGaussianDistribution +from ....ldm.models.autoencoder import IdentityFirstStage, AutoencoderKL +from ....ldm.modules.diffusionmodules.util import make_beta_schedule, extract_into_tensor, noise_like +from ....ldm.models.diffusion.ddim import DDIMSampler +from ....model.mixins import ImageLoggerMixin __conditioning_keys__ = {'concat': 'c_concat', diff --git a/ldm/models/diffusion/ddpm_ccsr_stage2.py b/ldm/models/diffusion/ddpm_ccsr_stage2.py index aa766f6..9a27ba2 100644 --- a/ldm/models/diffusion/ddpm_ccsr_stage2.py +++ b/ldm/models/diffusion/ddpm_ccsr_stage2.py @@ -17,18 +17,18 @@ from functools import partial import itertools from tqdm import tqdm from torchvision.utils import make_grid -from pytorch_lightning.utilities.distributed import rank_zero_only +from pytorch_lightning.utilities.rank_zero import rank_zero_only from omegaconf import ListConfig -from ldm.util import log_txt_as_img, exists, default, ismap, isimage, mean_flat, count_params, instantiate_from_config -from ldm.modules.ema import LitEma -from ldm.modules.distributions.distributions import normal_kl, DiagonalGaussianDistribution -from ldm.models.autoencoder import IdentityFirstStage, AutoencoderKL -from ldm.modules.diffusionmodules.util import make_beta_schedule, extract_into_tensor, noise_like -from ldm.models.diffusion.ddim import DDIMSampler -from model.mixins import ImageLoggerMixin +from ....ldm.util import log_txt_as_img, exists, default, ismap, isimage, mean_flat, count_params, instantiate_from_config +from ....ldm.modules.ema import LitEma +from ....ldm.modules.distributions.distributions import normal_kl, DiagonalGaussianDistribution +from ....ldm.models.autoencoder import IdentityFirstStage, AutoencoderKL +from ....ldm.modules.diffusionmodules.util import make_beta_schedule, extract_into_tensor, noise_like +from ....ldm.models.diffusion.ddim import DDIMSampler +from ....model.mixins import ImageLoggerMixin -from model.q_sampler import space_timesteps +from ....model.q_sampler import space_timesteps __conditioning_keys__ = {'concat': 'c_concat', diff --git a/ldm/modules/__pycache__/attention.cpython-310.pyc b/ldm/modules/__pycache__/attention.cpython-310.pyc index 6bdf92a..871b724 100644 Binary files a/ldm/modules/__pycache__/attention.cpython-310.pyc and b/ldm/modules/__pycache__/attention.cpython-310.pyc differ diff --git a/ldm/modules/__pycache__/ema.cpython-310.pyc b/ldm/modules/__pycache__/ema.cpython-310.pyc index df87120..43f43a8 100644 Binary files a/ldm/modules/__pycache__/ema.cpython-310.pyc and b/ldm/modules/__pycache__/ema.cpython-310.pyc differ diff --git a/ldm/modules/attention.py b/ldm/modules/attention.py index ac45171..f72de57 100644 --- a/ldm/modules/attention.py +++ b/ldm/modules/attention.py @@ -6,8 +6,8 @@ from torch import nn, einsum from einops import rearrange, repeat from typing import Optional, Any -from ldm.modules.diffusionmodules.util import checkpoint -from ldm import xformers_state +from ...ldm.modules.diffusionmodules.util import checkpoint +from ...ldm import xformers_state # try: # import xformers diff --git a/ldm/modules/diffusionmodules/__pycache__/__init__.cpython-310.pyc b/ldm/modules/diffusionmodules/__pycache__/__init__.cpython-310.pyc index 3efbede..f68359c 100644 Binary files a/ldm/modules/diffusionmodules/__pycache__/__init__.cpython-310.pyc and b/ldm/modules/diffusionmodules/__pycache__/__init__.cpython-310.pyc differ diff --git a/ldm/modules/diffusionmodules/__pycache__/openaimodel.cpython-310.pyc b/ldm/modules/diffusionmodules/__pycache__/openaimodel.cpython-310.pyc index e61d899..1ce5f7a 100644 Binary files a/ldm/modules/diffusionmodules/__pycache__/openaimodel.cpython-310.pyc and b/ldm/modules/diffusionmodules/__pycache__/openaimodel.cpython-310.pyc differ diff --git a/ldm/modules/diffusionmodules/__pycache__/util.cpython-310.pyc b/ldm/modules/diffusionmodules/__pycache__/util.cpython-310.pyc index 6c68efe..61b1af1 100644 Binary files a/ldm/modules/diffusionmodules/__pycache__/util.cpython-310.pyc and b/ldm/modules/diffusionmodules/__pycache__/util.cpython-310.pyc differ diff --git a/ldm/modules/diffusionmodules/openaimodel.py b/ldm/modules/diffusionmodules/openaimodel.py index 7df6b5a..ef58639 100644 --- a/ldm/modules/diffusionmodules/openaimodel.py +++ b/ldm/modules/diffusionmodules/openaimodel.py @@ -6,7 +6,7 @@ import torch as th import torch.nn as nn import torch.nn.functional as F -from ldm.modules.diffusionmodules.util import ( +from ....ldm.modules.diffusionmodules.util import ( checkpoint, conv_nd, linear, @@ -15,8 +15,8 @@ from ldm.modules.diffusionmodules.util import ( normalization, timestep_embedding, ) -from ldm.modules.attention import SpatialTransformer -from ldm.util import exists +from ....ldm.modules.attention import SpatialTransformer +from ....ldm.util import exists # dummy replace diff --git a/ldm/modules/distributions/__pycache__/__init__.cpython-310.pyc b/ldm/modules/distributions/__pycache__/__init__.cpython-310.pyc index 12659b1..7c8924d 100644 Binary files a/ldm/modules/distributions/__pycache__/__init__.cpython-310.pyc and b/ldm/modules/distributions/__pycache__/__init__.cpython-310.pyc differ diff --git a/ldm/modules/distributions/__pycache__/distributions.cpython-310.pyc b/ldm/modules/distributions/__pycache__/distributions.cpython-310.pyc index a14e79d..a5317e1 100644 Binary files a/ldm/modules/distributions/__pycache__/distributions.cpython-310.pyc and b/ldm/modules/distributions/__pycache__/distributions.cpython-310.pyc differ diff --git a/ldm/modules/encoders/__pycache__/__init__.cpython-310.pyc b/ldm/modules/encoders/__pycache__/__init__.cpython-310.pyc index ccdc342..c11fe71 100644 Binary files a/ldm/modules/encoders/__pycache__/__init__.cpython-310.pyc and b/ldm/modules/encoders/__pycache__/__init__.cpython-310.pyc differ diff --git a/ldm/modules/encoders/__pycache__/modules.cpython-310.pyc b/ldm/modules/encoders/__pycache__/modules.cpython-310.pyc index bc3f3b9..2214dc4 100644 Binary files a/ldm/modules/encoders/__pycache__/modules.cpython-310.pyc and b/ldm/modules/encoders/__pycache__/modules.cpython-310.pyc differ diff --git a/ldm/modules/losses/__init__.py b/ldm/modules/losses/__init__.py index 876d7c5..b8b8ba7 100644 --- a/ldm/modules/losses/__init__.py +++ b/ldm/modules/losses/__init__.py @@ -1 +1 @@ -from ldm.modules.losses.contperceptual import LPIPSWithDiscriminator \ No newline at end of file +from ....ldm.modules.losses.contperceptual import LPIPSWithDiscriminator \ No newline at end of file diff --git a/ldm/modules/losses/__pycache__/__init__.cpython-310.pyc b/ldm/modules/losses/__pycache__/__init__.cpython-310.pyc index 0cbd09b..2c05b4f 100644 Binary files a/ldm/modules/losses/__pycache__/__init__.cpython-310.pyc and b/ldm/modules/losses/__pycache__/__init__.cpython-310.pyc differ diff --git a/ldm/modules/losses/__pycache__/contperceptual.cpython-310.pyc b/ldm/modules/losses/__pycache__/contperceptual.cpython-310.pyc index ffdaaff..093d104 100644 Binary files a/ldm/modules/losses/__pycache__/contperceptual.cpython-310.pyc and b/ldm/modules/losses/__pycache__/contperceptual.cpython-310.pyc differ diff --git a/model/__pycache__/ccsr_stage1.cpython-310.pyc b/model/__pycache__/ccsr_stage1.cpython-310.pyc index 1409dfb..856565e 100644 Binary files a/model/__pycache__/ccsr_stage1.cpython-310.pyc and b/model/__pycache__/ccsr_stage1.cpython-310.pyc differ diff --git a/model/__pycache__/ccsr_stage2.cpython-310.pyc b/model/__pycache__/ccsr_stage2.cpython-310.pyc index 8e796a4..ea9f9a3 100644 Binary files a/model/__pycache__/ccsr_stage2.cpython-310.pyc and b/model/__pycache__/ccsr_stage2.cpython-310.pyc differ diff --git a/model/__pycache__/cond_fn.cpython-310.pyc b/model/__pycache__/cond_fn.cpython-310.pyc index 284b18e..55d89bd 100644 Binary files a/model/__pycache__/cond_fn.cpython-310.pyc and b/model/__pycache__/cond_fn.cpython-310.pyc differ diff --git a/model/__pycache__/mixins.cpython-310.pyc b/model/__pycache__/mixins.cpython-310.pyc index a0f5286..3517afe 100644 Binary files a/model/__pycache__/mixins.cpython-310.pyc and b/model/__pycache__/mixins.cpython-310.pyc differ diff --git a/model/__pycache__/q_sampler.cpython-310.pyc b/model/__pycache__/q_sampler.cpython-310.pyc index 8677eff..eda7029 100644 Binary files a/model/__pycache__/q_sampler.cpython-310.pyc and b/model/__pycache__/q_sampler.cpython-310.pyc differ diff --git a/model/__pycache__/spaced_sampler.cpython-310.pyc b/model/__pycache__/spaced_sampler.cpython-310.pyc index 523e4d0..be68873 100644 Binary files a/model/__pycache__/spaced_sampler.cpython-310.pyc and b/model/__pycache__/spaced_sampler.cpython-310.pyc differ diff --git a/model/ccsr_stage1.py b/model/ccsr_stage1.py index a19e97e..24a9b8c 100644 --- a/model/ccsr_stage1.py +++ b/model/ccsr_stage1.py @@ -7,18 +7,18 @@ import torch import torch as th import torch.nn as nn -from ldm.modules.diffusionmodules.util import ( +from ..ldm.modules.diffusionmodules.util import ( conv_nd, linear, zero_module, timestep_embedding, ) -from ldm.modules.attention import SpatialTransformer -from ldm.modules.diffusionmodules.openaimodel import TimestepEmbedSequential, ResBlock, Downsample, AttentionBlock, UNetModel -from ldm.models.diffusion.ddpm_ccsr_stage1 import LatentDiffusion -from ldm.util import log_txt_as_img, exists, instantiate_from_config -from ldm.modules.distributions.distributions import DiagonalGaussianDistribution -from utils.common import frozen_module +from ..ldm.modules.attention import SpatialTransformer +from ..ldm.modules.diffusionmodules.openaimodel import TimestepEmbedSequential, ResBlock, Downsample, AttentionBlock, UNetModel +from ..ldm.models.diffusion.ddpm_ccsr_stage1 import LatentDiffusion +from ..ldm.util import log_txt_as_img, exists, instantiate_from_config +from ..ldm.modules.distributions.distributions import DiagonalGaussianDistribution +from ..utils.common import frozen_module from .spaced_sampler import SpacedSampler class ControlledUnetModel(UNetModel): diff --git a/model/ccsr_stage2.py b/model/ccsr_stage2.py index 9beaf31..986bcba 100644 --- a/model/ccsr_stage2.py +++ b/model/ccsr_stage2.py @@ -7,18 +7,18 @@ import torch import torch as th import torch.nn as nn -from ldm.modules.diffusionmodules.util import ( +from ..ldm.modules.diffusionmodules.util import ( conv_nd, linear, zero_module, timestep_embedding, ) -from ldm.modules.attention import SpatialTransformer -from ldm.modules.diffusionmodules.openaimodel import TimestepEmbedSequential, ResBlock, Downsample, AttentionBlock, UNetModel -from ldm.models.diffusion.ddpm_ccsr_stage2 import LatentDiffusion -from ldm.util import log_txt_as_img, exists, instantiate_from_config -from ldm.modules.distributions.distributions import DiagonalGaussianDistribution -from utils.common import frozen_module +from ..ldm.modules.attention import SpatialTransformer +from ..ldm.modules.diffusionmodules.openaimodel import TimestepEmbedSequential, ResBlock, Downsample, AttentionBlock, UNetModel +from ..ldm.models.diffusion.ddpm_ccsr_stage2 import LatentDiffusion +from ..ldm.util import log_txt_as_img, exists, instantiate_from_config +from ..ldm.modules.distributions.distributions import DiagonalGaussianDistribution +from ..utils.common import frozen_module from .spaced_sampler import SpacedSampler class ControlledUnetModel(UNetModel): diff --git a/model/q_sampler.py b/model/q_sampler.py index 8f0c9dc..8a78c5b 100644 --- a/model/q_sampler.py +++ b/model/q_sampler.py @@ -7,9 +7,9 @@ import einops import os from PIL import Image -from ldm.modules.diffusionmodules.util import make_beta_schedule -from model.cond_fn import Guidance -from utils.image import ( +from ..ldm.modules.diffusionmodules.util import make_beta_schedule +from ..model.cond_fn import Guidance +from ..utils.image import ( wavelet_reconstruction, adaptive_instance_normalization ) diff --git a/model/spaced_sampler.py b/model/spaced_sampler.py index fca59b9..1d6d766 100644 --- a/model/spaced_sampler.py +++ b/model/spaced_sampler.py @@ -7,11 +7,12 @@ import einops import os from PIL import Image -from ldm.modules.diffusionmodules.util import make_beta_schedule -from model.cond_fn import Guidance -from utils.image import ( +from ..ldm.modules.diffusionmodules.util import make_beta_schedule +from ..model.cond_fn import Guidance +from ..utils.image import ( wavelet_reconstruction, adaptive_instance_normalization ) +import comfy.utils # https://github.com/openai/guided-diffusion/blob/main/guided_diffusion/respace.py def space_timesteps(num_timesteps, section_counts): @@ -470,6 +471,8 @@ class SpacedSampler: total_steps = len(self.timesteps) iterator = tqdm(time_range, desc="Spaced Sampler", total=total_steps) + pbar = comfy.utils.ProgressBar(total_steps) + # sampling loop for i, step in enumerate(iterator): ts = torch.full((b,), step, device=device, dtype=torch.long) @@ -520,6 +523,7 @@ class SpacedSampler: noise_buffer.zero_() count.zero_() + pbar.update(1) # decode samples of each diffusion process img_buffer = torch.zeros_like(cond_img) diff --git a/nodes.py b/nodes.py new file mode 100644 index 0000000..ca62625 --- /dev/null +++ b/nodes.py @@ -0,0 +1,155 @@ +import os +import sys +import torch +import numpy as np + +import numpy as np +import torch +import einops +from torch.nn import functional as F +import pytorch_lightning as pl +from PIL import Image +from omegaconf import OmegaConf + +from .model.q_sampler import SpacedSampler +from .model.ccsr_stage1 import ControlLDM +from .model.cond_fn import MSEGuidance + +from .utils.common import instantiate_from_config, load_state_dict + +script_directory = os.path.dirname(os.path.abspath(__file__)) +project_directory = os.path.join(script_directory, '..') # Adjust the path as necessary +sys.path.insert(0, project_directory) + +class CCSR_Upscale: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "image": ("IMAGE", ), + #"sr_scale": ("INT", {"default": 4, "min": 1, "max": 12, "step": 1}), + "steps": ("INT", {"default": 10, "min": 1, "max": 4096, "step": 1}), + "t_max": ("FLOAT", {"default": 0.6667,"min": 0, "max": 1, "step": 0.01}), + "t_min": ("FLOAT", {"default": 0.3333,"min": 0, "max": 1, "step": 0.01}), + "tile_size": ("INT", {"default": 512, "min": 1, "max": 4096, "step": 1}), + "tile_stride": ("INT", {"default": 256, "min": 1, "max": 4096, "step": 1}), + "tiled": ("BOOLEAN", {"default": False}), + "color_fix_type": ( + [ + 'none', + 'adain', + 'wavelet', + ], { + "default": 'adain' + }), + + }, + + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES =("upscaled_image",) + FUNCTION = "process" + + CATEGORY = "CCSR" + + @torch.no_grad() + + + + def process(self, image, steps, t_max, t_min, tiled,tile_size, tile_stride, color_fix_type): + """ + Apply CCSR model on a list of low-quality images. + + Args: + model (ControlLDM): Model. + control_imgs (List[np.ndarray]): A list of low-quality images (HWC, RGB, range in [0, 255]). + steps (int): Sampling steps. + t_max (float): + t_min (float): + strength (float): Control strength. Set to 1.0 during training. + color_fix_type (str): Type of color correction for samples. + cond_fn (Guidance | None): Guidance function that returns gradient to guide the predicted x_0. + tiled (bool): If specified, a patch-based sampling strategy will be used for sampling. + tile_size (int): Size of patch. + tile_stride (int): Stride of sliding patch. + + Returns: + preds (List[np.ndarray]): Restoration results (HWC, RGB, range in [0, 255]). + """ + checkpoint_path = os.path.join(script_directory, "../../models/checkpoints/real-world_ccsr.ckpt") + config_path = os.path.join(script_directory, "configs/model/ccsr_stage2.yaml") + + config = OmegaConf.load(config_path) + model = instantiate_from_config(config) + + + load_state_dict(model, torch.load(checkpoint_path, map_location="cpu"), strict=True) + # reload preprocess model if specified + + model.freeze() + model.to("cuda") + + sampler = SpacedSampler(model, var_type="fixed_small") + + # Assuming 'image' is a PyTorch tensor with shape [B, H, W, C] and you want to resize it. + B, H, W, C = image.shape + + # Calculate the new height and width, rounding down to the nearest multiple of 64. + new_height = H // 64 * 64 + new_width = W // 64 * 64 + + # Reorder to [B, C, H, W] before using interpolate. + image = image.permute(0, 3, 1, 2).contiguous() + + # Resize the image tensor. + resized_image = F.interpolate(image, size=(new_height, new_width), mode='bicubic', align_corners=False) + + # Move the tensor to the GPU. + resized_image = resized_image.to("cuda") + + + strength = 1.0 + model.control_scales = [strength] * 13 + + cond_fn = None + + height, width = resized_image.size(-2), resized_image.size(-1) + shape = (1, 4, height // 8, width // 8) + x_T = torch.randn(shape, device=model.device, dtype=torch.float32) + print(resized_image.shape) + print(x_T.shape) + if not tiled: + # samples = sampler.sample_ccsr_stage1( + # steps=steps, t_max=t_max, shape=shape, cond_img=control, + # positive_prompt="", negative_prompt="", x_T=x_T, + # cfg_scale=1.0, cond_fn=cond_fn, + # color_fix_type=color_fix_type + # ) + samples = sampler.sample_ccsr( + steps=steps, t_max=t_max, t_min=t_min, shape=shape, cond_img=resized_image, + positive_prompt="", negative_prompt="", x_T=x_T, + cfg_scale=1.0, cond_fn=cond_fn, + color_fix_type=color_fix_type + ) + else: + samples = sampler.sample_with_mixdiff_ccsr( + tile_size=tile_size, tile_stride=tile_stride, + steps=steps, t_max=t_max, t_min=t_min, shape=shape, cond_img=resized_image, + positive_prompt="", negative_prompt="", x_T=x_T, + cfg_scale=1.0, cond_fn=cond_fn, + color_fix_type=color_fix_type + ) + + x_samples = samples.clamp(0, 1) + #x_samples = (einops.rearrange(x_samples, "b c h w -> b h w c") * 255).cpu().numpy().clip(0, 255).astype(np.uint8) + x_samples = (einops.rearrange(x_samples, "b c h w -> b h w c")).cpu() + print(x_samples.shape) + return (x_samples,) + + +NODE_CLASS_MAPPINGS = { + "CCSR_Upscale": CCSR_Upscale, +} +NODE_DISPLAY_NAME_MAPPINGS = { + "CCSR_Upscale": "CCSR_Upscale", +} \ No newline at end of file diff --git a/preset/test_datasets/19.jpg b/preset/test_datasets/19.jpg deleted file mode 100644 index 69a5675..0000000 Binary files a/preset/test_datasets/19.jpg and /dev/null differ diff --git a/preset/test_datasets/49.jpg b/preset/test_datasets/49.jpg deleted file mode 100644 index de46b72..0000000 Binary files a/preset/test_datasets/49.jpg and /dev/null differ diff --git a/readme.md b/readme.md new file mode 100644 index 0000000..0399469 --- /dev/null +++ b/readme.md @@ -0,0 +1,9 @@ +# ComfyUI- CCSR upscaler node + +This is a simple wrapper node for https://github.com/csslc/CCSR + +NOT a proper ComfyUI implementation, so not very efficient and there might be memory issues, tested on 4090 and 4x upscale tiled worked well. + +Upscale the input first with another node for the desired end scale. + +The model (https://drive.google.com/drive/folders/1jM1mxDryPk9CTuFTvYcraP2XIVzbPiw_?usp=drive_link) goes to ComfyUI/models/checkpoints \ No newline at end of file diff --git a/requirements.txt b/requirements.txt index 1871195..d06356d 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,19 +1,3 @@ -torch==2.0.1 -torchvision -xformers -pytorch_lightning -einops -open-clip-torch +taming-transformers omegaconf -torchmetrics -triton==2.0.0 -opencv-python-headless -scipy -matplotlib -lpips -gradio -chardet -transformers -facexlib -pyiqa -pydantic==1.10.11 \ No newline at end of file +einops \ No newline at end of file diff --git a/scripts/make_file_list.py b/scripts/make_file_list.py deleted file mode 100644 index df4ebb8..0000000 --- a/scripts/make_file_list.py +++ /dev/null @@ -1,36 +0,0 @@ -import sys -sys.path.append(".") -import os -from argparse import ArgumentParser - -from utils.file import list_image_files - - -parser = ArgumentParser() -# parser.add_argument("--img_folder", type=str, required=True) -parser.add_argument("--img_folder", nargs="+", required=True) -parser.add_argument("--val_size", type=int, required=True) -parser.add_argument("--save_folder", type=str, required=True) -parser.add_argument("--follow_links", action="store_true") -args = parser.parse_args() - -files = list_image_files( - args.img_folder, exts=(".jpg", ".png", ".jpeg"), follow_links=args.follow_links, - log_progress=True, log_every_n_files=10000 -) - -print(f"find {len(files)} images in {args.img_folder}") -assert args.val_size < len(files) - -val_files = files[:args.val_size] -train_files = files[args.val_size:] - -os.makedirs(args.save_folder, exist_ok=True) - -with open(os.path.join(args.save_folder, "train.list"), "w") as fp: - for file_path in train_files: - fp.write(f"{file_path}\n") - -with open(os.path.join(args.save_folder, "val.list"), "w") as fp: - for file_path in val_files: - fp.write(f"{file_path}\n") diff --git a/scripts/make_stage2_init_weight.py b/scripts/make_stage2_init_weight.py deleted file mode 100644 index d8f57a6..0000000 --- a/scripts/make_stage2_init_weight.py +++ /dev/null @@ -1,80 +0,0 @@ -import sys -sys.path.append(".") -from argparse import ArgumentParser -from typing import Dict - -import torch -from omegaconf import OmegaConf - -from utils.common import instantiate_from_config - - -def load_weight(weight_path: str) -> Dict[str, torch.Tensor]: - weight = torch.load(weight_path) - if "state_dict" in weight: - weight = weight["state_dict"] - - pure_weight = {} - for key, val in weight.items(): - if key.startswith("module."): - key = key[len("module."):] - pure_weight[key] = val - - return pure_weight - -parser = ArgumentParser() -parser.add_argument("--cldm_config", type=str, default = 'configs/model/ccsr_stage1.yaml') -parser.add_argument("--sd_weight", type=str, default='****/v2-1_512-ema-pruned.ckpt') -parser.add_argument("--swinir_weight", type=str, required=False) -parser.add_argument("--output", type=str, default='weights/init_weight_ccsr.ckpt') -args = parser.parse_args() - -model = instantiate_from_config(OmegaConf.load(args.cldm_config)) - -sd_weights = load_weight(args.sd_weight) -# swinir_weights = load_weight(args.swinir_weight) -scratch_weights = model.state_dict() - -init_weights = {} -for weight_name in scratch_weights.keys(): - # find target pretrained weights for this weight - if weight_name.startswith("control_"): - suffix = weight_name[len("control_"):] - target_name = f"model.diffusion_{suffix}" - target_model_weights = sd_weights - # elif weight_name.startswith("preprocess_model."): - # suffix = weight_name[len("preprocess_model."):] - # target_name = suffix - # target_model_weights = swinir_weights - elif weight_name.startswith("cond_encoder."): - suffix = weight_name[len("cond_encoder."):] - target_name = F"first_stage_model.{suffix}" - target_model_weights = sd_weights - else: - target_name = weight_name - target_model_weights = sd_weights - - # if target weight exist in pretrained model - print(f"copy weights: {target_name} -> {weight_name}") - if target_name in target_model_weights: - # get pretrained weight - target_weight = target_model_weights[target_name] - target_shape = target_weight.shape - model_shape = scratch_weights[weight_name].shape - # if pretrained weight has the same shape with model weight, we make a copy - if model_shape == target_shape: - init_weights[weight_name] = target_weight.clone() - # else we copy pretrained weight with additional channels initialized to zero - else: - newly_added_channels = model_shape[1] - target_shape[1] - oc, _, h, w = target_shape - zero_weight = torch.zeros((oc, newly_added_channels, h, w)).type_as(target_weight) - init_weights[weight_name] = torch.cat((target_weight.clone(), zero_weight), dim=1) - print(f"add zero weight to {target_name} in pretrained weights, newly added channels = {newly_added_channels}") - else: - init_weights[weight_name] = scratch_weights[weight_name].clone() - print(f"These weights are newly added: {weight_name}") - -model.load_state_dict(init_weights, strict=True) -torch.save(model.state_dict(), args.output) -print("Done.") diff --git a/scripts/test.sh b/scripts/test.sh deleted file mode 100644 index e69de29..0000000 diff --git a/sh/1_pip.sh b/sh/1_pip.sh deleted file mode 100644 index 11de99b..0000000 --- a/sh/1_pip.sh +++ /dev/null @@ -1,3 +0,0 @@ -cd CCSR -pip install -r requirements.txt -pip install -e git+https://github.com/CompVis/taming-transformers.git@master#egg=taming-transformers \ No newline at end of file diff --git a/sh/2_inference.sh b/sh/2_inference.sh deleted file mode 100644 index 18e735a..0000000 --- a/sh/2_inference.sh +++ /dev/null @@ -1,14 +0,0 @@ - - -python inference_ccsr.py \ ---input preset/test_datasets \ ---config configs/model/ccsr_stage2.yaml \ ---ckpt weights/real-world_ccsr.ckpt \ ---steps 45 \ ---sr_scale 4 \ ---t_max 0.6667 \ ---t_min 0.3333 \ ---color_fix_type adain \ ---output experiments/test \ ---device cuda \ ---repeat_times 1 diff --git a/sh/3_create_weight.sh b/sh/3_create_weight.sh deleted file mode 100644 index 3ac895d..0000000 --- a/sh/3_create_weight.sh +++ /dev/null @@ -1,11 +0,0 @@ - - - - - - -python scripts/make_stage2_init_weight.py \ ---cldm_config configs/model/ccsr_stage1.yaml \ ---sd_weight ****/v2-1_512-ema-pruned.ckpt \ ---output weights/init_weight_ccsr.ckpt - diff --git a/sh/4_generate_data_file.sh b/sh/4_generate_data_file.sh deleted file mode 100644 index fe8fd24..0000000 --- a/sh/4_generate_data_file.sh +++ /dev/null @@ -1,10 +0,0 @@ - - - - -python scripts/make_file_list.py \ ---img_folder ****\ -**** \ ---val_size 10 \ ---save_folder 'preset/train_datasets' \ ---follow_links \ No newline at end of file diff --git a/sh/5_train_stage1.sh b/sh/5_train_stage1.sh deleted file mode 100644 index b5abd7a..0000000 --- a/sh/5_train_stage1.sh +++ /dev/null @@ -1,3 +0,0 @@ - - -python train.py --config configs/train_ccsr_stage1.yaml \ No newline at end of file diff --git a/sh/6_train_stage2.sh b/sh/6_train_stage2.sh deleted file mode 100644 index 8638f3d..0000000 --- a/sh/6_train_stage2.sh +++ /dev/null @@ -1,3 +0,0 @@ - - -python train.py --config configs/train_ccsr_stage2.yaml \ No newline at end of file diff --git a/taming/data/ade20k.py b/taming/data/ade20k.py new file mode 100644 index 0000000..366dae9 --- /dev/null +++ b/taming/data/ade20k.py @@ -0,0 +1,124 @@ +import os +import numpy as np +import cv2 +import albumentations +from PIL import Image +from torch.utils.data import Dataset + +from taming.data.sflckr import SegmentationBase # for examples included in repo + + +class Examples(SegmentationBase): + def __init__(self, size=256, random_crop=False, interpolation="bicubic"): + super().__init__(data_csv="data/ade20k_examples.txt", + data_root="data/ade20k_images", + segmentation_root="data/ade20k_segmentations", + size=size, random_crop=random_crop, + interpolation=interpolation, + n_labels=151, shift_segmentation=False) + + +# With semantic map and scene label +class ADE20kBase(Dataset): + def __init__(self, config=None, size=None, random_crop=False, interpolation="bicubic", crop_size=None): + self.split = self.get_split() + self.n_labels = 151 # unknown + 150 + self.data_csv = {"train": "data/ade20k_train.txt", + "validation": "data/ade20k_test.txt"}[self.split] + self.data_root = "data/ade20k_root" + with open(os.path.join(self.data_root, "sceneCategories.txt"), "r") as f: + self.scene_categories = f.read().splitlines() + self.scene_categories = dict(line.split() for line in self.scene_categories) + with open(self.data_csv, "r") as f: + self.image_paths = f.read().splitlines() + self._length = len(self.image_paths) + self.labels = { + "relative_file_path_": [l for l in self.image_paths], + "file_path_": [os.path.join(self.data_root, "images", l) + for l in self.image_paths], + "relative_segmentation_path_": [l.replace(".jpg", ".png") + for l in self.image_paths], + "segmentation_path_": [os.path.join(self.data_root, "annotations", + l.replace(".jpg", ".png")) + for l in self.image_paths], + "scene_category": [self.scene_categories[l.split("/")[1].replace(".jpg", "")] + for l in self.image_paths], + } + + size = None if size is not None and size<=0 else size + self.size = size + if crop_size is None: + self.crop_size = size if size is not None else None + else: + self.crop_size = crop_size + if self.size is not None: + self.interpolation = interpolation + self.interpolation = { + "nearest": cv2.INTER_NEAREST, + "bilinear": cv2.INTER_LINEAR, + "bicubic": cv2.INTER_CUBIC, + "area": cv2.INTER_AREA, + "lanczos": cv2.INTER_LANCZOS4}[self.interpolation] + self.image_rescaler = albumentations.SmallestMaxSize(max_size=self.size, + interpolation=self.interpolation) + self.segmentation_rescaler = albumentations.SmallestMaxSize(max_size=self.size, + interpolation=cv2.INTER_NEAREST) + + if crop_size is not None: + self.center_crop = not random_crop + if self.center_crop: + self.cropper = albumentations.CenterCrop(height=self.crop_size, width=self.crop_size) + else: + self.cropper = albumentations.RandomCrop(height=self.crop_size, width=self.crop_size) + self.preprocessor = self.cropper + + def __len__(self): + return self._length + + def __getitem__(self, i): + example = dict((k, self.labels[k][i]) for k in self.labels) + image = Image.open(example["file_path_"]) + if not image.mode == "RGB": + image = image.convert("RGB") + image = np.array(image).astype(np.uint8) + if self.size is not None: + image = self.image_rescaler(image=image)["image"] + segmentation = Image.open(example["segmentation_path_"]) + segmentation = np.array(segmentation).astype(np.uint8) + if self.size is not None: + segmentation = self.segmentation_rescaler(image=segmentation)["image"] + if self.size is not None: + processed = self.preprocessor(image=image, mask=segmentation) + else: + processed = {"image": image, "mask": segmentation} + example["image"] = (processed["image"]/127.5 - 1.0).astype(np.float32) + segmentation = processed["mask"] + onehot = np.eye(self.n_labels)[segmentation] + example["segmentation"] = onehot + return example + + +class ADE20kTrain(ADE20kBase): + # default to random_crop=True + def __init__(self, config=None, size=None, random_crop=True, interpolation="bicubic", crop_size=None): + super().__init__(config=config, size=size, random_crop=random_crop, + interpolation=interpolation, crop_size=crop_size) + + def get_split(self): + return "train" + + +class ADE20kValidation(ADE20kBase): + def get_split(self): + return "validation" + + +if __name__ == "__main__": + dset = ADE20kValidation() + ex = dset[0] + for k in ["image", "scene_category", "segmentation"]: + print(type(ex[k])) + try: + print(ex[k].shape) + except: + print(ex[k]) diff --git a/taming/data/annotated_objects_coco.py b/taming/data/annotated_objects_coco.py new file mode 100644 index 0000000..af000ec --- /dev/null +++ b/taming/data/annotated_objects_coco.py @@ -0,0 +1,139 @@ +import json +from itertools import chain +from pathlib import Path +from typing import Iterable, Dict, List, Callable, Any +from collections import defaultdict + +from tqdm import tqdm + +from taming.data.annotated_objects_dataset import AnnotatedObjectsDataset +from taming.data.helper_types import Annotation, ImageDescription, Category + +COCO_PATH_STRUCTURE = { + 'train': { + 'top_level': '', + 'instances_annotations': 'annotations/instances_train2017.json', + 'stuff_annotations': 'annotations/stuff_train2017.json', + 'files': 'train2017' + }, + 'validation': { + 'top_level': '', + 'instances_annotations': 'annotations/instances_val2017.json', + 'stuff_annotations': 'annotations/stuff_val2017.json', + 'files': 'val2017' + } +} + + +def load_image_descriptions(description_json: List[Dict]) -> Dict[str, ImageDescription]: + return { + str(img['id']): ImageDescription( + id=img['id'], + license=img.get('license'), + file_name=img['file_name'], + coco_url=img['coco_url'], + original_size=(img['width'], img['height']), + date_captured=img.get('date_captured'), + flickr_url=img.get('flickr_url') + ) + for img in description_json + } + + +def load_categories(category_json: Iterable) -> Dict[str, Category]: + return {str(cat['id']): Category(id=str(cat['id']), super_category=cat['supercategory'], name=cat['name']) + for cat in category_json if cat['name'] != 'other'} + + +def load_annotations(annotations_json: List[Dict], image_descriptions: Dict[str, ImageDescription], + category_no_for_id: Callable[[str], int], split: str) -> Dict[str, List[Annotation]]: + annotations = defaultdict(list) + total = sum(len(a) for a in annotations_json) + for ann in tqdm(chain(*annotations_json), f'Loading {split} annotations', total=total): + image_id = str(ann['image_id']) + if image_id not in image_descriptions: + raise ValueError(f'image_id [{image_id}] has no image description.') + category_id = ann['category_id'] + try: + category_no = category_no_for_id(str(category_id)) + except KeyError: + continue + + width, height = image_descriptions[image_id].original_size + bbox = (ann['bbox'][0] / width, ann['bbox'][1] / height, ann['bbox'][2] / width, ann['bbox'][3] / height) + + annotations[image_id].append( + Annotation( + id=ann['id'], + area=bbox[2]*bbox[3], # use bbox area + is_group_of=ann['iscrowd'], + image_id=ann['image_id'], + bbox=bbox, + category_id=str(category_id), + category_no=category_no + ) + ) + return dict(annotations) + + +class AnnotatedObjectsCoco(AnnotatedObjectsDataset): + def __init__(self, use_things: bool = True, use_stuff: bool = True, **kwargs): + """ + @param data_path: is the path to the following folder structure: + coco/ + ├── annotations + │ ├── instances_train2017.json + │ ├── instances_val2017.json + │ ├── stuff_train2017.json + │ └── stuff_val2017.json + ├── train2017 + │ ├── 000000000009.jpg + │ ├── 000000000025.jpg + │ └── ... + ├── val2017 + │ ├── 000000000139.jpg + │ ├── 000000000285.jpg + │ └── ... + @param: split: one of 'train' or 'validation' + @param: desired image size (give square images) + """ + super().__init__(**kwargs) + self.use_things = use_things + self.use_stuff = use_stuff + + with open(self.paths['instances_annotations']) as f: + inst_data_json = json.load(f) + with open(self.paths['stuff_annotations']) as f: + stuff_data_json = json.load(f) + + category_jsons = [] + annotation_jsons = [] + if self.use_things: + category_jsons.append(inst_data_json['categories']) + annotation_jsons.append(inst_data_json['annotations']) + if self.use_stuff: + category_jsons.append(stuff_data_json['categories']) + annotation_jsons.append(stuff_data_json['annotations']) + + self.categories = load_categories(chain(*category_jsons)) + self.filter_categories() + self.setup_category_id_and_number() + + self.image_descriptions = load_image_descriptions(inst_data_json['images']) + annotations = load_annotations(annotation_jsons, self.image_descriptions, self.get_category_number, self.split) + self.annotations = self.filter_object_number(annotations, self.min_object_area, + self.min_objects_per_image, self.max_objects_per_image) + self.image_ids = list(self.annotations.keys()) + self.clean_up_annotations_and_image_descriptions() + + def get_path_structure(self) -> Dict[str, str]: + if self.split not in COCO_PATH_STRUCTURE: + raise ValueError(f'Split [{self.split} does not exist for COCO data.]') + return COCO_PATH_STRUCTURE[self.split] + + def get_image_path(self, image_id: str) -> Path: + return self.paths['files'].joinpath(self.image_descriptions[str(image_id)].file_name) + + def get_image_description(self, image_id: str) -> Dict[str, Any]: + # noinspection PyProtectedMember + return self.image_descriptions[image_id]._asdict() diff --git a/taming/data/annotated_objects_dataset.py b/taming/data/annotated_objects_dataset.py new file mode 100644 index 0000000..53cc346 --- /dev/null +++ b/taming/data/annotated_objects_dataset.py @@ -0,0 +1,218 @@ +from pathlib import Path +from typing import Optional, List, Callable, Dict, Any, Union +import warnings + +import PIL.Image as pil_image +from torch import Tensor +from torch.utils.data import Dataset +from torchvision import transforms + +from taming.data.conditional_builder.objects_bbox import ObjectsBoundingBoxConditionalBuilder +from taming.data.conditional_builder.objects_center_points import ObjectsCenterPointsConditionalBuilder +from taming.data.conditional_builder.utils import load_object_from_string +from taming.data.helper_types import BoundingBox, CropMethodType, Image, Annotation, SplitType +from taming.data.image_transforms import CenterCropReturnCoordinates, RandomCrop1dReturnCoordinates, \ + Random2dCropReturnCoordinates, RandomHorizontalFlipReturn, convert_pil_to_tensor + + +class AnnotatedObjectsDataset(Dataset): + def __init__(self, data_path: Union[str, Path], split: SplitType, keys: List[str], target_image_size: int, + min_object_area: float, min_objects_per_image: int, max_objects_per_image: int, + crop_method: CropMethodType, random_flip: bool, no_tokens: int, use_group_parameter: bool, + encode_crop: bool, category_allow_list_target: str = "", category_mapping_target: str = "", + no_object_classes: Optional[int] = None): + self.data_path = data_path + self.split = split + self.keys = keys + self.target_image_size = target_image_size + self.min_object_area = min_object_area + self.min_objects_per_image = min_objects_per_image + self.max_objects_per_image = max_objects_per_image + self.crop_method = crop_method + self.random_flip = random_flip + self.no_tokens = no_tokens + self.use_group_parameter = use_group_parameter + self.encode_crop = encode_crop + + self.annotations = None + self.image_descriptions = None + self.categories = None + self.category_ids = None + self.category_number = None + self.image_ids = None + self.transform_functions: List[Callable] = self.setup_transform(target_image_size, crop_method, random_flip) + self.paths = self.build_paths(self.data_path) + self._conditional_builders = None + self.category_allow_list = None + if category_allow_list_target: + allow_list = load_object_from_string(category_allow_list_target) + self.category_allow_list = {name for name, _ in allow_list} + self.category_mapping = {} + if category_mapping_target: + self.category_mapping = load_object_from_string(category_mapping_target) + self.no_object_classes = no_object_classes + + def build_paths(self, top_level: Union[str, Path]) -> Dict[str, Path]: + top_level = Path(top_level) + sub_paths = {name: top_level.joinpath(sub_path) for name, sub_path in self.get_path_structure().items()} + for path in sub_paths.values(): + if not path.exists(): + raise FileNotFoundError(f'{type(self).__name__} data structure error: [{path}] does not exist.') + return sub_paths + + @staticmethod + def load_image_from_disk(path: Path) -> Image: + return pil_image.open(path).convert('RGB') + + @staticmethod + def setup_transform(target_image_size: int, crop_method: CropMethodType, random_flip: bool): + transform_functions = [] + if crop_method == 'none': + transform_functions.append(transforms.Resize((target_image_size, target_image_size))) + elif crop_method == 'center': + transform_functions.extend([ + transforms.Resize(target_image_size), + CenterCropReturnCoordinates(target_image_size) + ]) + elif crop_method == 'random-1d': + transform_functions.extend([ + transforms.Resize(target_image_size), + RandomCrop1dReturnCoordinates(target_image_size) + ]) + elif crop_method == 'random-2d': + transform_functions.extend([ + Random2dCropReturnCoordinates(target_image_size), + transforms.Resize(target_image_size) + ]) + elif crop_method is None: + return None + else: + raise ValueError(f'Received invalid crop method [{crop_method}].') + if random_flip: + transform_functions.append(RandomHorizontalFlipReturn()) + transform_functions.append(transforms.Lambda(lambda x: x / 127.5 - 1.)) + return transform_functions + + def image_transform(self, x: Tensor) -> (Optional[BoundingBox], Optional[bool], Tensor): + crop_bbox = None + flipped = None + for t in self.transform_functions: + if isinstance(t, (RandomCrop1dReturnCoordinates, CenterCropReturnCoordinates, Random2dCropReturnCoordinates)): + crop_bbox, x = t(x) + elif isinstance(t, RandomHorizontalFlipReturn): + flipped, x = t(x) + else: + x = t(x) + return crop_bbox, flipped, x + + @property + def no_classes(self) -> int: + return self.no_object_classes if self.no_object_classes else len(self.categories) + + @property + def conditional_builders(self) -> ObjectsCenterPointsConditionalBuilder: + # cannot set this up in init because no_classes is only known after loading data in init of superclass + if self._conditional_builders is None: + self._conditional_builders = { + 'objects_center_points': ObjectsCenterPointsConditionalBuilder( + self.no_classes, + self.max_objects_per_image, + self.no_tokens, + self.encode_crop, + self.use_group_parameter, + getattr(self, 'use_additional_parameters', False) + ), + 'objects_bbox': ObjectsBoundingBoxConditionalBuilder( + self.no_classes, + self.max_objects_per_image, + self.no_tokens, + self.encode_crop, + self.use_group_parameter, + getattr(self, 'use_additional_parameters', False) + ) + } + return self._conditional_builders + + def filter_categories(self) -> None: + if self.category_allow_list: + self.categories = {id_: cat for id_, cat in self.categories.items() if cat.name in self.category_allow_list} + if self.category_mapping: + self.categories = {id_: cat for id_, cat in self.categories.items() if cat.id not in self.category_mapping} + + def setup_category_id_and_number(self) -> None: + self.category_ids = list(self.categories.keys()) + self.category_ids.sort() + if '/m/01s55n' in self.category_ids: + self.category_ids.remove('/m/01s55n') + self.category_ids.append('/m/01s55n') + self.category_number = {category_id: i for i, category_id in enumerate(self.category_ids)} + if self.category_allow_list is not None and self.category_mapping is None \ + and len(self.category_ids) != len(self.category_allow_list): + warnings.warn('Unexpected number of categories: Mismatch with category_allow_list. ' + 'Make sure all names in category_allow_list exist.') + + def clean_up_annotations_and_image_descriptions(self) -> None: + image_id_set = set(self.image_ids) + self.annotations = {k: v for k, v in self.annotations.items() if k in image_id_set} + self.image_descriptions = {k: v for k, v in self.image_descriptions.items() if k in image_id_set} + + @staticmethod + def filter_object_number(all_annotations: Dict[str, List[Annotation]], min_object_area: float, + min_objects_per_image: int, max_objects_per_image: int) -> Dict[str, List[Annotation]]: + filtered = {} + for image_id, annotations in all_annotations.items(): + annotations_with_min_area = [a for a in annotations if a.area > min_object_area] + if min_objects_per_image <= len(annotations_with_min_area) <= max_objects_per_image: + filtered[image_id] = annotations_with_min_area + return filtered + + def __len__(self): + return len(self.image_ids) + + def __getitem__(self, n: int) -> Dict[str, Any]: + image_id = self.get_image_id(n) + sample = self.get_image_description(image_id) + sample['annotations'] = self.get_annotation(image_id) + + if 'image' in self.keys: + sample['image_path'] = str(self.get_image_path(image_id)) + sample['image'] = self.load_image_from_disk(sample['image_path']) + sample['image'] = convert_pil_to_tensor(sample['image']) + sample['crop_bbox'], sample['flipped'], sample['image'] = self.image_transform(sample['image']) + sample['image'] = sample['image'].permute(1, 2, 0) + + for conditional, builder in self.conditional_builders.items(): + if conditional in self.keys: + sample[conditional] = builder.build(sample['annotations'], sample['crop_bbox'], sample['flipped']) + + if self.keys: + # only return specified keys + sample = {key: sample[key] for key in self.keys} + return sample + + def get_image_id(self, no: int) -> str: + return self.image_ids[no] + + def get_annotation(self, image_id: str) -> str: + return self.annotations[image_id] + + def get_textual_label_for_category_id(self, category_id: str) -> str: + return self.categories[category_id].name + + def get_textual_label_for_category_no(self, category_no: int) -> str: + return self.categories[self.get_category_id(category_no)].name + + def get_category_number(self, category_id: str) -> int: + return self.category_number[category_id] + + def get_category_id(self, category_no: int) -> str: + return self.category_ids[category_no] + + def get_image_description(self, image_id: str) -> Dict[str, Any]: + raise NotImplementedError() + + def get_path_structure(self): + raise NotImplementedError + + def get_image_path(self, image_id: str) -> Path: + raise NotImplementedError diff --git a/taming/data/annotated_objects_open_images.py b/taming/data/annotated_objects_open_images.py new file mode 100644 index 0000000..aede680 --- /dev/null +++ b/taming/data/annotated_objects_open_images.py @@ -0,0 +1,137 @@ +from collections import defaultdict +from csv import DictReader, reader as TupleReader +from pathlib import Path +from typing import Dict, List, Any +import warnings + +from taming.data.annotated_objects_dataset import AnnotatedObjectsDataset +from taming.data.helper_types import Annotation, Category +from tqdm import tqdm + +OPEN_IMAGES_STRUCTURE = { + 'train': { + 'top_level': '', + 'class_descriptions': 'class-descriptions-boxable.csv', + 'annotations': 'oidv6-train-annotations-bbox.csv', + 'file_list': 'train-images-boxable.csv', + 'files': 'train' + }, + 'validation': { + 'top_level': '', + 'class_descriptions': 'class-descriptions-boxable.csv', + 'annotations': 'validation-annotations-bbox.csv', + 'file_list': 'validation-images.csv', + 'files': 'validation' + }, + 'test': { + 'top_level': '', + 'class_descriptions': 'class-descriptions-boxable.csv', + 'annotations': 'test-annotations-bbox.csv', + 'file_list': 'test-images.csv', + 'files': 'test' + } +} + + +def load_annotations(descriptor_path: Path, min_object_area: float, category_mapping: Dict[str, str], + category_no_for_id: Dict[str, int]) -> Dict[str, List[Annotation]]: + annotations: Dict[str, List[Annotation]] = defaultdict(list) + with open(descriptor_path) as file: + reader = DictReader(file) + for i, row in tqdm(enumerate(reader), total=14620000, desc='Loading OpenImages annotations'): + width = float(row['XMax']) - float(row['XMin']) + height = float(row['YMax']) - float(row['YMin']) + area = width * height + category_id = row['LabelName'] + if category_id in category_mapping: + category_id = category_mapping[category_id] + if area >= min_object_area and category_id in category_no_for_id: + annotations[row['ImageID']].append( + Annotation( + id=i, + image_id=row['ImageID'], + source=row['Source'], + category_id=category_id, + category_no=category_no_for_id[category_id], + confidence=float(row['Confidence']), + bbox=(float(row['XMin']), float(row['YMin']), width, height), + area=area, + is_occluded=bool(int(row['IsOccluded'])), + is_truncated=bool(int(row['IsTruncated'])), + is_group_of=bool(int(row['IsGroupOf'])), + is_depiction=bool(int(row['IsDepiction'])), + is_inside=bool(int(row['IsInside'])) + ) + ) + if 'train' in str(descriptor_path) and i < 14000000: + warnings.warn(f'Running with subset of Open Images. Train dataset has length [{len(annotations)}].') + return dict(annotations) + + +def load_image_ids(csv_path: Path) -> List[str]: + with open(csv_path) as file: + reader = DictReader(file) + return [row['image_name'] for row in reader] + + +def load_categories(csv_path: Path) -> Dict[str, Category]: + with open(csv_path) as file: + reader = TupleReader(file) + return {row[0]: Category(id=row[0], name=row[1], super_category=None) for row in reader} + + +class AnnotatedObjectsOpenImages(AnnotatedObjectsDataset): + def __init__(self, use_additional_parameters: bool, **kwargs): + """ + @param data_path: is the path to the following folder structure: + open_images/ + │ oidv6-train-annotations-bbox.csv + ├── class-descriptions-boxable.csv + ├── oidv6-train-annotations-bbox.csv + ├── test + │ ├── 000026e7ee790996.jpg + │ ├── 000062a39995e348.jpg + │ └── ... + ├── test-annotations-bbox.csv + ├── test-images.csv + ├── train + │ ├── 000002b66c9c498e.jpg + │ ├── 000002b97e5471a0.jpg + │ └── ... + ├── train-images-boxable.csv + ├── validation + │ ├── 0001eeaf4aed83f9.jpg + │ ├── 0004886b7d043cfd.jpg + │ └── ... + ├── validation-annotations-bbox.csv + └── validation-images.csv + @param: split: one of 'train', 'validation' or 'test' + @param: desired image size (returns square images) + """ + + super().__init__(**kwargs) + self.use_additional_parameters = use_additional_parameters + + self.categories = load_categories(self.paths['class_descriptions']) + self.filter_categories() + self.setup_category_id_and_number() + + self.image_descriptions = {} + annotations = load_annotations(self.paths['annotations'], self.min_object_area, self.category_mapping, + self.category_number) + self.annotations = self.filter_object_number(annotations, self.min_object_area, self.min_objects_per_image, + self.max_objects_per_image) + self.image_ids = list(self.annotations.keys()) + self.clean_up_annotations_and_image_descriptions() + + def get_path_structure(self) -> Dict[str, str]: + if self.split not in OPEN_IMAGES_STRUCTURE: + raise ValueError(f'Split [{self.split} does not exist for Open Images data.]') + return OPEN_IMAGES_STRUCTURE[self.split] + + def get_image_path(self, image_id: str) -> Path: + return self.paths['files'].joinpath(f'{image_id:0>16}.jpg') + + def get_image_description(self, image_id: str) -> Dict[str, Any]: + image_path = self.get_image_path(image_id) + return {'file_path': str(image_path), 'file_name': image_path.name} diff --git a/taming/data/base.py b/taming/data/base.py new file mode 100644 index 0000000..e21667d --- /dev/null +++ b/taming/data/base.py @@ -0,0 +1,70 @@ +import bisect +import numpy as np +import albumentations +from PIL import Image +from torch.utils.data import Dataset, ConcatDataset + + +class ConcatDatasetWithIndex(ConcatDataset): + """Modified from original pytorch code to return dataset idx""" + def __getitem__(self, idx): + if idx < 0: + if -idx > len(self): + raise ValueError("absolute value of index should not exceed dataset length") + idx = len(self) + idx + dataset_idx = bisect.bisect_right(self.cumulative_sizes, idx) + if dataset_idx == 0: + sample_idx = idx + else: + sample_idx = idx - self.cumulative_sizes[dataset_idx - 1] + return self.datasets[dataset_idx][sample_idx], dataset_idx + + +class ImagePaths(Dataset): + def __init__(self, paths, size=None, random_crop=False, labels=None): + self.size = size + self.random_crop = random_crop + + self.labels = dict() if labels is None else labels + self.labels["file_path_"] = paths + self._length = len(paths) + + if self.size is not None and self.size > 0: + self.rescaler = albumentations.SmallestMaxSize(max_size = self.size) + if not self.random_crop: + self.cropper = albumentations.CenterCrop(height=self.size,width=self.size) + else: + self.cropper = albumentations.RandomCrop(height=self.size,width=self.size) + self.preprocessor = albumentations.Compose([self.rescaler, self.cropper]) + else: + self.preprocessor = lambda **kwargs: kwargs + + def __len__(self): + return self._length + + def preprocess_image(self, image_path): + image = Image.open(image_path) + if not image.mode == "RGB": + image = image.convert("RGB") + image = np.array(image).astype(np.uint8) + image = self.preprocessor(image=image)["image"] + image = (image/127.5 - 1.0).astype(np.float32) + return image + + def __getitem__(self, i): + example = dict() + example["image"] = self.preprocess_image(self.labels["file_path_"][i]) + for k in self.labels: + example[k] = self.labels[k][i] + return example + + +class NumpyPaths(ImagePaths): + def preprocess_image(self, image_path): + image = np.load(image_path).squeeze(0) # 3 x 1024 x 1024 + image = np.transpose(image, (1,2,0)) + image = Image.fromarray(image, mode="RGB") + image = np.array(image).astype(np.uint8) + image = self.preprocessor(image=image)["image"] + image = (image/127.5 - 1.0).astype(np.float32) + return image diff --git a/taming/data/coco.py b/taming/data/coco.py new file mode 100644 index 0000000..2b2f783 --- /dev/null +++ b/taming/data/coco.py @@ -0,0 +1,176 @@ +import os +import json +import albumentations +import numpy as np +from PIL import Image +from tqdm import tqdm +from torch.utils.data import Dataset + +from taming.data.sflckr import SegmentationBase # for examples included in repo + + +class Examples(SegmentationBase): + def __init__(self, size=256, random_crop=False, interpolation="bicubic"): + super().__init__(data_csv="data/coco_examples.txt", + data_root="data/coco_images", + segmentation_root="data/coco_segmentations", + size=size, random_crop=random_crop, + interpolation=interpolation, + n_labels=183, shift_segmentation=True) + + +class CocoBase(Dataset): + """needed for (image, caption, segmentation) pairs""" + def __init__(self, size=None, dataroot="", datajson="", onehot_segmentation=False, use_stuffthing=False, + crop_size=None, force_no_crop=False, given_files=None): + self.split = self.get_split() + self.size = size + if crop_size is None: + self.crop_size = size + else: + self.crop_size = crop_size + + self.onehot = onehot_segmentation # return segmentation as rgb or one hot + self.stuffthing = use_stuffthing # include thing in segmentation + if self.onehot and not self.stuffthing: + raise NotImplemented("One hot mode is only supported for the " + "stuffthings version because labels are stored " + "a bit different.") + + data_json = datajson + with open(data_json) as json_file: + self.json_data = json.load(json_file) + self.img_id_to_captions = dict() + self.img_id_to_filepath = dict() + self.img_id_to_segmentation_filepath = dict() + + assert data_json.split("/")[-1] in ["captions_train2017.json", + "captions_val2017.json"] + if self.stuffthing: + self.segmentation_prefix = ( + "data/cocostuffthings/val2017" if + data_json.endswith("captions_val2017.json") else + "data/cocostuffthings/train2017") + else: + self.segmentation_prefix = ( + "data/coco/annotations/stuff_val2017_pixelmaps" if + data_json.endswith("captions_val2017.json") else + "data/coco/annotations/stuff_train2017_pixelmaps") + + imagedirs = self.json_data["images"] + self.labels = {"image_ids": list()} + for imgdir in tqdm(imagedirs, desc="ImgToPath"): + self.img_id_to_filepath[imgdir["id"]] = os.path.join(dataroot, imgdir["file_name"]) + self.img_id_to_captions[imgdir["id"]] = list() + pngfilename = imgdir["file_name"].replace("jpg", "png") + self.img_id_to_segmentation_filepath[imgdir["id"]] = os.path.join( + self.segmentation_prefix, pngfilename) + if given_files is not None: + if pngfilename in given_files: + self.labels["image_ids"].append(imgdir["id"]) + else: + self.labels["image_ids"].append(imgdir["id"]) + + capdirs = self.json_data["annotations"] + for capdir in tqdm(capdirs, desc="ImgToCaptions"): + # there are in average 5 captions per image + self.img_id_to_captions[capdir["image_id"]].append(np.array([capdir["caption"]])) + + self.rescaler = albumentations.SmallestMaxSize(max_size=self.size) + if self.split=="validation": + self.cropper = albumentations.CenterCrop(height=self.crop_size, width=self.crop_size) + else: + self.cropper = albumentations.RandomCrop(height=self.crop_size, width=self.crop_size) + self.preprocessor = albumentations.Compose( + [self.rescaler, self.cropper], + additional_targets={"segmentation": "image"}) + if force_no_crop: + self.rescaler = albumentations.Resize(height=self.size, width=self.size) + self.preprocessor = albumentations.Compose( + [self.rescaler], + additional_targets={"segmentation": "image"}) + + def __len__(self): + return len(self.labels["image_ids"]) + + def preprocess_image(self, image_path, segmentation_path): + image = Image.open(image_path) + if not image.mode == "RGB": + image = image.convert("RGB") + image = np.array(image).astype(np.uint8) + + segmentation = Image.open(segmentation_path) + if not self.onehot and not segmentation.mode == "RGB": + segmentation = segmentation.convert("RGB") + segmentation = np.array(segmentation).astype(np.uint8) + if self.onehot: + assert self.stuffthing + # stored in caffe format: unlabeled==255. stuff and thing from + # 0-181. to be compatible with the labels in + # https://github.com/nightrome/cocostuff/blob/master/labels.txt + # we shift stuffthing one to the right and put unlabeled in zero + # as long as segmentation is uint8 shifting to right handles the + # latter too + assert segmentation.dtype == np.uint8 + segmentation = segmentation + 1 + + processed = self.preprocessor(image=image, segmentation=segmentation) + image, segmentation = processed["image"], processed["segmentation"] + image = (image / 127.5 - 1.0).astype(np.float32) + + if self.onehot: + assert segmentation.dtype == np.uint8 + # make it one hot + n_labels = 183 + flatseg = np.ravel(segmentation) + onehot = np.zeros((flatseg.size, n_labels), dtype=np.bool) + onehot[np.arange(flatseg.size), flatseg] = True + onehot = onehot.reshape(segmentation.shape + (n_labels,)).astype(int) + segmentation = onehot + else: + segmentation = (segmentation / 127.5 - 1.0).astype(np.float32) + return image, segmentation + + def __getitem__(self, i): + img_path = self.img_id_to_filepath[self.labels["image_ids"][i]] + seg_path = self.img_id_to_segmentation_filepath[self.labels["image_ids"][i]] + image, segmentation = self.preprocess_image(img_path, seg_path) + captions = self.img_id_to_captions[self.labels["image_ids"][i]] + # randomly draw one of all available captions per image + caption = captions[np.random.randint(0, len(captions))] + example = {"image": image, + "caption": [str(caption[0])], + "segmentation": segmentation, + "img_path": img_path, + "seg_path": seg_path, + "filename_": img_path.split(os.sep)[-1] + } + return example + + +class CocoImagesAndCaptionsTrain(CocoBase): + """returns a pair of (image, caption)""" + def __init__(self, size, onehot_segmentation=False, use_stuffthing=False, crop_size=None, force_no_crop=False): + super().__init__(size=size, + dataroot="data/coco/train2017", + datajson="data/coco/annotations/captions_train2017.json", + onehot_segmentation=onehot_segmentation, + use_stuffthing=use_stuffthing, crop_size=crop_size, force_no_crop=force_no_crop) + + def get_split(self): + return "train" + + +class CocoImagesAndCaptionsValidation(CocoBase): + """returns a pair of (image, caption)""" + def __init__(self, size, onehot_segmentation=False, use_stuffthing=False, crop_size=None, force_no_crop=False, + given_files=None): + super().__init__(size=size, + dataroot="data/coco/val2017", + datajson="data/coco/annotations/captions_val2017.json", + onehot_segmentation=onehot_segmentation, + use_stuffthing=use_stuffthing, crop_size=crop_size, force_no_crop=force_no_crop, + given_files=given_files) + + def get_split(self): + return "validation" diff --git a/taming/data/conditional_builder/objects_bbox.py b/taming/data/conditional_builder/objects_bbox.py new file mode 100644 index 0000000..15881e7 --- /dev/null +++ b/taming/data/conditional_builder/objects_bbox.py @@ -0,0 +1,60 @@ +from itertools import cycle +from typing import List, Tuple, Callable, Optional + +from PIL import Image as pil_image, ImageDraw as pil_img_draw, ImageFont +from more_itertools.recipes import grouper +from taming.data.image_transforms import convert_pil_to_tensor +from torch import LongTensor, Tensor + +from taming.data.helper_types import BoundingBox, Annotation +from taming.data.conditional_builder.objects_center_points import ObjectsCenterPointsConditionalBuilder +from taming.data.conditional_builder.utils import COLOR_PALETTE, WHITE, GRAY_75, BLACK, additional_parameters_string, \ + pad_list, get_plot_font_size, absolute_bbox + + +class ObjectsBoundingBoxConditionalBuilder(ObjectsCenterPointsConditionalBuilder): + @property + def object_descriptor_length(self) -> int: + return 3 + + def _make_object_descriptors(self, annotations: List[Annotation]) -> List[Tuple[int, ...]]: + object_triples = [ + (self.object_representation(ann), *self.token_pair_from_bbox(ann.bbox)) + for ann in annotations + ] + empty_triple = (self.none, self.none, self.none) + object_triples = pad_list(object_triples, empty_triple, self.no_max_objects) + return object_triples + + def inverse_build(self, conditional: LongTensor) -> Tuple[List[Tuple[int, BoundingBox]], Optional[BoundingBox]]: + conditional_list = conditional.tolist() + crop_coordinates = None + if self.encode_crop: + crop_coordinates = self.bbox_from_token_pair(conditional_list[-2], conditional_list[-1]) + conditional_list = conditional_list[:-2] + object_triples = grouper(conditional_list, 3) + assert conditional.shape[0] == self.embedding_dim + return [ + (object_triple[0], self.bbox_from_token_pair(object_triple[1], object_triple[2])) + for object_triple in object_triples if object_triple[0] != self.none + ], crop_coordinates + + def plot(self, conditional: LongTensor, label_for_category_no: Callable[[int], str], figure_size: Tuple[int, int], + line_width: int = 3, font_size: Optional[int] = None) -> Tensor: + plot = pil_image.new('RGB', figure_size, WHITE) + draw = pil_img_draw.Draw(plot) + font = ImageFont.truetype( + "/usr/share/fonts/truetype/lato/Lato-Regular.ttf", + size=get_plot_font_size(font_size, figure_size) + ) + width, height = plot.size + description, crop_coordinates = self.inverse_build(conditional) + for (representation, bbox), color in zip(description, cycle(COLOR_PALETTE)): + annotation = self.representation_to_annotation(representation) + class_label = label_for_category_no(annotation.category_no) + ' ' + additional_parameters_string(annotation) + bbox = absolute_bbox(bbox, width, height) + draw.rectangle(bbox, outline=color, width=line_width) + draw.text((bbox[0] + line_width, bbox[1] + line_width), class_label, anchor='la', fill=BLACK, font=font) + if crop_coordinates is not None: + draw.rectangle(absolute_bbox(crop_coordinates, width, height), outline=GRAY_75, width=line_width) + return convert_pil_to_tensor(plot) / 127.5 - 1. diff --git a/taming/data/conditional_builder/objects_center_points.py b/taming/data/conditional_builder/objects_center_points.py new file mode 100644 index 0000000..9a48032 --- /dev/null +++ b/taming/data/conditional_builder/objects_center_points.py @@ -0,0 +1,168 @@ +import math +import random +import warnings +from itertools import cycle +from typing import List, Optional, Tuple, Callable + +from PIL import Image as pil_image, ImageDraw as pil_img_draw, ImageFont +from more_itertools.recipes import grouper +from taming.data.conditional_builder.utils import COLOR_PALETTE, WHITE, GRAY_75, BLACK, FULL_CROP, filter_annotations, \ + additional_parameters_string, horizontally_flip_bbox, pad_list, get_circle_size, get_plot_font_size, \ + absolute_bbox, rescale_annotations +from taming.data.helper_types import BoundingBox, Annotation +from taming.data.image_transforms import convert_pil_to_tensor +from torch import LongTensor, Tensor + + +class ObjectsCenterPointsConditionalBuilder: + def __init__(self, no_object_classes: int, no_max_objects: int, no_tokens: int, encode_crop: bool, + use_group_parameter: bool, use_additional_parameters: bool): + self.no_object_classes = no_object_classes + self.no_max_objects = no_max_objects + self.no_tokens = no_tokens + self.encode_crop = encode_crop + self.no_sections = int(math.sqrt(self.no_tokens)) + self.use_group_parameter = use_group_parameter + self.use_additional_parameters = use_additional_parameters + + @property + def none(self) -> int: + return self.no_tokens - 1 + + @property + def object_descriptor_length(self) -> int: + return 2 + + @property + def embedding_dim(self) -> int: + extra_length = 2 if self.encode_crop else 0 + return self.no_max_objects * self.object_descriptor_length + extra_length + + def tokenize_coordinates(self, x: float, y: float) -> int: + """ + Express 2d coordinates with one number. + Example: assume self.no_tokens = 16, then no_sections = 4: + 0 0 0 0 + 0 0 # 0 + 0 0 0 0 + 0 0 0 x + Then the # position corresponds to token 6, the x position to token 15. + @param x: float in [0, 1] + @param y: float in [0, 1] + @return: discrete tokenized coordinate + """ + x_discrete = int(round(x * (self.no_sections - 1))) + y_discrete = int(round(y * (self.no_sections - 1))) + return y_discrete * self.no_sections + x_discrete + + def coordinates_from_token(self, token: int) -> (float, float): + x = token % self.no_sections + y = token // self.no_sections + return x / (self.no_sections - 1), y / (self.no_sections - 1) + + def bbox_from_token_pair(self, token1: int, token2: int) -> BoundingBox: + x0, y0 = self.coordinates_from_token(token1) + x1, y1 = self.coordinates_from_token(token2) + return x0, y0, x1 - x0, y1 - y0 + + def token_pair_from_bbox(self, bbox: BoundingBox) -> Tuple[int, int]: + return self.tokenize_coordinates(bbox[0], bbox[1]), \ + self.tokenize_coordinates(bbox[0] + bbox[2], bbox[1] + bbox[3]) + + def inverse_build(self, conditional: LongTensor) \ + -> Tuple[List[Tuple[int, Tuple[float, float]]], Optional[BoundingBox]]: + conditional_list = conditional.tolist() + crop_coordinates = None + if self.encode_crop: + crop_coordinates = self.bbox_from_token_pair(conditional_list[-2], conditional_list[-1]) + conditional_list = conditional_list[:-2] + table_of_content = grouper(conditional_list, self.object_descriptor_length) + assert conditional.shape[0] == self.embedding_dim + return [ + (object_tuple[0], self.coordinates_from_token(object_tuple[1])) + for object_tuple in table_of_content if object_tuple[0] != self.none + ], crop_coordinates + + def plot(self, conditional: LongTensor, label_for_category_no: Callable[[int], str], figure_size: Tuple[int, int], + line_width: int = 3, font_size: Optional[int] = None) -> Tensor: + plot = pil_image.new('RGB', figure_size, WHITE) + draw = pil_img_draw.Draw(plot) + circle_size = get_circle_size(figure_size) + font = ImageFont.truetype('/usr/share/fonts/truetype/lato/Lato-Regular.ttf', + size=get_plot_font_size(font_size, figure_size)) + width, height = plot.size + description, crop_coordinates = self.inverse_build(conditional) + for (representation, (x, y)), color in zip(description, cycle(COLOR_PALETTE)): + x_abs, y_abs = x * width, y * height + ann = self.representation_to_annotation(representation) + label = label_for_category_no(ann.category_no) + ' ' + additional_parameters_string(ann) + ellipse_bbox = [x_abs - circle_size, y_abs - circle_size, x_abs + circle_size, y_abs + circle_size] + draw.ellipse(ellipse_bbox, fill=color, width=0) + draw.text((x_abs, y_abs), label, anchor='md', fill=BLACK, font=font) + if crop_coordinates is not None: + draw.rectangle(absolute_bbox(crop_coordinates, width, height), outline=GRAY_75, width=line_width) + return convert_pil_to_tensor(plot) / 127.5 - 1. + + def object_representation(self, annotation: Annotation) -> int: + modifier = 0 + if self.use_group_parameter: + modifier |= 1 * (annotation.is_group_of is True) + if self.use_additional_parameters: + modifier |= 2 * (annotation.is_occluded is True) + modifier |= 4 * (annotation.is_depiction is True) + modifier |= 8 * (annotation.is_inside is True) + return annotation.category_no + self.no_object_classes * modifier + + def representation_to_annotation(self, representation: int) -> Annotation: + category_no = representation % self.no_object_classes + modifier = representation // self.no_object_classes + # noinspection PyTypeChecker + return Annotation( + area=None, image_id=None, bbox=None, category_id=None, id=None, source=None, confidence=None, + category_no=category_no, + is_group_of=bool((modifier & 1) * self.use_group_parameter), + is_occluded=bool((modifier & 2) * self.use_additional_parameters), + is_depiction=bool((modifier & 4) * self.use_additional_parameters), + is_inside=bool((modifier & 8) * self.use_additional_parameters) + ) + + def _crop_encoder(self, crop_coordinates: BoundingBox) -> List[int]: + return list(self.token_pair_from_bbox(crop_coordinates)) + + def _make_object_descriptors(self, annotations: List[Annotation]) -> List[Tuple[int, ...]]: + object_tuples = [ + (self.object_representation(a), + self.tokenize_coordinates(a.bbox[0] + a.bbox[2] / 2, a.bbox[1] + a.bbox[3] / 2)) + for a in annotations + ] + empty_tuple = (self.none, self.none) + object_tuples = pad_list(object_tuples, empty_tuple, self.no_max_objects) + return object_tuples + + def build(self, annotations: List, crop_coordinates: Optional[BoundingBox] = None, horizontal_flip: bool = False) \ + -> LongTensor: + if len(annotations) == 0: + warnings.warn('Did not receive any annotations.') + if len(annotations) > self.no_max_objects: + warnings.warn('Received more annotations than allowed.') + annotations = annotations[:self.no_max_objects] + + if not crop_coordinates: + crop_coordinates = FULL_CROP + + random.shuffle(annotations) + annotations = filter_annotations(annotations, crop_coordinates) + if self.encode_crop: + annotations = rescale_annotations(annotations, FULL_CROP, horizontal_flip) + if horizontal_flip: + crop_coordinates = horizontally_flip_bbox(crop_coordinates) + extra = self._crop_encoder(crop_coordinates) + else: + annotations = rescale_annotations(annotations, crop_coordinates, horizontal_flip) + extra = [] + + object_tuples = self._make_object_descriptors(annotations) + flattened = [token for tuple_ in object_tuples for token in tuple_] + extra + assert len(flattened) == self.embedding_dim + assert all(0 <= value < self.no_tokens for value in flattened) + return LongTensor(flattened) diff --git a/taming/data/conditional_builder/utils.py b/taming/data/conditional_builder/utils.py new file mode 100644 index 0000000..d0ee175 --- /dev/null +++ b/taming/data/conditional_builder/utils.py @@ -0,0 +1,105 @@ +import importlib +from typing import List, Any, Tuple, Optional + +from taming.data.helper_types import BoundingBox, Annotation + +# source: seaborn, color palette tab10 +COLOR_PALETTE = [(30, 118, 179), (255, 126, 13), (43, 159, 43), (213, 38, 39), (147, 102, 188), + (139, 85, 74), (226, 118, 193), (126, 126, 126), (187, 188, 33), (22, 189, 206)] +BLACK = (0, 0, 0) +GRAY_75 = (63, 63, 63) +GRAY_50 = (127, 127, 127) +GRAY_25 = (191, 191, 191) +WHITE = (255, 255, 255) +FULL_CROP = (0., 0., 1., 1.) + + +def intersection_area(rectangle1: BoundingBox, rectangle2: BoundingBox) -> float: + """ + Give intersection area of two rectangles. + @param rectangle1: (x0, y0, w, h) of first rectangle + @param rectangle2: (x0, y0, w, h) of second rectangle + """ + rectangle1 = rectangle1[0], rectangle1[1], rectangle1[0] + rectangle1[2], rectangle1[1] + rectangle1[3] + rectangle2 = rectangle2[0], rectangle2[1], rectangle2[0] + rectangle2[2], rectangle2[1] + rectangle2[3] + x_overlap = max(0., min(rectangle1[2], rectangle2[2]) - max(rectangle1[0], rectangle2[0])) + y_overlap = max(0., min(rectangle1[3], rectangle2[3]) - max(rectangle1[1], rectangle2[1])) + return x_overlap * y_overlap + + +def horizontally_flip_bbox(bbox: BoundingBox) -> BoundingBox: + return 1 - (bbox[0] + bbox[2]), bbox[1], bbox[2], bbox[3] + + +def absolute_bbox(relative_bbox: BoundingBox, width: int, height: int) -> Tuple[int, int, int, int]: + bbox = relative_bbox + bbox = bbox[0] * width, bbox[1] * height, (bbox[0] + bbox[2]) * width, (bbox[1] + bbox[3]) * height + return int(bbox[0]), int(bbox[1]), int(bbox[2]), int(bbox[3]) + + +def pad_list(list_: List, pad_element: Any, pad_to_length: int) -> List: + return list_ + [pad_element for _ in range(pad_to_length - len(list_))] + + +def rescale_annotations(annotations: List[Annotation], crop_coordinates: BoundingBox, flip: bool) -> \ + List[Annotation]: + def clamp(x: float): + return max(min(x, 1.), 0.) + + def rescale_bbox(bbox: BoundingBox) -> BoundingBox: + x0 = clamp((bbox[0] - crop_coordinates[0]) / crop_coordinates[2]) + y0 = clamp((bbox[1] - crop_coordinates[1]) / crop_coordinates[3]) + w = min(bbox[2] / crop_coordinates[2], 1 - x0) + h = min(bbox[3] / crop_coordinates[3], 1 - y0) + if flip: + x0 = 1 - (x0 + w) + return x0, y0, w, h + + return [a._replace(bbox=rescale_bbox(a.bbox)) for a in annotations] + + +def filter_annotations(annotations: List[Annotation], crop_coordinates: BoundingBox) -> List: + return [a for a in annotations if intersection_area(a.bbox, crop_coordinates) > 0.0] + + +def additional_parameters_string(annotation: Annotation, short: bool = True) -> str: + sl = slice(1) if short else slice(None) + string = '' + if not (annotation.is_group_of or annotation.is_occluded or annotation.is_depiction or annotation.is_inside): + return string + if annotation.is_group_of: + string += 'group'[sl] + ',' + if annotation.is_occluded: + string += 'occluded'[sl] + ',' + if annotation.is_depiction: + string += 'depiction'[sl] + ',' + if annotation.is_inside: + string += 'inside'[sl] + return '(' + string.strip(",") + ')' + + +def get_plot_font_size(font_size: Optional[int], figure_size: Tuple[int, int]) -> int: + if font_size is None: + font_size = 10 + if max(figure_size) >= 256: + font_size = 12 + if max(figure_size) >= 512: + font_size = 15 + return font_size + + +def get_circle_size(figure_size: Tuple[int, int]) -> int: + circle_size = 2 + if max(figure_size) >= 256: + circle_size = 3 + if max(figure_size) >= 512: + circle_size = 4 + return circle_size + + +def load_object_from_string(object_string: str) -> Any: + """ + Source: https://stackoverflow.com/a/10773699 + """ + module_name, class_name = object_string.rsplit(".", 1) + return getattr(importlib.import_module(module_name), class_name) diff --git a/taming/data/custom.py b/taming/data/custom.py new file mode 100644 index 0000000..33f302a --- /dev/null +++ b/taming/data/custom.py @@ -0,0 +1,38 @@ +import os +import numpy as np +import albumentations +from torch.utils.data import Dataset + +from taming.data.base import ImagePaths, NumpyPaths, ConcatDatasetWithIndex + + +class CustomBase(Dataset): + def __init__(self, *args, **kwargs): + super().__init__() + self.data = None + + def __len__(self): + return len(self.data) + + def __getitem__(self, i): + example = self.data[i] + return example + + + +class CustomTrain(CustomBase): + def __init__(self, size, training_images_list_file): + super().__init__() + with open(training_images_list_file, "r") as f: + paths = f.read().splitlines() + self.data = ImagePaths(paths=paths, size=size, random_crop=False) + + +class CustomTest(CustomBase): + def __init__(self, size, test_images_list_file): + super().__init__() + with open(test_images_list_file, "r") as f: + paths = f.read().splitlines() + self.data = ImagePaths(paths=paths, size=size, random_crop=False) + + diff --git a/taming/data/faceshq.py b/taming/data/faceshq.py new file mode 100644 index 0000000..6912d04 --- /dev/null +++ b/taming/data/faceshq.py @@ -0,0 +1,134 @@ +import os +import numpy as np +import albumentations +from torch.utils.data import Dataset + +from taming.data.base import ImagePaths, NumpyPaths, ConcatDatasetWithIndex + + +class FacesBase(Dataset): + def __init__(self, *args, **kwargs): + super().__init__() + self.data = None + self.keys = None + + def __len__(self): + return len(self.data) + + def __getitem__(self, i): + example = self.data[i] + ex = {} + if self.keys is not None: + for k in self.keys: + ex[k] = example[k] + else: + ex = example + return ex + + +class CelebAHQTrain(FacesBase): + def __init__(self, size, keys=None): + super().__init__() + root = "data/celebahq" + with open("data/celebahqtrain.txt", "r") as f: + relpaths = f.read().splitlines() + paths = [os.path.join(root, relpath) for relpath in relpaths] + self.data = NumpyPaths(paths=paths, size=size, random_crop=False) + self.keys = keys + + +class CelebAHQValidation(FacesBase): + def __init__(self, size, keys=None): + super().__init__() + root = "data/celebahq" + with open("data/celebahqvalidation.txt", "r") as f: + relpaths = f.read().splitlines() + paths = [os.path.join(root, relpath) for relpath in relpaths] + self.data = NumpyPaths(paths=paths, size=size, random_crop=False) + self.keys = keys + + +class FFHQTrain(FacesBase): + def __init__(self, size, keys=None): + super().__init__() + root = "data/ffhq" + with open("data/ffhqtrain.txt", "r") as f: + relpaths = f.read().splitlines() + paths = [os.path.join(root, relpath) for relpath in relpaths] + self.data = ImagePaths(paths=paths, size=size, random_crop=False) + self.keys = keys + + +class FFHQValidation(FacesBase): + def __init__(self, size, keys=None): + super().__init__() + root = "data/ffhq" + with open("data/ffhqvalidation.txt", "r") as f: + relpaths = f.read().splitlines() + paths = [os.path.join(root, relpath) for relpath in relpaths] + self.data = ImagePaths(paths=paths, size=size, random_crop=False) + self.keys = keys + + +class FacesHQTrain(Dataset): + # CelebAHQ [0] + FFHQ [1] + def __init__(self, size, keys=None, crop_size=None, coord=False): + d1 = CelebAHQTrain(size=size, keys=keys) + d2 = FFHQTrain(size=size, keys=keys) + self.data = ConcatDatasetWithIndex([d1, d2]) + self.coord = coord + if crop_size is not None: + self.cropper = albumentations.RandomCrop(height=crop_size,width=crop_size) + if self.coord: + self.cropper = albumentations.Compose([self.cropper], + additional_targets={"coord": "image"}) + + def __len__(self): + return len(self.data) + + def __getitem__(self, i): + ex, y = self.data[i] + if hasattr(self, "cropper"): + if not self.coord: + out = self.cropper(image=ex["image"]) + ex["image"] = out["image"] + else: + h,w,_ = ex["image"].shape + coord = np.arange(h*w).reshape(h,w,1)/(h*w) + out = self.cropper(image=ex["image"], coord=coord) + ex["image"] = out["image"] + ex["coord"] = out["coord"] + ex["class"] = y + return ex + + +class FacesHQValidation(Dataset): + # CelebAHQ [0] + FFHQ [1] + def __init__(self, size, keys=None, crop_size=None, coord=False): + d1 = CelebAHQValidation(size=size, keys=keys) + d2 = FFHQValidation(size=size, keys=keys) + self.data = ConcatDatasetWithIndex([d1, d2]) + self.coord = coord + if crop_size is not None: + self.cropper = albumentations.CenterCrop(height=crop_size,width=crop_size) + if self.coord: + self.cropper = albumentations.Compose([self.cropper], + additional_targets={"coord": "image"}) + + def __len__(self): + return len(self.data) + + def __getitem__(self, i): + ex, y = self.data[i] + if hasattr(self, "cropper"): + if not self.coord: + out = self.cropper(image=ex["image"]) + ex["image"] = out["image"] + else: + h,w,_ = ex["image"].shape + coord = np.arange(h*w).reshape(h,w,1)/(h*w) + out = self.cropper(image=ex["image"], coord=coord) + ex["image"] = out["image"] + ex["coord"] = out["coord"] + ex["class"] = y + return ex diff --git a/taming/data/helper_types.py b/taming/data/helper_types.py new file mode 100644 index 0000000..fb51e30 --- /dev/null +++ b/taming/data/helper_types.py @@ -0,0 +1,49 @@ +from typing import Dict, Tuple, Optional, NamedTuple, Union +from PIL.Image import Image as pil_image +from torch import Tensor + +try: + from typing import Literal +except ImportError: + from typing_extensions import Literal + +Image = Union[Tensor, pil_image] +BoundingBox = Tuple[float, float, float, float] # x0, y0, w, h +CropMethodType = Literal['none', 'random', 'center', 'random-2d'] +SplitType = Literal['train', 'validation', 'test'] + + +class ImageDescription(NamedTuple): + id: int + file_name: str + original_size: Tuple[int, int] # w, h + url: Optional[str] = None + license: Optional[int] = None + coco_url: Optional[str] = None + date_captured: Optional[str] = None + flickr_url: Optional[str] = None + flickr_id: Optional[str] = None + coco_id: Optional[str] = None + + +class Category(NamedTuple): + id: str + super_category: Optional[str] + name: str + + +class Annotation(NamedTuple): + area: float + image_id: str + bbox: BoundingBox + category_no: int + category_id: str + id: Optional[int] = None + source: Optional[str] = None + confidence: Optional[float] = None + is_group_of: Optional[bool] = None + is_truncated: Optional[bool] = None + is_occluded: Optional[bool] = None + is_depiction: Optional[bool] = None + is_inside: Optional[bool] = None + segmentation: Optional[Dict] = None diff --git a/taming/data/image_transforms.py b/taming/data/image_transforms.py new file mode 100644 index 0000000..657ac33 --- /dev/null +++ b/taming/data/image_transforms.py @@ -0,0 +1,132 @@ +import random +import warnings +from typing import Union + +import torch +from torch import Tensor +from torchvision.transforms import RandomCrop, functional as F, CenterCrop, RandomHorizontalFlip, PILToTensor +from torchvision.transforms.functional import _get_image_size as get_image_size + +from taming.data.helper_types import BoundingBox, Image + +pil_to_tensor = PILToTensor() + + +def convert_pil_to_tensor(image: Image) -> Tensor: + with warnings.catch_warnings(): + # to filter PyTorch UserWarning as described here: https://github.com/pytorch/vision/issues/2194 + warnings.simplefilter("ignore") + return pil_to_tensor(image) + + +class RandomCrop1dReturnCoordinates(RandomCrop): + def forward(self, img: Image) -> (BoundingBox, Image): + """ + Additionally to cropping, returns the relative coordinates of the crop bounding box. + Args: + img (PIL Image or Tensor): Image to be cropped. + + Returns: + Bounding box: x0, y0, w, h + PIL Image or Tensor: Cropped image. + + Based on: + torchvision.transforms.RandomCrop, torchvision 1.7.0 + """ + if self.padding is not None: + img = F.pad(img, self.padding, self.fill, self.padding_mode) + + width, height = get_image_size(img) + # pad the width if needed + if self.pad_if_needed and width < self.size[1]: + padding = [self.size[1] - width, 0] + img = F.pad(img, padding, self.fill, self.padding_mode) + # pad the height if needed + if self.pad_if_needed and height < self.size[0]: + padding = [0, self.size[0] - height] + img = F.pad(img, padding, self.fill, self.padding_mode) + + i, j, h, w = self.get_params(img, self.size) + bbox = (j / width, i / height, w / width, h / height) # x0, y0, w, h + return bbox, F.crop(img, i, j, h, w) + + +class Random2dCropReturnCoordinates(torch.nn.Module): + """ + Additionally to cropping, returns the relative coordinates of the crop bounding box. + Args: + img (PIL Image or Tensor): Image to be cropped. + + Returns: + Bounding box: x0, y0, w, h + PIL Image or Tensor: Cropped image. + + Based on: + torchvision.transforms.RandomCrop, torchvision 1.7.0 + """ + + def __init__(self, min_size: int): + super().__init__() + self.min_size = min_size + + def forward(self, img: Image) -> (BoundingBox, Image): + width, height = get_image_size(img) + max_size = min(width, height) + if max_size <= self.min_size: + size = max_size + else: + size = random.randint(self.min_size, max_size) + top = random.randint(0, height - size) + left = random.randint(0, width - size) + bbox = left / width, top / height, size / width, size / height + return bbox, F.crop(img, top, left, size, size) + + +class CenterCropReturnCoordinates(CenterCrop): + @staticmethod + def get_bbox_of_center_crop(width: int, height: int) -> BoundingBox: + if width > height: + w = height / width + h = 1.0 + x0 = 0.5 - w / 2 + y0 = 0. + else: + w = 1.0 + h = width / height + x0 = 0. + y0 = 0.5 - h / 2 + return x0, y0, w, h + + def forward(self, img: Union[Image, Tensor]) -> (BoundingBox, Union[Image, Tensor]): + """ + Additionally to cropping, returns the relative coordinates of the crop bounding box. + Args: + img (PIL Image or Tensor): Image to be cropped. + + Returns: + Bounding box: x0, y0, w, h + PIL Image or Tensor: Cropped image. + Based on: + torchvision.transforms.RandomHorizontalFlip (version 1.7.0) + """ + width, height = get_image_size(img) + return self.get_bbox_of_center_crop(width, height), F.center_crop(img, self.size) + + +class RandomHorizontalFlipReturn(RandomHorizontalFlip): + def forward(self, img: Image) -> (bool, Image): + """ + Additionally to flipping, returns a boolean whether it was flipped or not. + Args: + img (PIL Image or Tensor): Image to be flipped. + + Returns: + flipped: whether the image was flipped or not + PIL Image or Tensor: Randomly flipped image. + + Based on: + torchvision.transforms.RandomHorizontalFlip (version 1.7.0) + """ + if torch.rand(1) < self.p: + return True, F.hflip(img) + return False, img diff --git a/taming/data/imagenet.py b/taming/data/imagenet.py new file mode 100644 index 0000000..9a02ec4 --- /dev/null +++ b/taming/data/imagenet.py @@ -0,0 +1,558 @@ +import os, tarfile, glob, shutil +import yaml +import numpy as np +from tqdm import tqdm +from PIL import Image +import albumentations +from omegaconf import OmegaConf +from torch.utils.data import Dataset + +from taming.data.base import ImagePaths +from taming.util import download, retrieve +import taming.data.utils as bdu + + +def give_synsets_from_indices(indices, path_to_yaml="data/imagenet_idx_to_synset.yaml"): + synsets = [] + with open(path_to_yaml) as f: + di2s = yaml.load(f) + for idx in indices: + synsets.append(str(di2s[idx])) + print("Using {} different synsets for construction of Restriced Imagenet.".format(len(synsets))) + return synsets + + +def str_to_indices(string): + """Expects a string in the format '32-123, 256, 280-321'""" + assert not string.endswith(","), "provided string '{}' ends with a comma, pls remove it".format(string) + subs = string.split(",") + indices = [] + for sub in subs: + subsubs = sub.split("-") + assert len(subsubs) > 0 + if len(subsubs) == 1: + indices.append(int(subsubs[0])) + else: + rang = [j for j in range(int(subsubs[0]), int(subsubs[1]))] + indices.extend(rang) + return sorted(indices) + + +class ImageNetBase(Dataset): + def __init__(self, config=None): + self.config = config or OmegaConf.create() + if not type(self.config)==dict: + self.config = OmegaConf.to_container(self.config) + self._prepare() + self._prepare_synset_to_human() + self._prepare_idx_to_synset() + self._load() + + def __len__(self): + return len(self.data) + + def __getitem__(self, i): + return self.data[i] + + def _prepare(self): + raise NotImplementedError() + + def _filter_relpaths(self, relpaths): + ignore = set([ + "n06596364_9591.JPEG", + ]) + relpaths = [rpath for rpath in relpaths if not rpath.split("/")[-1] in ignore] + if "sub_indices" in self.config: + indices = str_to_indices(self.config["sub_indices"]) + synsets = give_synsets_from_indices(indices, path_to_yaml=self.idx2syn) # returns a list of strings + files = [] + for rpath in relpaths: + syn = rpath.split("/")[0] + if syn in synsets: + files.append(rpath) + return files + else: + return relpaths + + def _prepare_synset_to_human(self): + SIZE = 2655750 + URL = "https://heibox.uni-heidelberg.de/f/9f28e956cd304264bb82/?dl=1" + self.human_dict = os.path.join(self.root, "synset_human.txt") + if (not os.path.exists(self.human_dict) or + not os.path.getsize(self.human_dict)==SIZE): + download(URL, self.human_dict) + + def _prepare_idx_to_synset(self): + URL = "https://heibox.uni-heidelberg.de/f/d835d5b6ceda4d3aa910/?dl=1" + self.idx2syn = os.path.join(self.root, "index_synset.yaml") + if (not os.path.exists(self.idx2syn)): + download(URL, self.idx2syn) + + def _load(self): + with open(self.txt_filelist, "r") as f: + self.relpaths = f.read().splitlines() + l1 = len(self.relpaths) + self.relpaths = self._filter_relpaths(self.relpaths) + print("Removed {} files from filelist during filtering.".format(l1 - len(self.relpaths))) + + self.synsets = [p.split("/")[0] for p in self.relpaths] + self.abspaths = [os.path.join(self.datadir, p) for p in self.relpaths] + + unique_synsets = np.unique(self.synsets) + class_dict = dict((synset, i) for i, synset in enumerate(unique_synsets)) + self.class_labels = [class_dict[s] for s in self.synsets] + + with open(self.human_dict, "r") as f: + human_dict = f.read().splitlines() + human_dict = dict(line.split(maxsplit=1) for line in human_dict) + + self.human_labels = [human_dict[s] for s in self.synsets] + + labels = { + "relpath": np.array(self.relpaths), + "synsets": np.array(self.synsets), + "class_label": np.array(self.class_labels), + "human_label": np.array(self.human_labels), + } + self.data = ImagePaths(self.abspaths, + labels=labels, + size=retrieve(self.config, "size", default=0), + random_crop=self.random_crop) + + +class ImageNetTrain(ImageNetBase): + NAME = "ILSVRC2012_train" + URL = "http://www.image-net.org/challenges/LSVRC/2012/" + AT_HASH = "a306397ccf9c2ead27155983c254227c0fd938e2" + FILES = [ + "ILSVRC2012_img_train.tar", + ] + SIZES = [ + 147897477120, + ] + + def _prepare(self): + self.random_crop = retrieve(self.config, "ImageNetTrain/random_crop", + default=True) + cachedir = os.environ.get("XDG_CACHE_HOME", os.path.expanduser("~/.cache")) + self.root = os.path.join(cachedir, "autoencoders/data", self.NAME) + self.datadir = os.path.join(self.root, "data") + self.txt_filelist = os.path.join(self.root, "filelist.txt") + self.expected_length = 1281167 + if not bdu.is_prepared(self.root): + # prep + print("Preparing dataset {} in {}".format(self.NAME, self.root)) + + datadir = self.datadir + if not os.path.exists(datadir): + path = os.path.join(self.root, self.FILES[0]) + if not os.path.exists(path) or not os.path.getsize(path)==self.SIZES[0]: + import academictorrents as at + atpath = at.get(self.AT_HASH, datastore=self.root) + assert atpath == path + + print("Extracting {} to {}".format(path, datadir)) + os.makedirs(datadir, exist_ok=True) + with tarfile.open(path, "r:") as tar: + tar.extractall(path=datadir) + + print("Extracting sub-tars.") + subpaths = sorted(glob.glob(os.path.join(datadir, "*.tar"))) + for subpath in tqdm(subpaths): + subdir = subpath[:-len(".tar")] + os.makedirs(subdir, exist_ok=True) + with tarfile.open(subpath, "r:") as tar: + tar.extractall(path=subdir) + + + filelist = glob.glob(os.path.join(datadir, "**", "*.JPEG")) + filelist = [os.path.relpath(p, start=datadir) for p in filelist] + filelist = sorted(filelist) + filelist = "\n".join(filelist)+"\n" + with open(self.txt_filelist, "w") as f: + f.write(filelist) + + bdu.mark_prepared(self.root) + + +class ImageNetValidation(ImageNetBase): + NAME = "ILSVRC2012_validation" + URL = "http://www.image-net.org/challenges/LSVRC/2012/" + AT_HASH = "5d6d0df7ed81efd49ca99ea4737e0ae5e3a5f2e5" + VS_URL = "https://heibox.uni-heidelberg.de/f/3e0f6e9c624e45f2bd73/?dl=1" + FILES = [ + "ILSVRC2012_img_val.tar", + "validation_synset.txt", + ] + SIZES = [ + 6744924160, + 1950000, + ] + + def _prepare(self): + self.random_crop = retrieve(self.config, "ImageNetValidation/random_crop", + default=False) + cachedir = os.environ.get("XDG_CACHE_HOME", os.path.expanduser("~/.cache")) + self.root = os.path.join(cachedir, "autoencoders/data", self.NAME) + self.datadir = os.path.join(self.root, "data") + self.txt_filelist = os.path.join(self.root, "filelist.txt") + self.expected_length = 50000 + if not bdu.is_prepared(self.root): + # prep + print("Preparing dataset {} in {}".format(self.NAME, self.root)) + + datadir = self.datadir + if not os.path.exists(datadir): + path = os.path.join(self.root, self.FILES[0]) + if not os.path.exists(path) or not os.path.getsize(path)==self.SIZES[0]: + import academictorrents as at + atpath = at.get(self.AT_HASH, datastore=self.root) + assert atpath == path + + print("Extracting {} to {}".format(path, datadir)) + os.makedirs(datadir, exist_ok=True) + with tarfile.open(path, "r:") as tar: + tar.extractall(path=datadir) + + vspath = os.path.join(self.root, self.FILES[1]) + if not os.path.exists(vspath) or not os.path.getsize(vspath)==self.SIZES[1]: + download(self.VS_URL, vspath) + + with open(vspath, "r") as f: + synset_dict = f.read().splitlines() + synset_dict = dict(line.split() for line in synset_dict) + + print("Reorganizing into synset folders") + synsets = np.unique(list(synset_dict.values())) + for s in synsets: + os.makedirs(os.path.join(datadir, s), exist_ok=True) + for k, v in synset_dict.items(): + src = os.path.join(datadir, k) + dst = os.path.join(datadir, v) + shutil.move(src, dst) + + filelist = glob.glob(os.path.join(datadir, "**", "*.JPEG")) + filelist = [os.path.relpath(p, start=datadir) for p in filelist] + filelist = sorted(filelist) + filelist = "\n".join(filelist)+"\n" + with open(self.txt_filelist, "w") as f: + f.write(filelist) + + bdu.mark_prepared(self.root) + + +def get_preprocessor(size=None, random_crop=False, additional_targets=None, + crop_size=None): + if size is not None and size > 0: + transforms = list() + rescaler = albumentations.SmallestMaxSize(max_size = size) + transforms.append(rescaler) + if not random_crop: + cropper = albumentations.CenterCrop(height=size,width=size) + transforms.append(cropper) + else: + cropper = albumentations.RandomCrop(height=size,width=size) + transforms.append(cropper) + flipper = albumentations.HorizontalFlip() + transforms.append(flipper) + preprocessor = albumentations.Compose(transforms, + additional_targets=additional_targets) + elif crop_size is not None and crop_size > 0: + if not random_crop: + cropper = albumentations.CenterCrop(height=crop_size,width=crop_size) + else: + cropper = albumentations.RandomCrop(height=crop_size,width=crop_size) + transforms = [cropper] + preprocessor = albumentations.Compose(transforms, + additional_targets=additional_targets) + else: + preprocessor = lambda **kwargs: kwargs + return preprocessor + + +def rgba_to_depth(x): + assert x.dtype == np.uint8 + assert len(x.shape) == 3 and x.shape[2] == 4 + y = x.copy() + y.dtype = np.float32 + y = y.reshape(x.shape[:2]) + return np.ascontiguousarray(y) + + +class BaseWithDepth(Dataset): + DEFAULT_DEPTH_ROOT="data/imagenet_depth" + + def __init__(self, config=None, size=None, random_crop=False, + crop_size=None, root=None): + self.config = config + self.base_dset = self.get_base_dset() + self.preprocessor = get_preprocessor( + size=size, + crop_size=crop_size, + random_crop=random_crop, + additional_targets={"depth": "image"}) + self.crop_size = crop_size + if self.crop_size is not None: + self.rescaler = albumentations.Compose( + [albumentations.SmallestMaxSize(max_size = self.crop_size)], + additional_targets={"depth": "image"}) + if root is not None: + self.DEFAULT_DEPTH_ROOT = root + + def __len__(self): + return len(self.base_dset) + + def preprocess_depth(self, path): + rgba = np.array(Image.open(path)) + depth = rgba_to_depth(rgba) + depth = (depth - depth.min())/max(1e-8, depth.max()-depth.min()) + depth = 2.0*depth-1.0 + return depth + + def __getitem__(self, i): + e = self.base_dset[i] + e["depth"] = self.preprocess_depth(self.get_depth_path(e)) + # up if necessary + h,w,c = e["image"].shape + if self.crop_size and min(h,w) < self.crop_size: + # have to upscale to be able to crop - this just uses bilinear + out = self.rescaler(image=e["image"], depth=e["depth"]) + e["image"] = out["image"] + e["depth"] = out["depth"] + transformed = self.preprocessor(image=e["image"], depth=e["depth"]) + e["image"] = transformed["image"] + e["depth"] = transformed["depth"] + return e + + +class ImageNetTrainWithDepth(BaseWithDepth): + # default to random_crop=True + def __init__(self, random_crop=True, sub_indices=None, **kwargs): + self.sub_indices = sub_indices + super().__init__(random_crop=random_crop, **kwargs) + + def get_base_dset(self): + if self.sub_indices is None: + return ImageNetTrain() + else: + return ImageNetTrain({"sub_indices": self.sub_indices}) + + def get_depth_path(self, e): + fid = os.path.splitext(e["relpath"])[0]+".png" + fid = os.path.join(self.DEFAULT_DEPTH_ROOT, "train", fid) + return fid + + +class ImageNetValidationWithDepth(BaseWithDepth): + def __init__(self, sub_indices=None, **kwargs): + self.sub_indices = sub_indices + super().__init__(**kwargs) + + def get_base_dset(self): + if self.sub_indices is None: + return ImageNetValidation() + else: + return ImageNetValidation({"sub_indices": self.sub_indices}) + + def get_depth_path(self, e): + fid = os.path.splitext(e["relpath"])[0]+".png" + fid = os.path.join(self.DEFAULT_DEPTH_ROOT, "val", fid) + return fid + + +class RINTrainWithDepth(ImageNetTrainWithDepth): + def __init__(self, config=None, size=None, random_crop=True, crop_size=None): + sub_indices = "30-32, 33-37, 151-268, 281-285, 80-100, 365-382, 389-397, 118-121, 300-319" + super().__init__(config=config, size=size, random_crop=random_crop, + sub_indices=sub_indices, crop_size=crop_size) + + +class RINValidationWithDepth(ImageNetValidationWithDepth): + def __init__(self, config=None, size=None, random_crop=False, crop_size=None): + sub_indices = "30-32, 33-37, 151-268, 281-285, 80-100, 365-382, 389-397, 118-121, 300-319" + super().__init__(config=config, size=size, random_crop=random_crop, + sub_indices=sub_indices, crop_size=crop_size) + + +class DRINExamples(Dataset): + def __init__(self): + self.preprocessor = get_preprocessor(size=256, additional_targets={"depth": "image"}) + with open("data/drin_examples.txt", "r") as f: + relpaths = f.read().splitlines() + self.image_paths = [os.path.join("data/drin_images", + relpath) for relpath in relpaths] + self.depth_paths = [os.path.join("data/drin_depth", + relpath.replace(".JPEG", ".png")) for relpath in relpaths] + + def __len__(self): + return len(self.image_paths) + + def preprocess_image(self, image_path): + image = Image.open(image_path) + if not image.mode == "RGB": + image = image.convert("RGB") + image = np.array(image).astype(np.uint8) + image = self.preprocessor(image=image)["image"] + image = (image/127.5 - 1.0).astype(np.float32) + return image + + def preprocess_depth(self, path): + rgba = np.array(Image.open(path)) + depth = rgba_to_depth(rgba) + depth = (depth - depth.min())/max(1e-8, depth.max()-depth.min()) + depth = 2.0*depth-1.0 + return depth + + def __getitem__(self, i): + e = dict() + e["image"] = self.preprocess_image(self.image_paths[i]) + e["depth"] = self.preprocess_depth(self.depth_paths[i]) + transformed = self.preprocessor(image=e["image"], depth=e["depth"]) + e["image"] = transformed["image"] + e["depth"] = transformed["depth"] + return e + + +def imscale(x, factor, keepshapes=False, keepmode="bicubic"): + if factor is None or factor==1: + return x + + dtype = x.dtype + assert dtype in [np.float32, np.float64] + assert x.min() >= -1 + assert x.max() <= 1 + + keepmode = {"nearest": Image.NEAREST, "bilinear": Image.BILINEAR, + "bicubic": Image.BICUBIC}[keepmode] + + lr = (x+1.0)*127.5 + lr = lr.clip(0,255).astype(np.uint8) + lr = Image.fromarray(lr) + + h, w, _ = x.shape + nh = h//factor + nw = w//factor + assert nh > 0 and nw > 0, (nh, nw) + + lr = lr.resize((nw,nh), Image.BICUBIC) + if keepshapes: + lr = lr.resize((w,h), keepmode) + lr = np.array(lr)/127.5-1.0 + lr = lr.astype(dtype) + + return lr + + +class ImageNetScale(Dataset): + def __init__(self, size=None, crop_size=None, random_crop=False, + up_factor=None, hr_factor=None, keep_mode="bicubic"): + self.base = self.get_base() + + self.size = size + self.crop_size = crop_size if crop_size is not None else self.size + self.random_crop = random_crop + self.up_factor = up_factor + self.hr_factor = hr_factor + self.keep_mode = keep_mode + + transforms = list() + + if self.size is not None and self.size > 0: + rescaler = albumentations.SmallestMaxSize(max_size = self.size) + self.rescaler = rescaler + transforms.append(rescaler) + + if self.crop_size is not None and self.crop_size > 0: + if len(transforms) == 0: + self.rescaler = albumentations.SmallestMaxSize(max_size = self.crop_size) + + if not self.random_crop: + cropper = albumentations.CenterCrop(height=self.crop_size,width=self.crop_size) + else: + cropper = albumentations.RandomCrop(height=self.crop_size,width=self.crop_size) + transforms.append(cropper) + + if len(transforms) > 0: + if self.up_factor is not None: + additional_targets = {"lr": "image"} + else: + additional_targets = None + self.preprocessor = albumentations.Compose(transforms, + additional_targets=additional_targets) + else: + self.preprocessor = lambda **kwargs: kwargs + + def __len__(self): + return len(self.base) + + def __getitem__(self, i): + example = self.base[i] + image = example["image"] + # adjust resolution + image = imscale(image, self.hr_factor, keepshapes=False) + h,w,c = image.shape + if self.crop_size and min(h,w) < self.crop_size: + # have to upscale to be able to crop - this just uses bilinear + image = self.rescaler(image=image)["image"] + if self.up_factor is None: + image = self.preprocessor(image=image)["image"] + example["image"] = image + else: + lr = imscale(image, self.up_factor, keepshapes=True, + keepmode=self.keep_mode) + + out = self.preprocessor(image=image, lr=lr) + example["image"] = out["image"] + example["lr"] = out["lr"] + + return example + +class ImageNetScaleTrain(ImageNetScale): + def __init__(self, random_crop=True, **kwargs): + super().__init__(random_crop=random_crop, **kwargs) + + def get_base(self): + return ImageNetTrain() + +class ImageNetScaleValidation(ImageNetScale): + def get_base(self): + return ImageNetValidation() + + +from skimage.feature import canny +from skimage.color import rgb2gray + + +class ImageNetEdges(ImageNetScale): + def __init__(self, up_factor=1, **kwargs): + super().__init__(up_factor=1, **kwargs) + + def __getitem__(self, i): + example = self.base[i] + image = example["image"] + h,w,c = image.shape + if self.crop_size and min(h,w) < self.crop_size: + # have to upscale to be able to crop - this just uses bilinear + image = self.rescaler(image=image)["image"] + + lr = canny(rgb2gray(image), sigma=2) + lr = lr.astype(np.float32) + lr = lr[:,:,None][:,:,[0,0,0]] + + out = self.preprocessor(image=image, lr=lr) + example["image"] = out["image"] + example["lr"] = out["lr"] + + return example + + +class ImageNetEdgesTrain(ImageNetEdges): + def __init__(self, random_crop=True, **kwargs): + super().__init__(random_crop=random_crop, **kwargs) + + def get_base(self): + return ImageNetTrain() + +class ImageNetEdgesValidation(ImageNetEdges): + def get_base(self): + return ImageNetValidation() diff --git a/taming/data/open_images_helper.py b/taming/data/open_images_helper.py new file mode 100644 index 0000000..8feb7c6 --- /dev/null +++ b/taming/data/open_images_helper.py @@ -0,0 +1,379 @@ +open_images_unify_categories_for_coco = { + '/m/03bt1vf': '/m/01g317', + '/m/04yx4': '/m/01g317', + '/m/05r655': '/m/01g317', + '/m/01bl7v': '/m/01g317', + '/m/0cnyhnx': '/m/01xq0k1', + '/m/01226z': '/m/018xm', + '/m/05ctyq': '/m/018xm', + '/m/058qzx': '/m/04ctx', + '/m/06pcq': '/m/0l515', + '/m/03m3pdh': '/m/02crq1', + '/m/046dlr': '/m/01x3z', + '/m/0h8mzrc': '/m/01x3z', +} + + +top_300_classes_plus_coco_compatibility = [ + ('Man', 1060962), + ('Clothing', 986610), + ('Tree', 748162), + ('Woman', 611896), + ('Person', 610294), + ('Human face', 442948), + ('Girl', 175399), + ('Building', 162147), + ('Car', 159135), + ('Plant', 155704), + ('Human body', 137073), + ('Flower', 133128), + ('Window', 127485), + ('Human arm', 118380), + ('House', 114365), + ('Wheel', 111684), + ('Suit', 99054), + ('Human hair', 98089), + ('Human head', 92763), + ('Chair', 88624), + ('Boy', 79849), + ('Table', 73699), + ('Jeans', 57200), + ('Tire', 55725), + ('Skyscraper', 53321), + ('Food', 52400), + ('Footwear', 50335), + ('Dress', 50236), + ('Human leg', 47124), + ('Toy', 46636), + ('Tower', 45605), + ('Boat', 43486), + ('Land vehicle', 40541), + ('Bicycle wheel', 34646), + ('Palm tree', 33729), + ('Fashion accessory', 32914), + ('Glasses', 31940), + ('Bicycle', 31409), + ('Furniture', 30656), + ('Sculpture', 29643), + ('Bottle', 27558), + ('Dog', 26980), + ('Snack', 26796), + ('Human hand', 26664), + ('Bird', 25791), + ('Book', 25415), + ('Guitar', 24386), + ('Jacket', 23998), + ('Poster', 22192), + ('Dessert', 21284), + ('Baked goods', 20657), + ('Drink', 19754), + ('Flag', 18588), + ('Houseplant', 18205), + ('Tableware', 17613), + ('Airplane', 17218), + ('Door', 17195), + ('Sports uniform', 17068), + ('Shelf', 16865), + ('Drum', 16612), + ('Vehicle', 16542), + ('Microphone', 15269), + ('Street light', 14957), + ('Cat', 14879), + ('Fruit', 13684), + ('Fast food', 13536), + ('Animal', 12932), + ('Vegetable', 12534), + ('Train', 12358), + ('Horse', 11948), + ('Flowerpot', 11728), + ('Motorcycle', 11621), + ('Fish', 11517), + ('Desk', 11405), + ('Helmet', 10996), + ('Truck', 10915), + ('Bus', 10695), + ('Hat', 10532), + ('Auto part', 10488), + ('Musical instrument', 10303), + ('Sunglasses', 10207), + ('Picture frame', 10096), + ('Sports equipment', 10015), + ('Shorts', 9999), + ('Wine glass', 9632), + ('Duck', 9242), + ('Wine', 9032), + ('Rose', 8781), + ('Tie', 8693), + ('Butterfly', 8436), + ('Beer', 7978), + ('Cabinetry', 7956), + ('Laptop', 7907), + ('Insect', 7497), + ('Goggles', 7363), + ('Shirt', 7098), + ('Dairy Product', 7021), + ('Marine invertebrates', 7014), + ('Cattle', 7006), + ('Trousers', 6903), + ('Van', 6843), + ('Billboard', 6777), + ('Balloon', 6367), + ('Human nose', 6103), + ('Tent', 6073), + ('Camera', 6014), + ('Doll', 6002), + ('Coat', 5951), + ('Mobile phone', 5758), + ('Swimwear', 5729), + ('Strawberry', 5691), + ('Stairs', 5643), + ('Goose', 5599), + ('Umbrella', 5536), + ('Cake', 5508), + ('Sun hat', 5475), + ('Bench', 5310), + ('Bookcase', 5163), + ('Bee', 5140), + ('Computer monitor', 5078), + ('Hiking equipment', 4983), + ('Office building', 4981), + ('Coffee cup', 4748), + ('Curtain', 4685), + ('Plate', 4651), + ('Box', 4621), + ('Tomato', 4595), + ('Coffee table', 4529), + ('Office supplies', 4473), + ('Maple', 4416), + ('Muffin', 4365), + ('Cocktail', 4234), + ('Castle', 4197), + ('Couch', 4134), + ('Pumpkin', 3983), + ('Computer keyboard', 3960), + ('Human mouth', 3926), + ('Christmas tree', 3893), + ('Mushroom', 3883), + ('Swimming pool', 3809), + ('Pastry', 3799), + ('Lavender (Plant)', 3769), + ('Football helmet', 3732), + ('Bread', 3648), + ('Traffic sign', 3628), + ('Common sunflower', 3597), + ('Television', 3550), + ('Bed', 3525), + ('Cookie', 3485), + ('Fountain', 3484), + ('Paddle', 3447), + ('Bicycle helmet', 3429), + ('Porch', 3420), + ('Deer', 3387), + ('Fedora', 3339), + ('Canoe', 3338), + ('Carnivore', 3266), + ('Bowl', 3202), + ('Human eye', 3166), + ('Ball', 3118), + ('Pillow', 3077), + ('Salad', 3061), + ('Beetle', 3060), + ('Orange', 3050), + ('Drawer', 2958), + ('Platter', 2937), + ('Elephant', 2921), + ('Seafood', 2921), + ('Monkey', 2915), + ('Countertop', 2879), + ('Watercraft', 2831), + ('Helicopter', 2805), + ('Kitchen appliance', 2797), + ('Personal flotation device', 2781), + ('Swan', 2739), + ('Lamp', 2711), + ('Boot', 2695), + ('Bronze sculpture', 2693), + ('Chicken', 2677), + ('Taxi', 2643), + ('Juice', 2615), + ('Cowboy hat', 2604), + ('Apple', 2600), + ('Tin can', 2590), + ('Necklace', 2564), + ('Ice cream', 2560), + ('Human beard', 2539), + ('Coin', 2536), + ('Candle', 2515), + ('Cart', 2512), + ('High heels', 2441), + ('Weapon', 2433), + ('Handbag', 2406), + ('Penguin', 2396), + ('Rifle', 2352), + ('Violin', 2336), + ('Skull', 2304), + ('Lantern', 2285), + ('Scarf', 2269), + ('Saucer', 2225), + ('Sheep', 2215), + ('Vase', 2189), + ('Lily', 2180), + ('Mug', 2154), + ('Parrot', 2140), + ('Human ear', 2137), + ('Sandal', 2115), + ('Lizard', 2100), + ('Kitchen & dining room table', 2063), + ('Spider', 1977), + ('Coffee', 1974), + ('Goat', 1926), + ('Squirrel', 1922), + ('Cello', 1913), + ('Sushi', 1881), + ('Tortoise', 1876), + ('Pizza', 1870), + ('Studio couch', 1864), + ('Barrel', 1862), + ('Cosmetics', 1841), + ('Moths and butterflies', 1841), + ('Convenience store', 1817), + ('Watch', 1792), + ('Home appliance', 1786), + ('Harbor seal', 1780), + ('Luggage and bags', 1756), + ('Vehicle registration plate', 1754), + ('Shrimp', 1751), + ('Jellyfish', 1730), + ('French fries', 1723), + ('Egg (Food)', 1698), + ('Football', 1697), + ('Musical keyboard', 1683), + ('Falcon', 1674), + ('Candy', 1660), + ('Medical equipment', 1654), + ('Eagle', 1651), + ('Dinosaur', 1634), + ('Surfboard', 1630), + ('Tank', 1628), + ('Grape', 1624), + ('Lion', 1624), + ('Owl', 1622), + ('Ski', 1613), + ('Waste container', 1606), + ('Frog', 1591), + ('Sparrow', 1585), + ('Rabbit', 1581), + ('Pen', 1546), + ('Sea lion', 1537), + ('Spoon', 1521), + ('Sink', 1512), + ('Teddy bear', 1507), + ('Bull', 1495), + ('Sofa bed', 1490), + ('Dragonfly', 1479), + ('Brassiere', 1478), + ('Chest of drawers', 1472), + ('Aircraft', 1466), + ('Human foot', 1463), + ('Pig', 1455), + ('Fork', 1454), + ('Antelope', 1438), + ('Tripod', 1427), + ('Tool', 1424), + ('Cheese', 1422), + ('Lemon', 1397), + ('Hamburger', 1393), + ('Dolphin', 1390), + ('Mirror', 1390), + ('Marine mammal', 1387), + ('Giraffe', 1385), + ('Snake', 1368), + ('Gondola', 1364), + ('Wheelchair', 1360), + ('Piano', 1358), + ('Cupboard', 1348), + ('Banana', 1345), + ('Trumpet', 1335), + ('Lighthouse', 1333), + ('Invertebrate', 1317), + ('Carrot', 1268), + ('Sock', 1260), + ('Tiger', 1241), + ('Camel', 1224), + ('Parachute', 1224), + ('Bathroom accessory', 1223), + ('Earrings', 1221), + ('Headphones', 1218), + ('Skirt', 1198), + ('Skateboard', 1190), + ('Sandwich', 1148), + ('Saxophone', 1141), + ('Goldfish', 1136), + ('Stool', 1104), + ('Traffic light', 1097), + ('Shellfish', 1081), + ('Backpack', 1079), + ('Sea turtle', 1078), + ('Cucumber', 1075), + ('Tea', 1051), + ('Toilet', 1047), + ('Roller skates', 1040), + ('Mule', 1039), + ('Bust', 1031), + ('Broccoli', 1030), + ('Crab', 1020), + ('Oyster', 1019), + ('Cannon', 1012), + ('Zebra', 1012), + ('French horn', 1008), + ('Grapefruit', 998), + ('Whiteboard', 997), + ('Zucchini', 997), + ('Crocodile', 992), + + ('Clock', 960), + ('Wall clock', 958), + + ('Doughnut', 869), + ('Snail', 868), + + ('Baseball glove', 859), + + ('Panda', 830), + ('Tennis racket', 830), + + ('Pear', 652), + + ('Bagel', 617), + ('Oven', 616), + ('Ladybug', 615), + ('Shark', 615), + ('Polar bear', 614), + ('Ostrich', 609), + + ('Hot dog', 473), + ('Microwave oven', 467), + ('Fire hydrant', 20), + ('Stop sign', 20), + ('Parking meter', 20), + ('Bear', 20), + ('Flying disc', 20), + ('Snowboard', 20), + ('Tennis ball', 20), + ('Kite', 20), + ('Baseball bat', 20), + ('Kitchen knife', 20), + ('Knife', 20), + ('Submarine sandwich', 20), + ('Computer mouse', 20), + ('Remote control', 20), + ('Toaster', 20), + ('Sink', 20), + ('Refrigerator', 20), + ('Alarm clock', 20), + ('Wall clock', 20), + ('Scissors', 20), + ('Hair dryer', 20), + ('Toothbrush', 20), + ('Suitcase', 20) +] diff --git a/taming/data/sflckr.py b/taming/data/sflckr.py new file mode 100644 index 0000000..91101be --- /dev/null +++ b/taming/data/sflckr.py @@ -0,0 +1,91 @@ +import os +import numpy as np +import cv2 +import albumentations +from PIL import Image +from torch.utils.data import Dataset + + +class SegmentationBase(Dataset): + def __init__(self, + data_csv, data_root, segmentation_root, + size=None, random_crop=False, interpolation="bicubic", + n_labels=182, shift_segmentation=False, + ): + self.n_labels = n_labels + self.shift_segmentation = shift_segmentation + self.data_csv = data_csv + self.data_root = data_root + self.segmentation_root = segmentation_root + with open(self.data_csv, "r") as f: + self.image_paths = f.read().splitlines() + self._length = len(self.image_paths) + self.labels = { + "relative_file_path_": [l for l in self.image_paths], + "file_path_": [os.path.join(self.data_root, l) + for l in self.image_paths], + "segmentation_path_": [os.path.join(self.segmentation_root, l.replace(".jpg", ".png")) + for l in self.image_paths] + } + + size = None if size is not None and size<=0 else size + self.size = size + if self.size is not None: + self.interpolation = interpolation + self.interpolation = { + "nearest": cv2.INTER_NEAREST, + "bilinear": cv2.INTER_LINEAR, + "bicubic": cv2.INTER_CUBIC, + "area": cv2.INTER_AREA, + "lanczos": cv2.INTER_LANCZOS4}[self.interpolation] + self.image_rescaler = albumentations.SmallestMaxSize(max_size=self.size, + interpolation=self.interpolation) + self.segmentation_rescaler = albumentations.SmallestMaxSize(max_size=self.size, + interpolation=cv2.INTER_NEAREST) + self.center_crop = not random_crop + if self.center_crop: + self.cropper = albumentations.CenterCrop(height=self.size, width=self.size) + else: + self.cropper = albumentations.RandomCrop(height=self.size, width=self.size) + self.preprocessor = self.cropper + + def __len__(self): + return self._length + + def __getitem__(self, i): + example = dict((k, self.labels[k][i]) for k in self.labels) + image = Image.open(example["file_path_"]) + if not image.mode == "RGB": + image = image.convert("RGB") + image = np.array(image).astype(np.uint8) + if self.size is not None: + image = self.image_rescaler(image=image)["image"] + segmentation = Image.open(example["segmentation_path_"]) + assert segmentation.mode == "L", segmentation.mode + segmentation = np.array(segmentation).astype(np.uint8) + if self.shift_segmentation: + # used to support segmentations containing unlabeled==255 label + segmentation = segmentation+1 + if self.size is not None: + segmentation = self.segmentation_rescaler(image=segmentation)["image"] + if self.size is not None: + processed = self.preprocessor(image=image, + mask=segmentation + ) + else: + processed = {"image": image, + "mask": segmentation + } + example["image"] = (processed["image"]/127.5 - 1.0).astype(np.float32) + segmentation = processed["mask"] + onehot = np.eye(self.n_labels)[segmentation] + example["segmentation"] = onehot + return example + + +class Examples(SegmentationBase): + def __init__(self, size=None, random_crop=False, interpolation="bicubic"): + super().__init__(data_csv="data/sflckr_examples.txt", + data_root="data/sflckr_images", + segmentation_root="data/sflckr_segmentations", + size=size, random_crop=random_crop, interpolation=interpolation) diff --git a/taming/data/utils.py b/taming/data/utils.py new file mode 100644 index 0000000..2b3c3d5 --- /dev/null +++ b/taming/data/utils.py @@ -0,0 +1,169 @@ +import collections +import os +import tarfile +import urllib +import zipfile +from pathlib import Path + +import numpy as np +import torch +from taming.data.helper_types import Annotation +from torch._six import string_classes +from torch.utils.data._utils.collate import np_str_obj_array_pattern, default_collate_err_msg_format +from tqdm import tqdm + + +def unpack(path): + if path.endswith("tar.gz"): + with tarfile.open(path, "r:gz") as tar: + tar.extractall(path=os.path.split(path)[0]) + elif path.endswith("tar"): + with tarfile.open(path, "r:") as tar: + tar.extractall(path=os.path.split(path)[0]) + elif path.endswith("zip"): + with zipfile.ZipFile(path, "r") as f: + f.extractall(path=os.path.split(path)[0]) + else: + raise NotImplementedError( + "Unknown file extension: {}".format(os.path.splitext(path)[1]) + ) + + +def reporthook(bar): + """tqdm progress bar for downloads.""" + + def hook(b=1, bsize=1, tsize=None): + if tsize is not None: + bar.total = tsize + bar.update(b * bsize - bar.n) + + return hook + + +def get_root(name): + base = "data/" + root = os.path.join(base, name) + os.makedirs(root, exist_ok=True) + return root + + +def is_prepared(root): + return Path(root).joinpath(".ready").exists() + + +def mark_prepared(root): + Path(root).joinpath(".ready").touch() + + +def prompt_download(file_, source, target_dir, content_dir=None): + targetpath = os.path.join(target_dir, file_) + while not os.path.exists(targetpath): + if content_dir is not None and os.path.exists( + os.path.join(target_dir, content_dir) + ): + break + print( + "Please download '{}' from '{}' to '{}'.".format(file_, source, targetpath) + ) + if content_dir is not None: + print( + "Or place its content into '{}'.".format( + os.path.join(target_dir, content_dir) + ) + ) + input("Press Enter when done...") + return targetpath + + +def download_url(file_, url, target_dir): + targetpath = os.path.join(target_dir, file_) + os.makedirs(target_dir, exist_ok=True) + with tqdm( + unit="B", unit_scale=True, unit_divisor=1024, miniters=1, desc=file_ + ) as bar: + urllib.request.urlretrieve(url, targetpath, reporthook=reporthook(bar)) + return targetpath + + +def download_urls(urls, target_dir): + paths = dict() + for fname, url in urls.items(): + outpath = download_url(fname, url, target_dir) + paths[fname] = outpath + return paths + + +def quadratic_crop(x, bbox, alpha=1.0): + """bbox is xmin, ymin, xmax, ymax""" + im_h, im_w = x.shape[:2] + bbox = np.array(bbox, dtype=np.float32) + bbox = np.clip(bbox, 0, max(im_h, im_w)) + center = 0.5 * (bbox[0] + bbox[2]), 0.5 * (bbox[1] + bbox[3]) + w = bbox[2] - bbox[0] + h = bbox[3] - bbox[1] + l = int(alpha * max(w, h)) + l = max(l, 2) + + required_padding = -1 * min( + center[0] - l, center[1] - l, im_w - (center[0] + l), im_h - (center[1] + l) + ) + required_padding = int(np.ceil(required_padding)) + if required_padding > 0: + padding = [ + [required_padding, required_padding], + [required_padding, required_padding], + ] + padding += [[0, 0]] * (len(x.shape) - 2) + x = np.pad(x, padding, "reflect") + center = center[0] + required_padding, center[1] + required_padding + xmin = int(center[0] - l / 2) + ymin = int(center[1] - l / 2) + return np.array(x[ymin : ymin + l, xmin : xmin + l, ...]) + + +def custom_collate(batch): + r"""source: pytorch 1.9.0, only one modification to original code """ + + elem = batch[0] + elem_type = type(elem) + if isinstance(elem, torch.Tensor): + out = None + if torch.utils.data.get_worker_info() is not None: + # If we're in a background process, concatenate directly into a + # shared memory tensor to avoid an extra copy + numel = sum([x.numel() for x in batch]) + storage = elem.storage()._new_shared(numel) + out = elem.new(storage) + return torch.stack(batch, 0, out=out) + elif elem_type.__module__ == 'numpy' and elem_type.__name__ != 'str_' \ + and elem_type.__name__ != 'string_': + if elem_type.__name__ == 'ndarray' or elem_type.__name__ == 'memmap': + # array of string classes and object + if np_str_obj_array_pattern.search(elem.dtype.str) is not None: + raise TypeError(default_collate_err_msg_format.format(elem.dtype)) + + return custom_collate([torch.as_tensor(b) for b in batch]) + elif elem.shape == (): # scalars + return torch.as_tensor(batch) + elif isinstance(elem, float): + return torch.tensor(batch, dtype=torch.float64) + elif isinstance(elem, int): + return torch.tensor(batch) + elif isinstance(elem, string_classes): + return batch + elif isinstance(elem, collections.abc.Mapping): + return {key: custom_collate([d[key] for d in batch]) for key in elem} + elif isinstance(elem, tuple) and hasattr(elem, '_fields'): # namedtuple + return elem_type(*(custom_collate(samples) for samples in zip(*batch))) + if isinstance(elem, collections.abc.Sequence) and isinstance(elem[0], Annotation): # added + return batch # added + elif isinstance(elem, collections.abc.Sequence): + # check to make sure that the elements in batch have consistent size + it = iter(batch) + elem_size = len(next(it)) + if not all(len(elem) == elem_size for elem in it): + raise RuntimeError('each element in list of batch should be of equal size') + transposed = zip(*batch) + return [custom_collate(samples) for samples in transposed] + + raise TypeError(default_collate_err_msg_format.format(elem_type)) diff --git a/taming/lr_scheduler.py b/taming/lr_scheduler.py new file mode 100644 index 0000000..e598ed1 --- /dev/null +++ b/taming/lr_scheduler.py @@ -0,0 +1,34 @@ +import numpy as np + + +class LambdaWarmUpCosineScheduler: + """ + note: use with a base_lr of 1.0 + """ + def __init__(self, warm_up_steps, lr_min, lr_max, lr_start, max_decay_steps, verbosity_interval=0): + self.lr_warm_up_steps = warm_up_steps + self.lr_start = lr_start + self.lr_min = lr_min + self.lr_max = lr_max + self.lr_max_decay_steps = max_decay_steps + self.last_lr = 0. + self.verbosity_interval = verbosity_interval + + def schedule(self, n): + if self.verbosity_interval > 0: + if n % self.verbosity_interval == 0: print(f"current step: {n}, recent lr-multiplier: {self.last_lr}") + if n < self.lr_warm_up_steps: + lr = (self.lr_max - self.lr_start) / self.lr_warm_up_steps * n + self.lr_start + self.last_lr = lr + return lr + else: + t = (n - self.lr_warm_up_steps) / (self.lr_max_decay_steps - self.lr_warm_up_steps) + t = min(t, 1.0) + lr = self.lr_min + 0.5 * (self.lr_max - self.lr_min) * ( + 1 + np.cos(t * np.pi)) + self.last_lr = lr + return lr + + def __call__(self, n): + return self.schedule(n) + diff --git a/taming/models/cond_transformer.py b/taming/models/cond_transformer.py new file mode 100644 index 0000000..e4c6373 --- /dev/null +++ b/taming/models/cond_transformer.py @@ -0,0 +1,352 @@ +import os, math +import torch +import torch.nn.functional as F +import pytorch_lightning as pl + +from main import instantiate_from_config +from taming.modules.util import SOSProvider + + +def disabled_train(self, mode=True): + """Overwrite model.train with this function to make sure train/eval mode + does not change anymore.""" + return self + + +class Net2NetTransformer(pl.LightningModule): + def __init__(self, + transformer_config, + first_stage_config, + cond_stage_config, + permuter_config=None, + ckpt_path=None, + ignore_keys=[], + first_stage_key="image", + cond_stage_key="depth", + downsample_cond_size=-1, + pkeep=1.0, + sos_token=0, + unconditional=False, + ): + super().__init__() + self.be_unconditional = unconditional + self.sos_token = sos_token + self.first_stage_key = first_stage_key + self.cond_stage_key = cond_stage_key + self.init_first_stage_from_ckpt(first_stage_config) + self.init_cond_stage_from_ckpt(cond_stage_config) + if permuter_config is None: + permuter_config = {"target": "taming.modules.transformer.permuter.Identity"} + self.permuter = instantiate_from_config(config=permuter_config) + self.transformer = instantiate_from_config(config=transformer_config) + + if ckpt_path is not None: + self.init_from_ckpt(ckpt_path, ignore_keys=ignore_keys) + self.downsample_cond_size = downsample_cond_size + self.pkeep = pkeep + + def init_from_ckpt(self, path, ignore_keys=list()): + sd = torch.load(path, map_location="cpu")["state_dict"] + for k in sd.keys(): + for ik in ignore_keys: + if k.startswith(ik): + self.print("Deleting key {} from state_dict.".format(k)) + del sd[k] + self.load_state_dict(sd, strict=False) + print(f"Restored from {path}") + + def init_first_stage_from_ckpt(self, config): + model = instantiate_from_config(config) + model = model.eval() + model.train = disabled_train + self.first_stage_model = model + + def init_cond_stage_from_ckpt(self, config): + if config == "__is_first_stage__": + print("Using first stage also as cond stage.") + self.cond_stage_model = self.first_stage_model + elif config == "__is_unconditional__" or self.be_unconditional: + print(f"Using no cond stage. Assuming the training is intended to be unconditional. " + f"Prepending {self.sos_token} as a sos token.") + self.be_unconditional = True + self.cond_stage_key = self.first_stage_key + self.cond_stage_model = SOSProvider(self.sos_token) + else: + model = instantiate_from_config(config) + model = model.eval() + model.train = disabled_train + self.cond_stage_model = model + + def forward(self, x, c): + # one step to produce the logits + _, z_indices = self.encode_to_z(x) + _, c_indices = self.encode_to_c(c) + + if self.training and self.pkeep < 1.0: + mask = torch.bernoulli(self.pkeep*torch.ones(z_indices.shape, + device=z_indices.device)) + mask = mask.round().to(dtype=torch.int64) + r_indices = torch.randint_like(z_indices, self.transformer.config.vocab_size) + a_indices = mask*z_indices+(1-mask)*r_indices + else: + a_indices = z_indices + + cz_indices = torch.cat((c_indices, a_indices), dim=1) + + # target includes all sequence elements (no need to handle first one + # differently because we are conditioning) + target = z_indices + # make the prediction + logits, _ = self.transformer(cz_indices[:, :-1]) + # cut off conditioning outputs - output i corresponds to p(z_i | z_{ -1: + c = F.interpolate(c, size=(self.downsample_cond_size, self.downsample_cond_size)) + quant_c, _, [_,_,indices] = self.cond_stage_model.encode(c) + if len(indices.shape) > 2: + indices = indices.view(c.shape[0], -1) + return quant_c, indices + + @torch.no_grad() + def decode_to_img(self, index, zshape): + index = self.permuter(index, reverse=True) + bhwc = (zshape[0],zshape[2],zshape[3],zshape[1]) + quant_z = self.first_stage_model.quantize.get_codebook_entry( + index.reshape(-1), shape=bhwc) + x = self.first_stage_model.decode(quant_z) + return x + + @torch.no_grad() + def log_images(self, batch, temperature=None, top_k=None, callback=None, lr_interface=False, **kwargs): + log = dict() + + N = 4 + if lr_interface: + x, c = self.get_xc(batch, N, diffuse=False, upsample_factor=8) + else: + x, c = self.get_xc(batch, N) + x = x.to(device=self.device) + c = c.to(device=self.device) + + quant_z, z_indices = self.encode_to_z(x) + quant_c, c_indices = self.encode_to_c(c) + + # create a "half"" sample + z_start_indices = z_indices[:,:z_indices.shape[1]//2] + index_sample = self.sample(z_start_indices, c_indices, + steps=z_indices.shape[1]-z_start_indices.shape[1], + temperature=temperature if temperature is not None else 1.0, + sample=True, + top_k=top_k if top_k is not None else 100, + callback=callback if callback is not None else lambda k: None) + x_sample = self.decode_to_img(index_sample, quant_z.shape) + + # sample + z_start_indices = z_indices[:, :0] + index_sample = self.sample(z_start_indices, c_indices, + steps=z_indices.shape[1], + temperature=temperature if temperature is not None else 1.0, + sample=True, + top_k=top_k if top_k is not None else 100, + callback=callback if callback is not None else lambda k: None) + x_sample_nopix = self.decode_to_img(index_sample, quant_z.shape) + + # det sample + z_start_indices = z_indices[:, :0] + index_sample = self.sample(z_start_indices, c_indices, + steps=z_indices.shape[1], + sample=False, + callback=callback if callback is not None else lambda k: None) + x_sample_det = self.decode_to_img(index_sample, quant_z.shape) + + # reconstruction + x_rec = self.decode_to_img(z_indices, quant_z.shape) + + log["inputs"] = x + log["reconstructions"] = x_rec + + if self.cond_stage_key in ["objects_bbox", "objects_center_points"]: + figure_size = (x_rec.shape[2], x_rec.shape[3]) + dataset = kwargs["pl_module"].trainer.datamodule.datasets["validation"] + label_for_category_no = dataset.get_textual_label_for_category_no + plotter = dataset.conditional_builders[self.cond_stage_key].plot + log["conditioning"] = torch.zeros_like(log["reconstructions"]) + for i in range(quant_c.shape[0]): + log["conditioning"][i] = plotter(quant_c[i], label_for_category_no, figure_size) + log["conditioning_rec"] = log["conditioning"] + elif self.cond_stage_key != "image": + cond_rec = self.cond_stage_model.decode(quant_c) + if self.cond_stage_key == "segmentation": + # get image from segmentation mask + num_classes = cond_rec.shape[1] + + c = torch.argmax(c, dim=1, keepdim=True) + c = F.one_hot(c, num_classes=num_classes) + c = c.squeeze(1).permute(0, 3, 1, 2).float() + c = self.cond_stage_model.to_rgb(c) + + cond_rec = torch.argmax(cond_rec, dim=1, keepdim=True) + cond_rec = F.one_hot(cond_rec, num_classes=num_classes) + cond_rec = cond_rec.squeeze(1).permute(0, 3, 1, 2).float() + cond_rec = self.cond_stage_model.to_rgb(cond_rec) + log["conditioning_rec"] = cond_rec + log["conditioning"] = c + + log["samples_half"] = x_sample + log["samples_nopix"] = x_sample_nopix + log["samples_det"] = x_sample_det + return log + + def get_input(self, key, batch): + x = batch[key] + if len(x.shape) == 3: + x = x[..., None] + if len(x.shape) == 4: + x = x.permute(0, 3, 1, 2).to(memory_format=torch.contiguous_format) + if x.dtype == torch.double: + x = x.float() + return x + + def get_xc(self, batch, N=None): + x = self.get_input(self.first_stage_key, batch) + c = self.get_input(self.cond_stage_key, batch) + if N is not None: + x = x[:N] + c = c[:N] + return x, c + + def shared_step(self, batch, batch_idx): + x, c = self.get_xc(batch) + logits, target = self(x, c) + loss = F.cross_entropy(logits.reshape(-1, logits.size(-1)), target.reshape(-1)) + return loss + + def training_step(self, batch, batch_idx): + loss = self.shared_step(batch, batch_idx) + self.log("train/loss", loss, prog_bar=True, logger=True, on_step=True, on_epoch=True) + return loss + + def validation_step(self, batch, batch_idx): + loss = self.shared_step(batch, batch_idx) + self.log("val/loss", loss, prog_bar=True, logger=True, on_step=True, on_epoch=True) + return loss + + def configure_optimizers(self): + """ + Following minGPT: + This long function is unfortunately doing something very simple and is being very defensive: + We are separating out all parameters of the model into two buckets: those that will experience + weight decay for regularization and those that won't (biases, and layernorm/embedding weights). + We are then returning the PyTorch optimizer object. + """ + # separate out all parameters to those that will and won't experience regularizing weight decay + decay = set() + no_decay = set() + whitelist_weight_modules = (torch.nn.Linear, ) + blacklist_weight_modules = (torch.nn.LayerNorm, torch.nn.Embedding) + for mn, m in self.transformer.named_modules(): + for pn, p in m.named_parameters(): + fpn = '%s.%s' % (mn, pn) if mn else pn # full param name + + if pn.endswith('bias'): + # all biases will not be decayed + no_decay.add(fpn) + elif pn.endswith('weight') and isinstance(m, whitelist_weight_modules): + # weights of whitelist modules will be weight decayed + decay.add(fpn) + elif pn.endswith('weight') and isinstance(m, blacklist_weight_modules): + # weights of blacklist modules will NOT be weight decayed + no_decay.add(fpn) + + # special case the position embedding parameter in the root GPT module as not decayed + no_decay.add('pos_emb') + + # validate that we considered every parameter + param_dict = {pn: p for pn, p in self.transformer.named_parameters()} + inter_params = decay & no_decay + union_params = decay | no_decay + assert len(inter_params) == 0, "parameters %s made it into both decay/no_decay sets!" % (str(inter_params), ) + assert len(param_dict.keys() - union_params) == 0, "parameters %s were not separated into either decay/no_decay set!" \ + % (str(param_dict.keys() - union_params), ) + + # create the pytorch optimizer object + optim_groups = [ + {"params": [param_dict[pn] for pn in sorted(list(decay))], "weight_decay": 0.01}, + {"params": [param_dict[pn] for pn in sorted(list(no_decay))], "weight_decay": 0.0}, + ] + optimizer = torch.optim.AdamW(optim_groups, lr=self.learning_rate, betas=(0.9, 0.95)) + return optimizer diff --git a/taming/models/dummy_cond_stage.py b/taming/models/dummy_cond_stage.py new file mode 100644 index 0000000..6e19938 --- /dev/null +++ b/taming/models/dummy_cond_stage.py @@ -0,0 +1,22 @@ +from torch import Tensor + + +class DummyCondStage: + def __init__(self, conditional_key): + self.conditional_key = conditional_key + self.train = None + + def eval(self): + return self + + @staticmethod + def encode(c: Tensor): + return c, None, (None, None, c) + + @staticmethod + def decode(c: Tensor): + return c + + @staticmethod + def to_rgb(c: Tensor): + return c diff --git a/taming/models/vqgan.py b/taming/models/vqgan.py new file mode 100644 index 0000000..a6950ba --- /dev/null +++ b/taming/models/vqgan.py @@ -0,0 +1,404 @@ +import torch +import torch.nn.functional as F +import pytorch_lightning as pl + +from main import instantiate_from_config + +from taming.modules.diffusionmodules.model import Encoder, Decoder +from taming.modules.vqvae.quantize import VectorQuantizer2 as VectorQuantizer +from taming.modules.vqvae.quantize import GumbelQuantize +from taming.modules.vqvae.quantize import EMAVectorQuantizer + +class VQModel(pl.LightningModule): + def __init__(self, + ddconfig, + lossconfig, + n_embed, + embed_dim, + ckpt_path=None, + ignore_keys=[], + image_key="image", + colorize_nlabels=None, + monitor=None, + remap=None, + sane_index_shape=False, # tell vector quantizer to return indices as bhw + ): + super().__init__() + self.image_key = image_key + self.encoder = Encoder(**ddconfig) + self.decoder = Decoder(**ddconfig) + self.loss = instantiate_from_config(lossconfig) + self.quantize = VectorQuantizer(n_embed, embed_dim, beta=0.25, + remap=remap, sane_index_shape=sane_index_shape) + self.quant_conv = torch.nn.Conv2d(ddconfig["z_channels"], embed_dim, 1) + self.post_quant_conv = torch.nn.Conv2d(embed_dim, ddconfig["z_channels"], 1) + if ckpt_path is not None: + self.init_from_ckpt(ckpt_path, ignore_keys=ignore_keys) + self.image_key = image_key + if colorize_nlabels is not None: + assert type(colorize_nlabels)==int + self.register_buffer("colorize", torch.randn(3, colorize_nlabels, 1, 1)) + if monitor is not None: + self.monitor = monitor + + def init_from_ckpt(self, path, ignore_keys=list()): + sd = torch.load(path, map_location="cpu")["state_dict"] + keys = list(sd.keys()) + for k in keys: + for ik in ignore_keys: + if k.startswith(ik): + print("Deleting key {} from state_dict.".format(k)) + del sd[k] + self.load_state_dict(sd, strict=False) + print(f"Restored from {path}") + + def encode(self, x): + h = self.encoder(x) + h = self.quant_conv(h) + quant, emb_loss, info = self.quantize(h) + return quant, emb_loss, info + + def decode(self, quant): + quant = self.post_quant_conv(quant) + dec = self.decoder(quant) + return dec + + def decode_code(self, code_b): + quant_b = self.quantize.embed_code(code_b) + dec = self.decode(quant_b) + return dec + + def forward(self, input): + quant, diff, _ = self.encode(input) + dec = self.decode(quant) + return dec, diff + + def get_input(self, batch, k): + x = batch[k] + if len(x.shape) == 3: + x = x[..., None] + x = x.permute(0, 3, 1, 2).to(memory_format=torch.contiguous_format) + return x.float() + + def training_step(self, batch, batch_idx, optimizer_idx): + x = self.get_input(batch, self.image_key) + xrec, qloss = self(x) + + if optimizer_idx == 0: + # autoencode + aeloss, log_dict_ae = self.loss(qloss, x, xrec, optimizer_idx, self.global_step, + last_layer=self.get_last_layer(), split="train") + + self.log("train/aeloss", aeloss, prog_bar=True, logger=True, on_step=True, on_epoch=True) + self.log_dict(log_dict_ae, prog_bar=False, logger=True, on_step=True, on_epoch=True) + return aeloss + + if optimizer_idx == 1: + # discriminator + discloss, log_dict_disc = self.loss(qloss, x, xrec, optimizer_idx, self.global_step, + last_layer=self.get_last_layer(), split="train") + self.log("train/discloss", discloss, prog_bar=True, logger=True, on_step=True, on_epoch=True) + self.log_dict(log_dict_disc, prog_bar=False, logger=True, on_step=True, on_epoch=True) + return discloss + + def validation_step(self, batch, batch_idx): + x = self.get_input(batch, self.image_key) + xrec, qloss = self(x) + aeloss, log_dict_ae = self.loss(qloss, x, xrec, 0, self.global_step, + last_layer=self.get_last_layer(), split="val") + + discloss, log_dict_disc = self.loss(qloss, x, xrec, 1, self.global_step, + last_layer=self.get_last_layer(), split="val") + rec_loss = log_dict_ae["val/rec_loss"] + self.log("val/rec_loss", rec_loss, + prog_bar=True, logger=True, on_step=True, on_epoch=True, sync_dist=True) + self.log("val/aeloss", aeloss, + prog_bar=True, logger=True, on_step=True, on_epoch=True, sync_dist=True) + self.log_dict(log_dict_ae) + self.log_dict(log_dict_disc) + return self.log_dict + + def configure_optimizers(self): + lr = self.learning_rate + opt_ae = torch.optim.Adam(list(self.encoder.parameters())+ + list(self.decoder.parameters())+ + list(self.quantize.parameters())+ + list(self.quant_conv.parameters())+ + list(self.post_quant_conv.parameters()), + lr=lr, betas=(0.5, 0.9)) + opt_disc = torch.optim.Adam(self.loss.discriminator.parameters(), + lr=lr, betas=(0.5, 0.9)) + return [opt_ae, opt_disc], [] + + def get_last_layer(self): + return self.decoder.conv_out.weight + + def log_images(self, batch, **kwargs): + log = dict() + x = self.get_input(batch, self.image_key) + x = x.to(self.device) + xrec, _ = self(x) + if x.shape[1] > 3: + # colorize with random projection + assert xrec.shape[1] > 3 + x = self.to_rgb(x) + xrec = self.to_rgb(xrec) + log["inputs"] = x + log["reconstructions"] = xrec + return log + + def to_rgb(self, x): + assert self.image_key == "segmentation" + if not hasattr(self, "colorize"): + self.register_buffer("colorize", torch.randn(3, x.shape[1], 1, 1).to(x)) + x = F.conv2d(x, weight=self.colorize) + x = 2.*(x-x.min())/(x.max()-x.min()) - 1. + return x + + +class VQSegmentationModel(VQModel): + def __init__(self, n_labels, *args, **kwargs): + super().__init__(*args, **kwargs) + self.register_buffer("colorize", torch.randn(3, n_labels, 1, 1)) + + def configure_optimizers(self): + lr = self.learning_rate + opt_ae = torch.optim.Adam(list(self.encoder.parameters())+ + list(self.decoder.parameters())+ + list(self.quantize.parameters())+ + list(self.quant_conv.parameters())+ + list(self.post_quant_conv.parameters()), + lr=lr, betas=(0.5, 0.9)) + return opt_ae + + def training_step(self, batch, batch_idx): + x = self.get_input(batch, self.image_key) + xrec, qloss = self(x) + aeloss, log_dict_ae = self.loss(qloss, x, xrec, split="train") + self.log_dict(log_dict_ae, prog_bar=False, logger=True, on_step=True, on_epoch=True) + return aeloss + + def validation_step(self, batch, batch_idx): + x = self.get_input(batch, self.image_key) + xrec, qloss = self(x) + aeloss, log_dict_ae = self.loss(qloss, x, xrec, split="val") + self.log_dict(log_dict_ae, prog_bar=False, logger=True, on_step=True, on_epoch=True) + total_loss = log_dict_ae["val/total_loss"] + self.log("val/total_loss", total_loss, + prog_bar=True, logger=True, on_step=True, on_epoch=True, sync_dist=True) + return aeloss + + @torch.no_grad() + def log_images(self, batch, **kwargs): + log = dict() + x = self.get_input(batch, self.image_key) + x = x.to(self.device) + xrec, _ = self(x) + if x.shape[1] > 3: + # colorize with random projection + assert xrec.shape[1] > 3 + # convert logits to indices + xrec = torch.argmax(xrec, dim=1, keepdim=True) + xrec = F.one_hot(xrec, num_classes=x.shape[1]) + xrec = xrec.squeeze(1).permute(0, 3, 1, 2).float() + x = self.to_rgb(x) + xrec = self.to_rgb(xrec) + log["inputs"] = x + log["reconstructions"] = xrec + return log + + +class VQNoDiscModel(VQModel): + def __init__(self, + ddconfig, + lossconfig, + n_embed, + embed_dim, + ckpt_path=None, + ignore_keys=[], + image_key="image", + colorize_nlabels=None + ): + super().__init__(ddconfig=ddconfig, lossconfig=lossconfig, n_embed=n_embed, embed_dim=embed_dim, + ckpt_path=ckpt_path, ignore_keys=ignore_keys, image_key=image_key, + colorize_nlabels=colorize_nlabels) + + def training_step(self, batch, batch_idx): + x = self.get_input(batch, self.image_key) + xrec, qloss = self(x) + # autoencode + aeloss, log_dict_ae = self.loss(qloss, x, xrec, self.global_step, split="train") + output = pl.TrainResult(minimize=aeloss) + output.log("train/aeloss", aeloss, + prog_bar=True, logger=True, on_step=True, on_epoch=True) + output.log_dict(log_dict_ae, prog_bar=False, logger=True, on_step=True, on_epoch=True) + return output + + def validation_step(self, batch, batch_idx): + x = self.get_input(batch, self.image_key) + xrec, qloss = self(x) + aeloss, log_dict_ae = self.loss(qloss, x, xrec, self.global_step, split="val") + rec_loss = log_dict_ae["val/rec_loss"] + output = pl.EvalResult(checkpoint_on=rec_loss) + output.log("val/rec_loss", rec_loss, + prog_bar=True, logger=True, on_step=True, on_epoch=True) + output.log("val/aeloss", aeloss, + prog_bar=True, logger=True, on_step=True, on_epoch=True) + output.log_dict(log_dict_ae) + + return output + + def configure_optimizers(self): + optimizer = torch.optim.Adam(list(self.encoder.parameters())+ + list(self.decoder.parameters())+ + list(self.quantize.parameters())+ + list(self.quant_conv.parameters())+ + list(self.post_quant_conv.parameters()), + lr=self.learning_rate, betas=(0.5, 0.9)) + return optimizer + + +class GumbelVQ(VQModel): + def __init__(self, + ddconfig, + lossconfig, + n_embed, + embed_dim, + temperature_scheduler_config, + ckpt_path=None, + ignore_keys=[], + image_key="image", + colorize_nlabels=None, + monitor=None, + kl_weight=1e-8, + remap=None, + ): + + z_channels = ddconfig["z_channels"] + super().__init__(ddconfig, + lossconfig, + n_embed, + embed_dim, + ckpt_path=None, + ignore_keys=ignore_keys, + image_key=image_key, + colorize_nlabels=colorize_nlabels, + monitor=monitor, + ) + + self.loss.n_classes = n_embed + self.vocab_size = n_embed + + self.quantize = GumbelQuantize(z_channels, embed_dim, + n_embed=n_embed, + kl_weight=kl_weight, temp_init=1.0, + remap=remap) + + self.temperature_scheduler = instantiate_from_config(temperature_scheduler_config) # annealing of temp + + if ckpt_path is not None: + self.init_from_ckpt(ckpt_path, ignore_keys=ignore_keys) + + def temperature_scheduling(self): + self.quantize.temperature = self.temperature_scheduler(self.global_step) + + def encode_to_prequant(self, x): + h = self.encoder(x) + h = self.quant_conv(h) + return h + + def decode_code(self, code_b): + raise NotImplementedError + + def training_step(self, batch, batch_idx, optimizer_idx): + self.temperature_scheduling() + x = self.get_input(batch, self.image_key) + xrec, qloss = self(x) + + if optimizer_idx == 0: + # autoencode + aeloss, log_dict_ae = self.loss(qloss, x, xrec, optimizer_idx, self.global_step, + last_layer=self.get_last_layer(), split="train") + + self.log_dict(log_dict_ae, prog_bar=False, logger=True, on_step=True, on_epoch=True) + self.log("temperature", self.quantize.temperature, prog_bar=False, logger=True, on_step=True, on_epoch=True) + return aeloss + + if optimizer_idx == 1: + # discriminator + discloss, log_dict_disc = self.loss(qloss, x, xrec, optimizer_idx, self.global_step, + last_layer=self.get_last_layer(), split="train") + self.log_dict(log_dict_disc, prog_bar=False, logger=True, on_step=True, on_epoch=True) + return discloss + + def validation_step(self, batch, batch_idx): + x = self.get_input(batch, self.image_key) + xrec, qloss = self(x, return_pred_indices=True) + aeloss, log_dict_ae = self.loss(qloss, x, xrec, 0, self.global_step, + last_layer=self.get_last_layer(), split="val") + + discloss, log_dict_disc = self.loss(qloss, x, xrec, 1, self.global_step, + last_layer=self.get_last_layer(), split="val") + rec_loss = log_dict_ae["val/rec_loss"] + self.log("val/rec_loss", rec_loss, + prog_bar=True, logger=True, on_step=False, on_epoch=True, sync_dist=True) + self.log("val/aeloss", aeloss, + prog_bar=True, logger=True, on_step=False, on_epoch=True, sync_dist=True) + self.log_dict(log_dict_ae) + self.log_dict(log_dict_disc) + return self.log_dict + + def log_images(self, batch, **kwargs): + log = dict() + x = self.get_input(batch, self.image_key) + x = x.to(self.device) + # encode + h = self.encoder(x) + h = self.quant_conv(h) + quant, _, _ = self.quantize(h) + # decode + x_rec = self.decode(quant) + log["inputs"] = x + log["reconstructions"] = x_rec + return log + + +class EMAVQ(VQModel): + def __init__(self, + ddconfig, + lossconfig, + n_embed, + embed_dim, + ckpt_path=None, + ignore_keys=[], + image_key="image", + colorize_nlabels=None, + monitor=None, + remap=None, + sane_index_shape=False, # tell vector quantizer to return indices as bhw + ): + super().__init__(ddconfig, + lossconfig, + n_embed, + embed_dim, + ckpt_path=None, + ignore_keys=ignore_keys, + image_key=image_key, + colorize_nlabels=colorize_nlabels, + monitor=monitor, + ) + self.quantize = EMAVectorQuantizer(n_embed=n_embed, + embedding_dim=embed_dim, + beta=0.25, + remap=remap) + def configure_optimizers(self): + lr = self.learning_rate + #Remove self.quantize from parameter list since it is updated via EMA + opt_ae = torch.optim.Adam(list(self.encoder.parameters())+ + list(self.decoder.parameters())+ + list(self.quant_conv.parameters())+ + list(self.post_quant_conv.parameters()), + lr=lr, betas=(0.5, 0.9)) + opt_disc = torch.optim.Adam(self.loss.discriminator.parameters(), + lr=lr, betas=(0.5, 0.9)) + return [opt_ae, opt_disc], [] \ No newline at end of file diff --git a/taming/modules/diffusionmodules/model.py b/taming/modules/diffusionmodules/model.py new file mode 100644 index 0000000..d3a5db6 --- /dev/null +++ b/taming/modules/diffusionmodules/model.py @@ -0,0 +1,776 @@ +# pytorch_diffusion + derived encoder decoder +import math +import torch +import torch.nn as nn +import numpy as np + + +def get_timestep_embedding(timesteps, embedding_dim): + """ + This matches the implementation in Denoising Diffusion Probabilistic Models: + From Fairseq. + Build sinusoidal embeddings. + This matches the implementation in tensor2tensor, but differs slightly + from the description in Section 3.5 of "Attention Is All You Need". + """ + assert len(timesteps.shape) == 1 + + half_dim = embedding_dim // 2 + emb = math.log(10000) / (half_dim - 1) + emb = torch.exp(torch.arange(half_dim, dtype=torch.float32) * -emb) + emb = emb.to(device=timesteps.device) + emb = timesteps.float()[:, None] * emb[None, :] + emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=1) + if embedding_dim % 2 == 1: # zero pad + emb = torch.nn.functional.pad(emb, (0,1,0,0)) + return emb + + +def nonlinearity(x): + # swish + return x*torch.sigmoid(x) + + +def Normalize(in_channels): + return torch.nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True) + + +class Upsample(nn.Module): + def __init__(self, in_channels, with_conv): + super().__init__() + self.with_conv = with_conv + if self.with_conv: + self.conv = torch.nn.Conv2d(in_channels, + in_channels, + kernel_size=3, + stride=1, + padding=1) + + def forward(self, x): + x = torch.nn.functional.interpolate(x, scale_factor=2.0, mode="nearest") + if self.with_conv: + x = self.conv(x) + return x + + +class Downsample(nn.Module): + def __init__(self, in_channels, with_conv): + super().__init__() + self.with_conv = with_conv + if self.with_conv: + # no asymmetric padding in torch conv, must do it ourselves + self.conv = torch.nn.Conv2d(in_channels, + in_channels, + kernel_size=3, + stride=2, + padding=0) + + def forward(self, x): + if self.with_conv: + pad = (0,1,0,1) + x = torch.nn.functional.pad(x, pad, mode="constant", value=0) + x = self.conv(x) + else: + x = torch.nn.functional.avg_pool2d(x, kernel_size=2, stride=2) + return x + + +class ResnetBlock(nn.Module): + def __init__(self, *, in_channels, out_channels=None, conv_shortcut=False, + dropout, temb_channels=512): + super().__init__() + self.in_channels = in_channels + out_channels = in_channels if out_channels is None else out_channels + self.out_channels = out_channels + self.use_conv_shortcut = conv_shortcut + + self.norm1 = Normalize(in_channels) + self.conv1 = torch.nn.Conv2d(in_channels, + out_channels, + kernel_size=3, + stride=1, + padding=1) + if temb_channels > 0: + self.temb_proj = torch.nn.Linear(temb_channels, + out_channels) + self.norm2 = Normalize(out_channels) + self.dropout = torch.nn.Dropout(dropout) + self.conv2 = torch.nn.Conv2d(out_channels, + out_channels, + kernel_size=3, + stride=1, + padding=1) + if self.in_channels != self.out_channels: + if self.use_conv_shortcut: + self.conv_shortcut = torch.nn.Conv2d(in_channels, + out_channels, + kernel_size=3, + stride=1, + padding=1) + else: + self.nin_shortcut = torch.nn.Conv2d(in_channels, + out_channels, + kernel_size=1, + stride=1, + padding=0) + + def forward(self, x, temb): + h = x + h = self.norm1(h) + h = nonlinearity(h) + h = self.conv1(h) + + if temb is not None: + h = h + self.temb_proj(nonlinearity(temb))[:,:,None,None] + + h = self.norm2(h) + h = nonlinearity(h) + h = self.dropout(h) + h = self.conv2(h) + + if self.in_channels != self.out_channels: + if self.use_conv_shortcut: + x = self.conv_shortcut(x) + else: + x = self.nin_shortcut(x) + + return x+h + + +class AttnBlock(nn.Module): + def __init__(self, in_channels): + super().__init__() + self.in_channels = in_channels + + self.norm = Normalize(in_channels) + self.q = torch.nn.Conv2d(in_channels, + in_channels, + kernel_size=1, + stride=1, + padding=0) + self.k = torch.nn.Conv2d(in_channels, + in_channels, + kernel_size=1, + stride=1, + padding=0) + self.v = torch.nn.Conv2d(in_channels, + in_channels, + kernel_size=1, + stride=1, + padding=0) + self.proj_out = torch.nn.Conv2d(in_channels, + in_channels, + kernel_size=1, + stride=1, + padding=0) + + + def forward(self, x): + h_ = x + h_ = self.norm(h_) + q = self.q(h_) + k = self.k(h_) + v = self.v(h_) + + # compute attention + b,c,h,w = q.shape + q = q.reshape(b,c,h*w) + q = q.permute(0,2,1) # b,hw,c + k = k.reshape(b,c,h*w) # b,c,hw + w_ = torch.bmm(q,k) # b,hw,hw w[b,i,j]=sum_c q[b,i,c]k[b,c,j] + w_ = w_ * (int(c)**(-0.5)) + w_ = torch.nn.functional.softmax(w_, dim=2) + + # attend to values + v = v.reshape(b,c,h*w) + w_ = w_.permute(0,2,1) # b,hw,hw (first hw of k, second of q) + h_ = torch.bmm(v,w_) # b, c,hw (hw of q) h_[b,c,j] = sum_i v[b,c,i] w_[b,i,j] + h_ = h_.reshape(b,c,h,w) + + h_ = self.proj_out(h_) + + return x+h_ + + +class Model(nn.Module): + def __init__(self, *, ch, out_ch, ch_mult=(1,2,4,8), num_res_blocks, + attn_resolutions, dropout=0.0, resamp_with_conv=True, in_channels, + resolution, use_timestep=True): + super().__init__() + self.ch = ch + self.temb_ch = self.ch*4 + self.num_resolutions = len(ch_mult) + self.num_res_blocks = num_res_blocks + self.resolution = resolution + self.in_channels = in_channels + + self.use_timestep = use_timestep + if self.use_timestep: + # timestep embedding + self.temb = nn.Module() + self.temb.dense = nn.ModuleList([ + torch.nn.Linear(self.ch, + self.temb_ch), + torch.nn.Linear(self.temb_ch, + self.temb_ch), + ]) + + # downsampling + self.conv_in = torch.nn.Conv2d(in_channels, + self.ch, + kernel_size=3, + stride=1, + padding=1) + + curr_res = resolution + in_ch_mult = (1,)+tuple(ch_mult) + self.down = nn.ModuleList() + for i_level in range(self.num_resolutions): + block = nn.ModuleList() + attn = nn.ModuleList() + block_in = ch*in_ch_mult[i_level] + block_out = ch*ch_mult[i_level] + for i_block in range(self.num_res_blocks): + block.append(ResnetBlock(in_channels=block_in, + out_channels=block_out, + temb_channels=self.temb_ch, + dropout=dropout)) + block_in = block_out + if curr_res in attn_resolutions: + attn.append(AttnBlock(block_in)) + down = nn.Module() + down.block = block + down.attn = attn + if i_level != self.num_resolutions-1: + down.downsample = Downsample(block_in, resamp_with_conv) + curr_res = curr_res // 2 + self.down.append(down) + + # middle + self.mid = nn.Module() + self.mid.block_1 = ResnetBlock(in_channels=block_in, + out_channels=block_in, + temb_channels=self.temb_ch, + dropout=dropout) + self.mid.attn_1 = AttnBlock(block_in) + self.mid.block_2 = ResnetBlock(in_channels=block_in, + out_channels=block_in, + temb_channels=self.temb_ch, + dropout=dropout) + + # upsampling + self.up = nn.ModuleList() + for i_level in reversed(range(self.num_resolutions)): + block = nn.ModuleList() + attn = nn.ModuleList() + block_out = ch*ch_mult[i_level] + skip_in = ch*ch_mult[i_level] + for i_block in range(self.num_res_blocks+1): + if i_block == self.num_res_blocks: + skip_in = ch*in_ch_mult[i_level] + block.append(ResnetBlock(in_channels=block_in+skip_in, + out_channels=block_out, + temb_channels=self.temb_ch, + dropout=dropout)) + block_in = block_out + if curr_res in attn_resolutions: + attn.append(AttnBlock(block_in)) + up = nn.Module() + up.block = block + up.attn = attn + if i_level != 0: + up.upsample = Upsample(block_in, resamp_with_conv) + curr_res = curr_res * 2 + self.up.insert(0, up) # prepend to get consistent order + + # end + self.norm_out = Normalize(block_in) + self.conv_out = torch.nn.Conv2d(block_in, + out_ch, + kernel_size=3, + stride=1, + padding=1) + + + def forward(self, x, t=None): + #assert x.shape[2] == x.shape[3] == self.resolution + + if self.use_timestep: + # timestep embedding + assert t is not None + temb = get_timestep_embedding(t, self.ch) + temb = self.temb.dense[0](temb) + temb = nonlinearity(temb) + temb = self.temb.dense[1](temb) + else: + temb = None + + # downsampling + hs = [self.conv_in(x)] + for i_level in range(self.num_resolutions): + for i_block in range(self.num_res_blocks): + h = self.down[i_level].block[i_block](hs[-1], temb) + if len(self.down[i_level].attn) > 0: + h = self.down[i_level].attn[i_block](h) + hs.append(h) + if i_level != self.num_resolutions-1: + hs.append(self.down[i_level].downsample(hs[-1])) + + # middle + h = hs[-1] + h = self.mid.block_1(h, temb) + h = self.mid.attn_1(h) + h = self.mid.block_2(h, temb) + + # upsampling + for i_level in reversed(range(self.num_resolutions)): + for i_block in range(self.num_res_blocks+1): + h = self.up[i_level].block[i_block]( + torch.cat([h, hs.pop()], dim=1), temb) + if len(self.up[i_level].attn) > 0: + h = self.up[i_level].attn[i_block](h) + if i_level != 0: + h = self.up[i_level].upsample(h) + + # end + h = self.norm_out(h) + h = nonlinearity(h) + h = self.conv_out(h) + return h + + +class Encoder(nn.Module): + def __init__(self, *, ch, out_ch, ch_mult=(1,2,4,8), num_res_blocks, + attn_resolutions, dropout=0.0, resamp_with_conv=True, in_channels, + resolution, z_channels, double_z=True, **ignore_kwargs): + super().__init__() + self.ch = ch + self.temb_ch = 0 + self.num_resolutions = len(ch_mult) + self.num_res_blocks = num_res_blocks + self.resolution = resolution + self.in_channels = in_channels + + # downsampling + self.conv_in = torch.nn.Conv2d(in_channels, + self.ch, + kernel_size=3, + stride=1, + padding=1) + + curr_res = resolution + in_ch_mult = (1,)+tuple(ch_mult) + self.down = nn.ModuleList() + for i_level in range(self.num_resolutions): + block = nn.ModuleList() + attn = nn.ModuleList() + block_in = ch*in_ch_mult[i_level] + block_out = ch*ch_mult[i_level] + for i_block in range(self.num_res_blocks): + block.append(ResnetBlock(in_channels=block_in, + out_channels=block_out, + temb_channels=self.temb_ch, + dropout=dropout)) + block_in = block_out + if curr_res in attn_resolutions: + attn.append(AttnBlock(block_in)) + down = nn.Module() + down.block = block + down.attn = attn + if i_level != self.num_resolutions-1: + down.downsample = Downsample(block_in, resamp_with_conv) + curr_res = curr_res // 2 + self.down.append(down) + + # middle + self.mid = nn.Module() + self.mid.block_1 = ResnetBlock(in_channels=block_in, + out_channels=block_in, + temb_channels=self.temb_ch, + dropout=dropout) + self.mid.attn_1 = AttnBlock(block_in) + self.mid.block_2 = ResnetBlock(in_channels=block_in, + out_channels=block_in, + temb_channels=self.temb_ch, + dropout=dropout) + + # end + self.norm_out = Normalize(block_in) + self.conv_out = torch.nn.Conv2d(block_in, + 2*z_channels if double_z else z_channels, + kernel_size=3, + stride=1, + padding=1) + + + def forward(self, x): + #assert x.shape[2] == x.shape[3] == self.resolution, "{}, {}, {}".format(x.shape[2], x.shape[3], self.resolution) + + # timestep embedding + temb = None + + # downsampling + hs = [self.conv_in(x)] + for i_level in range(self.num_resolutions): + for i_block in range(self.num_res_blocks): + h = self.down[i_level].block[i_block](hs[-1], temb) + if len(self.down[i_level].attn) > 0: + h = self.down[i_level].attn[i_block](h) + hs.append(h) + if i_level != self.num_resolutions-1: + hs.append(self.down[i_level].downsample(hs[-1])) + + # middle + h = hs[-1] + h = self.mid.block_1(h, temb) + h = self.mid.attn_1(h) + h = self.mid.block_2(h, temb) + + # end + h = self.norm_out(h) + h = nonlinearity(h) + h = self.conv_out(h) + return h + + +class Decoder(nn.Module): + def __init__(self, *, ch, out_ch, ch_mult=(1,2,4,8), num_res_blocks, + attn_resolutions, dropout=0.0, resamp_with_conv=True, in_channels, + resolution, z_channels, give_pre_end=False, **ignorekwargs): + super().__init__() + self.ch = ch + self.temb_ch = 0 + self.num_resolutions = len(ch_mult) + self.num_res_blocks = num_res_blocks + self.resolution = resolution + self.in_channels = in_channels + self.give_pre_end = give_pre_end + + # compute in_ch_mult, block_in and curr_res at lowest res + in_ch_mult = (1,)+tuple(ch_mult) + block_in = ch*ch_mult[self.num_resolutions-1] + curr_res = resolution // 2**(self.num_resolutions-1) + self.z_shape = (1,z_channels,curr_res,curr_res) + print("Working with z of shape {} = {} dimensions.".format( + self.z_shape, np.prod(self.z_shape))) + + # z to block_in + self.conv_in = torch.nn.Conv2d(z_channels, + block_in, + kernel_size=3, + stride=1, + padding=1) + + # middle + self.mid = nn.Module() + self.mid.block_1 = ResnetBlock(in_channels=block_in, + out_channels=block_in, + temb_channels=self.temb_ch, + dropout=dropout) + self.mid.attn_1 = AttnBlock(block_in) + self.mid.block_2 = ResnetBlock(in_channels=block_in, + out_channels=block_in, + temb_channels=self.temb_ch, + dropout=dropout) + + # upsampling + self.up = nn.ModuleList() + for i_level in reversed(range(self.num_resolutions)): + block = nn.ModuleList() + attn = nn.ModuleList() + block_out = ch*ch_mult[i_level] + for i_block in range(self.num_res_blocks+1): + block.append(ResnetBlock(in_channels=block_in, + out_channels=block_out, + temb_channels=self.temb_ch, + dropout=dropout)) + block_in = block_out + if curr_res in attn_resolutions: + attn.append(AttnBlock(block_in)) + up = nn.Module() + up.block = block + up.attn = attn + if i_level != 0: + up.upsample = Upsample(block_in, resamp_with_conv) + curr_res = curr_res * 2 + self.up.insert(0, up) # prepend to get consistent order + + # end + self.norm_out = Normalize(block_in) + self.conv_out = torch.nn.Conv2d(block_in, + out_ch, + kernel_size=3, + stride=1, + padding=1) + + def forward(self, z): + #assert z.shape[1:] == self.z_shape[1:] + self.last_z_shape = z.shape + + # timestep embedding + temb = None + + # z to block_in + h = self.conv_in(z) + + # middle + h = self.mid.block_1(h, temb) + h = self.mid.attn_1(h) + h = self.mid.block_2(h, temb) + + # upsampling + for i_level in reversed(range(self.num_resolutions)): + for i_block in range(self.num_res_blocks+1): + h = self.up[i_level].block[i_block](h, temb) + if len(self.up[i_level].attn) > 0: + h = self.up[i_level].attn[i_block](h) + if i_level != 0: + h = self.up[i_level].upsample(h) + + # end + if self.give_pre_end: + return h + + h = self.norm_out(h) + h = nonlinearity(h) + h = self.conv_out(h) + return h + + +class VUNet(nn.Module): + def __init__(self, *, ch, out_ch, ch_mult=(1,2,4,8), num_res_blocks, + attn_resolutions, dropout=0.0, resamp_with_conv=True, + in_channels, c_channels, + resolution, z_channels, use_timestep=False, **ignore_kwargs): + super().__init__() + self.ch = ch + self.temb_ch = self.ch*4 + self.num_resolutions = len(ch_mult) + self.num_res_blocks = num_res_blocks + self.resolution = resolution + + self.use_timestep = use_timestep + if self.use_timestep: + # timestep embedding + self.temb = nn.Module() + self.temb.dense = nn.ModuleList([ + torch.nn.Linear(self.ch, + self.temb_ch), + torch.nn.Linear(self.temb_ch, + self.temb_ch), + ]) + + # downsampling + self.conv_in = torch.nn.Conv2d(c_channels, + self.ch, + kernel_size=3, + stride=1, + padding=1) + + curr_res = resolution + in_ch_mult = (1,)+tuple(ch_mult) + self.down = nn.ModuleList() + for i_level in range(self.num_resolutions): + block = nn.ModuleList() + attn = nn.ModuleList() + block_in = ch*in_ch_mult[i_level] + block_out = ch*ch_mult[i_level] + for i_block in range(self.num_res_blocks): + block.append(ResnetBlock(in_channels=block_in, + out_channels=block_out, + temb_channels=self.temb_ch, + dropout=dropout)) + block_in = block_out + if curr_res in attn_resolutions: + attn.append(AttnBlock(block_in)) + down = nn.Module() + down.block = block + down.attn = attn + if i_level != self.num_resolutions-1: + down.downsample = Downsample(block_in, resamp_with_conv) + curr_res = curr_res // 2 + self.down.append(down) + + self.z_in = torch.nn.Conv2d(z_channels, + block_in, + kernel_size=1, + stride=1, + padding=0) + # middle + self.mid = nn.Module() + self.mid.block_1 = ResnetBlock(in_channels=2*block_in, + out_channels=block_in, + temb_channels=self.temb_ch, + dropout=dropout) + self.mid.attn_1 = AttnBlock(block_in) + self.mid.block_2 = ResnetBlock(in_channels=block_in, + out_channels=block_in, + temb_channels=self.temb_ch, + dropout=dropout) + + # upsampling + self.up = nn.ModuleList() + for i_level in reversed(range(self.num_resolutions)): + block = nn.ModuleList() + attn = nn.ModuleList() + block_out = ch*ch_mult[i_level] + skip_in = ch*ch_mult[i_level] + for i_block in range(self.num_res_blocks+1): + if i_block == self.num_res_blocks: + skip_in = ch*in_ch_mult[i_level] + block.append(ResnetBlock(in_channels=block_in+skip_in, + out_channels=block_out, + temb_channels=self.temb_ch, + dropout=dropout)) + block_in = block_out + if curr_res in attn_resolutions: + attn.append(AttnBlock(block_in)) + up = nn.Module() + up.block = block + up.attn = attn + if i_level != 0: + up.upsample = Upsample(block_in, resamp_with_conv) + curr_res = curr_res * 2 + self.up.insert(0, up) # prepend to get consistent order + + # end + self.norm_out = Normalize(block_in) + self.conv_out = torch.nn.Conv2d(block_in, + out_ch, + kernel_size=3, + stride=1, + padding=1) + + + def forward(self, x, z): + #assert x.shape[2] == x.shape[3] == self.resolution + + if self.use_timestep: + # timestep embedding + assert t is not None + temb = get_timestep_embedding(t, self.ch) + temb = self.temb.dense[0](temb) + temb = nonlinearity(temb) + temb = self.temb.dense[1](temb) + else: + temb = None + + # downsampling + hs = [self.conv_in(x)] + for i_level in range(self.num_resolutions): + for i_block in range(self.num_res_blocks): + h = self.down[i_level].block[i_block](hs[-1], temb) + if len(self.down[i_level].attn) > 0: + h = self.down[i_level].attn[i_block](h) + hs.append(h) + if i_level != self.num_resolutions-1: + hs.append(self.down[i_level].downsample(hs[-1])) + + # middle + h = hs[-1] + z = self.z_in(z) + h = torch.cat((h,z),dim=1) + h = self.mid.block_1(h, temb) + h = self.mid.attn_1(h) + h = self.mid.block_2(h, temb) + + # upsampling + for i_level in reversed(range(self.num_resolutions)): + for i_block in range(self.num_res_blocks+1): + h = self.up[i_level].block[i_block]( + torch.cat([h, hs.pop()], dim=1), temb) + if len(self.up[i_level].attn) > 0: + h = self.up[i_level].attn[i_block](h) + if i_level != 0: + h = self.up[i_level].upsample(h) + + # end + h = self.norm_out(h) + h = nonlinearity(h) + h = self.conv_out(h) + return h + + +class SimpleDecoder(nn.Module): + def __init__(self, in_channels, out_channels, *args, **kwargs): + super().__init__() + self.model = nn.ModuleList([nn.Conv2d(in_channels, in_channels, 1), + ResnetBlock(in_channels=in_channels, + out_channels=2 * in_channels, + temb_channels=0, dropout=0.0), + ResnetBlock(in_channels=2 * in_channels, + out_channels=4 * in_channels, + temb_channels=0, dropout=0.0), + ResnetBlock(in_channels=4 * in_channels, + out_channels=2 * in_channels, + temb_channels=0, dropout=0.0), + nn.Conv2d(2*in_channels, in_channels, 1), + Upsample(in_channels, with_conv=True)]) + # end + self.norm_out = Normalize(in_channels) + self.conv_out = torch.nn.Conv2d(in_channels, + out_channels, + kernel_size=3, + stride=1, + padding=1) + + def forward(self, x): + for i, layer in enumerate(self.model): + if i in [1,2,3]: + x = layer(x, None) + else: + x = layer(x) + + h = self.norm_out(x) + h = nonlinearity(h) + x = self.conv_out(h) + return x + + +class UpsampleDecoder(nn.Module): + def __init__(self, in_channels, out_channels, ch, num_res_blocks, resolution, + ch_mult=(2,2), dropout=0.0): + super().__init__() + # upsampling + self.temb_ch = 0 + self.num_resolutions = len(ch_mult) + self.num_res_blocks = num_res_blocks + block_in = in_channels + curr_res = resolution // 2 ** (self.num_resolutions - 1) + self.res_blocks = nn.ModuleList() + self.upsample_blocks = nn.ModuleList() + for i_level in range(self.num_resolutions): + res_block = [] + block_out = ch * ch_mult[i_level] + for i_block in range(self.num_res_blocks + 1): + res_block.append(ResnetBlock(in_channels=block_in, + out_channels=block_out, + temb_channels=self.temb_ch, + dropout=dropout)) + block_in = block_out + self.res_blocks.append(nn.ModuleList(res_block)) + if i_level != self.num_resolutions - 1: + self.upsample_blocks.append(Upsample(block_in, True)) + curr_res = curr_res * 2 + + # end + self.norm_out = Normalize(block_in) + self.conv_out = torch.nn.Conv2d(block_in, + out_channels, + kernel_size=3, + stride=1, + padding=1) + + def forward(self, x): + # upsampling + h = x + for k, i_level in enumerate(range(self.num_resolutions)): + for i_block in range(self.num_res_blocks + 1): + h = self.res_blocks[i_level][i_block](h, None) + if i_level != self.num_resolutions - 1: + h = self.upsample_blocks[k](h) + h = self.norm_out(h) + h = nonlinearity(h) + h = self.conv_out(h) + return h + diff --git a/taming/modules/discriminator/model.py b/taming/modules/discriminator/model.py new file mode 100644 index 0000000..2aaa311 --- /dev/null +++ b/taming/modules/discriminator/model.py @@ -0,0 +1,67 @@ +import functools +import torch.nn as nn + + +from taming.modules.util import ActNorm + + +def weights_init(m): + classname = m.__class__.__name__ + if classname.find('Conv') != -1: + nn.init.normal_(m.weight.data, 0.0, 0.02) + elif classname.find('BatchNorm') != -1: + nn.init.normal_(m.weight.data, 1.0, 0.02) + nn.init.constant_(m.bias.data, 0) + + +class NLayerDiscriminator(nn.Module): + """Defines a PatchGAN discriminator as in Pix2Pix + --> see https://github.com/junyanz/pytorch-CycleGAN-and-pix2pix/blob/master/models/networks.py + """ + def __init__(self, input_nc=3, ndf=64, n_layers=3, use_actnorm=False): + """Construct a PatchGAN discriminator + Parameters: + input_nc (int) -- the number of channels in input images + ndf (int) -- the number of filters in the last conv layer + n_layers (int) -- the number of conv layers in the discriminator + norm_layer -- normalization layer + """ + super(NLayerDiscriminator, self).__init__() + if not use_actnorm: + norm_layer = nn.BatchNorm2d + else: + norm_layer = ActNorm + if type(norm_layer) == functools.partial: # no need to use bias as BatchNorm2d has affine parameters + use_bias = norm_layer.func != nn.BatchNorm2d + else: + use_bias = norm_layer != nn.BatchNorm2d + + kw = 4 + padw = 1 + sequence = [nn.Conv2d(input_nc, ndf, kernel_size=kw, stride=2, padding=padw), nn.LeakyReLU(0.2, True)] + nf_mult = 1 + nf_mult_prev = 1 + for n in range(1, n_layers): # gradually increase the number of filters + nf_mult_prev = nf_mult + nf_mult = min(2 ** n, 8) + sequence += [ + nn.Conv2d(ndf * nf_mult_prev, ndf * nf_mult, kernel_size=kw, stride=2, padding=padw, bias=use_bias), + norm_layer(ndf * nf_mult), + nn.LeakyReLU(0.2, True) + ] + + nf_mult_prev = nf_mult + nf_mult = min(2 ** n_layers, 8) + sequence += [ + nn.Conv2d(ndf * nf_mult_prev, ndf * nf_mult, kernel_size=kw, stride=1, padding=padw, bias=use_bias), + norm_layer(ndf * nf_mult), + nn.LeakyReLU(0.2, True) + ] + + sequence += [ + nn.Conv2d(ndf * nf_mult, 1, kernel_size=kw, stride=1, padding=padw)] # output 1 channel prediction map + self.main = nn.Sequential(*sequence) + + def forward(self, input): + """Standard forward.""" + return self.main(input) diff --git a/taming/modules/losses/__init__.py b/taming/modules/losses/__init__.py new file mode 100644 index 0000000..d09caf9 --- /dev/null +++ b/taming/modules/losses/__init__.py @@ -0,0 +1,2 @@ +from taming.modules.losses.vqperceptual import DummyLoss + diff --git a/taming/modules/losses/__pycache__/__init__.cpython-310.pyc b/taming/modules/losses/__pycache__/__init__.cpython-310.pyc new file mode 100644 index 0000000..9762841 Binary files /dev/null and b/taming/modules/losses/__pycache__/__init__.cpython-310.pyc differ diff --git a/taming/modules/losses/lpips.py b/taming/modules/losses/lpips.py new file mode 100644 index 0000000..a728044 --- /dev/null +++ b/taming/modules/losses/lpips.py @@ -0,0 +1,123 @@ +"""Stripped version of https://github.com/richzhang/PerceptualSimilarity/tree/master/models""" + +import torch +import torch.nn as nn +from torchvision import models +from collections import namedtuple + +from taming.util import get_ckpt_path + + +class LPIPS(nn.Module): + # Learned perceptual metric + def __init__(self, use_dropout=True): + super().__init__() + self.scaling_layer = ScalingLayer() + self.chns = [64, 128, 256, 512, 512] # vg16 features + self.net = vgg16(pretrained=True, requires_grad=False) + self.lin0 = NetLinLayer(self.chns[0], use_dropout=use_dropout) + self.lin1 = NetLinLayer(self.chns[1], use_dropout=use_dropout) + self.lin2 = NetLinLayer(self.chns[2], use_dropout=use_dropout) + self.lin3 = NetLinLayer(self.chns[3], use_dropout=use_dropout) + self.lin4 = NetLinLayer(self.chns[4], use_dropout=use_dropout) + self.load_from_pretrained() + for param in self.parameters(): + param.requires_grad = False + + def load_from_pretrained(self, name="vgg_lpips"): + ckpt = get_ckpt_path(name, "taming/modules/autoencoder/lpips") + self.load_state_dict(torch.load(ckpt, map_location=torch.device("cpu")), strict=False) + print("loaded pretrained LPIPS loss from {}".format(ckpt)) + + @classmethod + def from_pretrained(cls, name="vgg_lpips"): + if name != "vgg_lpips": + raise NotImplementedError + model = cls() + ckpt = get_ckpt_path(name) + model.load_state_dict(torch.load(ckpt, map_location=torch.device("cpu")), strict=False) + return model + + def forward(self, input, target): + in0_input, in1_input = (self.scaling_layer(input), self.scaling_layer(target)) + outs0, outs1 = self.net(in0_input), self.net(in1_input) + feats0, feats1, diffs = {}, {}, {} + lins = [self.lin0, self.lin1, self.lin2, self.lin3, self.lin4] + for kk in range(len(self.chns)): + feats0[kk], feats1[kk] = normalize_tensor(outs0[kk]), normalize_tensor(outs1[kk]) + diffs[kk] = (feats0[kk] - feats1[kk]) ** 2 + + res = [spatial_average(lins[kk].model(diffs[kk]), keepdim=True) for kk in range(len(self.chns))] + val = res[0] + for l in range(1, len(self.chns)): + val += res[l] + return val + + +class ScalingLayer(nn.Module): + def __init__(self): + super(ScalingLayer, self).__init__() + self.register_buffer('shift', torch.Tensor([-.030, -.088, -.188])[None, :, None, None]) + self.register_buffer('scale', torch.Tensor([.458, .448, .450])[None, :, None, None]) + + def forward(self, inp): + return (inp - self.shift) / self.scale + + +class NetLinLayer(nn.Module): + """ A single linear layer which does a 1x1 conv """ + def __init__(self, chn_in, chn_out=1, use_dropout=False): + super(NetLinLayer, self).__init__() + layers = [nn.Dropout(), ] if (use_dropout) else [] + layers += [nn.Conv2d(chn_in, chn_out, 1, stride=1, padding=0, bias=False), ] + self.model = nn.Sequential(*layers) + + +class vgg16(torch.nn.Module): + def __init__(self, requires_grad=False, pretrained=True): + super(vgg16, self).__init__() + vgg_pretrained_features = models.vgg16(pretrained=pretrained).features + self.slice1 = torch.nn.Sequential() + self.slice2 = torch.nn.Sequential() + self.slice3 = torch.nn.Sequential() + self.slice4 = torch.nn.Sequential() + self.slice5 = torch.nn.Sequential() + self.N_slices = 5 + for x in range(4): + self.slice1.add_module(str(x), vgg_pretrained_features[x]) + for x in range(4, 9): + self.slice2.add_module(str(x), vgg_pretrained_features[x]) + for x in range(9, 16): + self.slice3.add_module(str(x), vgg_pretrained_features[x]) + for x in range(16, 23): + self.slice4.add_module(str(x), vgg_pretrained_features[x]) + for x in range(23, 30): + self.slice5.add_module(str(x), vgg_pretrained_features[x]) + if not requires_grad: + for param in self.parameters(): + param.requires_grad = False + + def forward(self, X): + h = self.slice1(X) + h_relu1_2 = h + h = self.slice2(h) + h_relu2_2 = h + h = self.slice3(h) + h_relu3_3 = h + h = self.slice4(h) + h_relu4_3 = h + h = self.slice5(h) + h_relu5_3 = h + vgg_outputs = namedtuple("VggOutputs", ['relu1_2', 'relu2_2', 'relu3_3', 'relu4_3', 'relu5_3']) + out = vgg_outputs(h_relu1_2, h_relu2_2, h_relu3_3, h_relu4_3, h_relu5_3) + return out + + +def normalize_tensor(x,eps=1e-10): + norm_factor = torch.sqrt(torch.sum(x**2,dim=1,keepdim=True)) + return x/(norm_factor+eps) + + +def spatial_average(x, keepdim=True): + return x.mean([2,3],keepdim=keepdim) + diff --git a/taming/modules/losses/segmentation.py b/taming/modules/losses/segmentation.py new file mode 100644 index 0000000..4ba77de --- /dev/null +++ b/taming/modules/losses/segmentation.py @@ -0,0 +1,22 @@ +import torch.nn as nn +import torch.nn.functional as F + + +class BCELoss(nn.Module): + def forward(self, prediction, target): + loss = F.binary_cross_entropy_with_logits(prediction,target) + return loss, {} + + +class BCELossWithQuant(nn.Module): + def __init__(self, codebook_weight=1.): + super().__init__() + self.codebook_weight = codebook_weight + + def forward(self, qloss, target, prediction, split): + bce_loss = F.binary_cross_entropy_with_logits(prediction,target) + loss = bce_loss + self.codebook_weight*qloss + return loss, {"{}/total_loss".format(split): loss.clone().detach().mean(), + "{}/bce_loss".format(split): bce_loss.detach().mean(), + "{}/quant_loss".format(split): qloss.detach().mean() + } diff --git a/taming/modules/losses/vqperceptual.py b/taming/modules/losses/vqperceptual.py new file mode 100644 index 0000000..c2febd4 --- /dev/null +++ b/taming/modules/losses/vqperceptual.py @@ -0,0 +1,136 @@ +import torch +import torch.nn as nn +import torch.nn.functional as F + +from taming.modules.losses.lpips import LPIPS +from taming.modules.discriminator.model import NLayerDiscriminator, weights_init + + +class DummyLoss(nn.Module): + def __init__(self): + super().__init__() + + +def adopt_weight(weight, global_step, threshold=0, value=0.): + if global_step < threshold: + weight = value + return weight + + +def hinge_d_loss(logits_real, logits_fake): + loss_real = torch.mean(F.relu(1. - logits_real)) + loss_fake = torch.mean(F.relu(1. + logits_fake)) + d_loss = 0.5 * (loss_real + loss_fake) + return d_loss + + +def vanilla_d_loss(logits_real, logits_fake): + d_loss = 0.5 * ( + torch.mean(torch.nn.functional.softplus(-logits_real)) + + torch.mean(torch.nn.functional.softplus(logits_fake))) + return d_loss + + +class VQLPIPSWithDiscriminator(nn.Module): + def __init__(self, disc_start, codebook_weight=1.0, pixelloss_weight=1.0, + disc_num_layers=3, disc_in_channels=3, disc_factor=1.0, disc_weight=1.0, + perceptual_weight=1.0, use_actnorm=False, disc_conditional=False, + disc_ndf=64, disc_loss="hinge"): + super().__init__() + assert disc_loss in ["hinge", "vanilla"] + self.codebook_weight = codebook_weight + self.pixel_weight = pixelloss_weight + self.perceptual_loss = LPIPS().eval() + self.perceptual_weight = perceptual_weight + + self.discriminator = NLayerDiscriminator(input_nc=disc_in_channels, + n_layers=disc_num_layers, + use_actnorm=use_actnorm, + ndf=disc_ndf + ).apply(weights_init) + self.discriminator_iter_start = disc_start + if disc_loss == "hinge": + self.disc_loss = hinge_d_loss + elif disc_loss == "vanilla": + self.disc_loss = vanilla_d_loss + else: + raise ValueError(f"Unknown GAN loss '{disc_loss}'.") + print(f"VQLPIPSWithDiscriminator running with {disc_loss} loss.") + self.disc_factor = disc_factor + self.discriminator_weight = disc_weight + self.disc_conditional = disc_conditional + + def calculate_adaptive_weight(self, nll_loss, g_loss, last_layer=None): + if last_layer is not None: + nll_grads = torch.autograd.grad(nll_loss, last_layer, retain_graph=True)[0] + g_grads = torch.autograd.grad(g_loss, last_layer, retain_graph=True)[0] + else: + nll_grads = torch.autograd.grad(nll_loss, self.last_layer[0], retain_graph=True)[0] + g_grads = torch.autograd.grad(g_loss, self.last_layer[0], retain_graph=True)[0] + + d_weight = torch.norm(nll_grads) / (torch.norm(g_grads) + 1e-4) + d_weight = torch.clamp(d_weight, 0.0, 1e4).detach() + d_weight = d_weight * self.discriminator_weight + return d_weight + + def forward(self, codebook_loss, inputs, reconstructions, optimizer_idx, + global_step, last_layer=None, cond=None, split="train"): + rec_loss = torch.abs(inputs.contiguous() - reconstructions.contiguous()) + if self.perceptual_weight > 0: + p_loss = self.perceptual_loss(inputs.contiguous(), reconstructions.contiguous()) + rec_loss = rec_loss + self.perceptual_weight * p_loss + else: + p_loss = torch.tensor([0.0]) + + nll_loss = rec_loss + #nll_loss = torch.sum(nll_loss) / nll_loss.shape[0] + nll_loss = torch.mean(nll_loss) + + # now the GAN part + if optimizer_idx == 0: + # generator update + if cond is None: + assert not self.disc_conditional + logits_fake = self.discriminator(reconstructions.contiguous()) + else: + assert self.disc_conditional + logits_fake = self.discriminator(torch.cat((reconstructions.contiguous(), cond), dim=1)) + g_loss = -torch.mean(logits_fake) + + try: + d_weight = self.calculate_adaptive_weight(nll_loss, g_loss, last_layer=last_layer) + except RuntimeError: + assert not self.training + d_weight = torch.tensor(0.0) + + disc_factor = adopt_weight(self.disc_factor, global_step, threshold=self.discriminator_iter_start) + loss = nll_loss + d_weight * disc_factor * g_loss + self.codebook_weight * codebook_loss.mean() + + log = {"{}/total_loss".format(split): loss.clone().detach().mean(), + "{}/quant_loss".format(split): codebook_loss.detach().mean(), + "{}/nll_loss".format(split): nll_loss.detach().mean(), + "{}/rec_loss".format(split): rec_loss.detach().mean(), + "{}/p_loss".format(split): p_loss.detach().mean(), + "{}/d_weight".format(split): d_weight.detach(), + "{}/disc_factor".format(split): torch.tensor(disc_factor), + "{}/g_loss".format(split): g_loss.detach().mean(), + } + return loss, log + + if optimizer_idx == 1: + # second pass for discriminator update + if cond is None: + logits_real = self.discriminator(inputs.contiguous().detach()) + logits_fake = self.discriminator(reconstructions.contiguous().detach()) + else: + logits_real = self.discriminator(torch.cat((inputs.contiguous().detach(), cond), dim=1)) + logits_fake = self.discriminator(torch.cat((reconstructions.contiguous().detach(), cond), dim=1)) + + disc_factor = adopt_weight(self.disc_factor, global_step, threshold=self.discriminator_iter_start) + d_loss = disc_factor * self.disc_loss(logits_real, logits_fake) + + log = {"{}/disc_loss".format(split): d_loss.clone().detach().mean(), + "{}/logits_real".format(split): logits_real.detach().mean(), + "{}/logits_fake".format(split): logits_fake.detach().mean() + } + return d_loss, log diff --git a/taming/modules/misc/coord.py b/taming/modules/misc/coord.py new file mode 100644 index 0000000..ee69b0c --- /dev/null +++ b/taming/modules/misc/coord.py @@ -0,0 +1,31 @@ +import torch + +class CoordStage(object): + def __init__(self, n_embed, down_factor): + self.n_embed = n_embed + self.down_factor = down_factor + + def eval(self): + return self + + def encode(self, c): + """fake vqmodel interface""" + assert 0.0 <= c.min() and c.max() <= 1.0 + b,ch,h,w = c.shape + assert ch == 1 + + c = torch.nn.functional.interpolate(c, scale_factor=1/self.down_factor, + mode="area") + c = c.clamp(0.0, 1.0) + c = self.n_embed*c + c_quant = c.round() + c_ind = c_quant.to(dtype=torch.long) + + info = None, None, c_ind + return c_quant, None, info + + def decode(self, c): + c = c/self.n_embed + c = torch.nn.functional.interpolate(c, scale_factor=self.down_factor, + mode="nearest") + return c diff --git a/taming/modules/transformer/mingpt.py b/taming/modules/transformer/mingpt.py new file mode 100644 index 0000000..d14b7b6 --- /dev/null +++ b/taming/modules/transformer/mingpt.py @@ -0,0 +1,415 @@ +""" +taken from: https://github.com/karpathy/minGPT/ +GPT model: +- the initial stem consists of a combination of token encoding and a positional encoding +- the meat of it is a uniform sequence of Transformer blocks + - each Transformer is a sequential combination of a 1-hidden-layer MLP block and a self-attention block + - all blocks feed into a central residual pathway similar to resnets +- the final decoder is a linear projection into a vanilla Softmax classifier +""" + +import math +import logging + +import torch +import torch.nn as nn +from torch.nn import functional as F +from transformers import top_k_top_p_filtering + +logger = logging.getLogger(__name__) + + +class GPTConfig: + """ base GPT config, params common to all GPT versions """ + embd_pdrop = 0.1 + resid_pdrop = 0.1 + attn_pdrop = 0.1 + + def __init__(self, vocab_size, block_size, **kwargs): + self.vocab_size = vocab_size + self.block_size = block_size + for k,v in kwargs.items(): + setattr(self, k, v) + + +class GPT1Config(GPTConfig): + """ GPT-1 like network roughly 125M params """ + n_layer = 12 + n_head = 12 + n_embd = 768 + + +class CausalSelfAttention(nn.Module): + """ + A vanilla multi-head masked self-attention layer with a projection at the end. + It is possible to use torch.nn.MultiheadAttention here but I am including an + explicit implementation here to show that there is nothing too scary here. + """ + + def __init__(self, config): + super().__init__() + assert config.n_embd % config.n_head == 0 + # key, query, value projections for all heads + self.key = nn.Linear(config.n_embd, config.n_embd) + self.query = nn.Linear(config.n_embd, config.n_embd) + self.value = nn.Linear(config.n_embd, config.n_embd) + # regularization + self.attn_drop = nn.Dropout(config.attn_pdrop) + self.resid_drop = nn.Dropout(config.resid_pdrop) + # output projection + self.proj = nn.Linear(config.n_embd, config.n_embd) + # causal mask to ensure that attention is only applied to the left in the input sequence + mask = torch.tril(torch.ones(config.block_size, + config.block_size)) + if hasattr(config, "n_unmasked"): + mask[:config.n_unmasked, :config.n_unmasked] = 1 + self.register_buffer("mask", mask.view(1, 1, config.block_size, config.block_size)) + self.n_head = config.n_head + + def forward(self, x, layer_past=None): + B, T, C = x.size() + + # calculate query, key, values for all heads in batch and move head forward to be the batch dim + k = self.key(x).view(B, T, self.n_head, C // self.n_head).transpose(1, 2) # (B, nh, T, hs) + q = self.query(x).view(B, T, self.n_head, C // self.n_head).transpose(1, 2) # (B, nh, T, hs) + v = self.value(x).view(B, T, self.n_head, C // self.n_head).transpose(1, 2) # (B, nh, T, hs) + + present = torch.stack((k, v)) + if layer_past is not None: + past_key, past_value = layer_past + k = torch.cat((past_key, k), dim=-2) + v = torch.cat((past_value, v), dim=-2) + + # causal self-attention; Self-attend: (B, nh, T, hs) x (B, nh, hs, T) -> (B, nh, T, T) + att = (q @ k.transpose(-2, -1)) * (1.0 / math.sqrt(k.size(-1))) + if layer_past is None: + att = att.masked_fill(self.mask[:,:,:T,:T] == 0, float('-inf')) + + att = F.softmax(att, dim=-1) + att = self.attn_drop(att) + y = att @ v # (B, nh, T, T) x (B, nh, T, hs) -> (B, nh, T, hs) + y = y.transpose(1, 2).contiguous().view(B, T, C) # re-assemble all head outputs side by side + + # output projection + y = self.resid_drop(self.proj(y)) + return y, present # TODO: check that this does not break anything + + +class Block(nn.Module): + """ an unassuming Transformer block """ + def __init__(self, config): + super().__init__() + self.ln1 = nn.LayerNorm(config.n_embd) + self.ln2 = nn.LayerNorm(config.n_embd) + self.attn = CausalSelfAttention(config) + self.mlp = nn.Sequential( + nn.Linear(config.n_embd, 4 * config.n_embd), + nn.GELU(), # nice + nn.Linear(4 * config.n_embd, config.n_embd), + nn.Dropout(config.resid_pdrop), + ) + + def forward(self, x, layer_past=None, return_present=False): + # TODO: check that training still works + if return_present: assert not self.training + # layer past: tuple of length two with B, nh, T, hs + attn, present = self.attn(self.ln1(x), layer_past=layer_past) + + x = x + attn + x = x + self.mlp(self.ln2(x)) + if layer_past is not None or return_present: + return x, present + return x + + +class GPT(nn.Module): + """ the full GPT language model, with a context size of block_size """ + def __init__(self, vocab_size, block_size, n_layer=12, n_head=8, n_embd=256, + embd_pdrop=0., resid_pdrop=0., attn_pdrop=0., n_unmasked=0): + super().__init__() + config = GPTConfig(vocab_size=vocab_size, block_size=block_size, + embd_pdrop=embd_pdrop, resid_pdrop=resid_pdrop, attn_pdrop=attn_pdrop, + n_layer=n_layer, n_head=n_head, n_embd=n_embd, + n_unmasked=n_unmasked) + # input embedding stem + self.tok_emb = nn.Embedding(config.vocab_size, config.n_embd) + self.pos_emb = nn.Parameter(torch.zeros(1, config.block_size, config.n_embd)) + self.drop = nn.Dropout(config.embd_pdrop) + # transformer + self.blocks = nn.Sequential(*[Block(config) for _ in range(config.n_layer)]) + # decoder head + self.ln_f = nn.LayerNorm(config.n_embd) + self.head = nn.Linear(config.n_embd, config.vocab_size, bias=False) + self.block_size = config.block_size + self.apply(self._init_weights) + self.config = config + logger.info("number of parameters: %e", sum(p.numel() for p in self.parameters())) + + def get_block_size(self): + return self.block_size + + def _init_weights(self, module): + if isinstance(module, (nn.Linear, nn.Embedding)): + module.weight.data.normal_(mean=0.0, std=0.02) + if isinstance(module, nn.Linear) and module.bias is not None: + module.bias.data.zero_() + elif isinstance(module, nn.LayerNorm): + module.bias.data.zero_() + module.weight.data.fill_(1.0) + + def forward(self, idx, embeddings=None, targets=None): + # forward the GPT model + token_embeddings = self.tok_emb(idx) # each index maps to a (learnable) vector + + if embeddings is not None: # prepend explicit embeddings + token_embeddings = torch.cat((embeddings, token_embeddings), dim=1) + + t = token_embeddings.shape[1] + assert t <= self.block_size, "Cannot forward, model block size is exhausted." + position_embeddings = self.pos_emb[:, :t, :] # each position maps to a (learnable) vector + x = self.drop(token_embeddings + position_embeddings) + x = self.blocks(x) + x = self.ln_f(x) + logits = self.head(x) + + # if we are given some desired targets also calculate the loss + loss = None + if targets is not None: + loss = F.cross_entropy(logits.view(-1, logits.size(-1)), targets.view(-1)) + + return logits, loss + + def forward_with_past(self, idx, embeddings=None, targets=None, past=None, past_length=None): + # inference only + assert not self.training + token_embeddings = self.tok_emb(idx) # each index maps to a (learnable) vector + if embeddings is not None: # prepend explicit embeddings + token_embeddings = torch.cat((embeddings, token_embeddings), dim=1) + + if past is not None: + assert past_length is not None + past = torch.cat(past, dim=-2) # n_layer, 2, b, nh, len_past, dim_head + past_shape = list(past.shape) + expected_shape = [self.config.n_layer, 2, idx.shape[0], self.config.n_head, past_length, self.config.n_embd//self.config.n_head] + assert past_shape == expected_shape, f"{past_shape} =/= {expected_shape}" + position_embeddings = self.pos_emb[:, past_length, :] # each position maps to a (learnable) vector + else: + position_embeddings = self.pos_emb[:, :token_embeddings.shape[1], :] + + x = self.drop(token_embeddings + position_embeddings) + presents = [] # accumulate over layers + for i, block in enumerate(self.blocks): + x, present = block(x, layer_past=past[i, ...] if past is not None else None, return_present=True) + presents.append(present) + + x = self.ln_f(x) + logits = self.head(x) + # if we are given some desired targets also calculate the loss + loss = None + if targets is not None: + loss = F.cross_entropy(logits.view(-1, logits.size(-1)), targets.view(-1)) + + return logits, loss, torch.stack(presents) # _, _, n_layer, 2, b, nh, 1, dim_head + + +class DummyGPT(nn.Module): + # for debugging + def __init__(self, add_value=1): + super().__init__() + self.add_value = add_value + + def forward(self, idx): + return idx + self.add_value, None + + +class CodeGPT(nn.Module): + """Takes in semi-embeddings""" + def __init__(self, vocab_size, block_size, in_channels, n_layer=12, n_head=8, n_embd=256, + embd_pdrop=0., resid_pdrop=0., attn_pdrop=0., n_unmasked=0): + super().__init__() + config = GPTConfig(vocab_size=vocab_size, block_size=block_size, + embd_pdrop=embd_pdrop, resid_pdrop=resid_pdrop, attn_pdrop=attn_pdrop, + n_layer=n_layer, n_head=n_head, n_embd=n_embd, + n_unmasked=n_unmasked) + # input embedding stem + self.tok_emb = nn.Linear(in_channels, config.n_embd) + self.pos_emb = nn.Parameter(torch.zeros(1, config.block_size, config.n_embd)) + self.drop = nn.Dropout(config.embd_pdrop) + # transformer + self.blocks = nn.Sequential(*[Block(config) for _ in range(config.n_layer)]) + # decoder head + self.ln_f = nn.LayerNorm(config.n_embd) + self.head = nn.Linear(config.n_embd, config.vocab_size, bias=False) + self.block_size = config.block_size + self.apply(self._init_weights) + self.config = config + logger.info("number of parameters: %e", sum(p.numel() for p in self.parameters())) + + def get_block_size(self): + return self.block_size + + def _init_weights(self, module): + if isinstance(module, (nn.Linear, nn.Embedding)): + module.weight.data.normal_(mean=0.0, std=0.02) + if isinstance(module, nn.Linear) and module.bias is not None: + module.bias.data.zero_() + elif isinstance(module, nn.LayerNorm): + module.bias.data.zero_() + module.weight.data.fill_(1.0) + + def forward(self, idx, embeddings=None, targets=None): + # forward the GPT model + token_embeddings = self.tok_emb(idx) # each index maps to a (learnable) vector + + if embeddings is not None: # prepend explicit embeddings + token_embeddings = torch.cat((embeddings, token_embeddings), dim=1) + + t = token_embeddings.shape[1] + assert t <= self.block_size, "Cannot forward, model block size is exhausted." + position_embeddings = self.pos_emb[:, :t, :] # each position maps to a (learnable) vector + x = self.drop(token_embeddings + position_embeddings) + x = self.blocks(x) + x = self.taming_cinln_f(x) + logits = self.head(x) + + # if we are given some desired targets also calculate the loss + loss = None + if targets is not None: + loss = F.cross_entropy(logits.view(-1, logits.size(-1)), targets.view(-1)) + + return logits, loss + + + +#### sampling utils + +def top_k_logits(logits, k): + v, ix = torch.topk(logits, k) + out = logits.clone() + out[out < v[:, [-1]]] = -float('Inf') + return out + +@torch.no_grad() +def sample(model, x, steps, temperature=1.0, sample=False, top_k=None): + """ + take a conditioning sequence of indices in x (of shape (b,t)) and predict the next token in + the sequence, feeding the predictions back into the model each time. Clearly the sampling + has quadratic complexity unlike an RNN that is only linear, and has a finite context window + of block_size, unlike an RNN that has an infinite context window. + """ + block_size = model.get_block_size() + model.eval() + for k in range(steps): + x_cond = x if x.size(1) <= block_size else x[:, -block_size:] # crop context if needed + logits, _ = model(x_cond) + # pluck the logits at the final step and scale by temperature + logits = logits[:, -1, :] / temperature + # optionally crop probabilities to only the top k options + if top_k is not None: + logits = top_k_logits(logits, top_k) + # apply softmax to convert to probabilities + probs = F.softmax(logits, dim=-1) + # sample from the distribution or take the most likely + if sample: + ix = torch.multinomial(probs, num_samples=1) + else: + _, ix = torch.topk(probs, k=1, dim=-1) + # append to the sequence and continue + x = torch.cat((x, ix), dim=1) + + return x + + +@torch.no_grad() +def sample_with_past(x, model, steps, temperature=1., sample_logits=True, + top_k=None, top_p=None, callback=None): + # x is conditioning + sample = x + cond_len = x.shape[1] + past = None + for n in range(steps): + if callback is not None: + callback(n) + logits, _, present = model.forward_with_past(x, past=past, past_length=(n+cond_len-1)) + if past is None: + past = [present] + else: + past.append(present) + logits = logits[:, -1, :] / temperature + if top_k is not None: + logits = top_k_top_p_filtering(logits, top_k=top_k, top_p=top_p) + + probs = F.softmax(logits, dim=-1) + if not sample_logits: + _, x = torch.topk(probs, k=1, dim=-1) + else: + x = torch.multinomial(probs, num_samples=1) + # append to the sequence and continue + sample = torch.cat((sample, x), dim=1) + del past + sample = sample[:, cond_len:] # cut conditioning off + return sample + + +#### clustering utils + +class KMeans(nn.Module): + def __init__(self, ncluster=512, nc=3, niter=10): + super().__init__() + self.ncluster = ncluster + self.nc = nc + self.niter = niter + self.shape = (3,32,32) + self.register_buffer("C", torch.zeros(self.ncluster,nc)) + self.register_buffer('initialized', torch.tensor(0, dtype=torch.uint8)) + + def is_initialized(self): + return self.initialized.item() == 1 + + @torch.no_grad() + def initialize(self, x): + N, D = x.shape + assert D == self.nc, D + c = x[torch.randperm(N)[:self.ncluster]] # init clusters at random + for i in range(self.niter): + # assign all pixels to the closest codebook element + a = ((x[:, None, :] - c[None, :, :])**2).sum(-1).argmin(1) + # move each codebook element to be the mean of the pixels that assigned to it + c = torch.stack([x[a==k].mean(0) for k in range(self.ncluster)]) + # re-assign any poorly positioned codebook elements + nanix = torch.any(torch.isnan(c), dim=1) + ndead = nanix.sum().item() + print('done step %d/%d, re-initialized %d dead clusters' % (i+1, self.niter, ndead)) + c[nanix] = x[torch.randperm(N)[:ndead]] # re-init dead clusters + + self.C.copy_(c) + self.initialized.fill_(1) + + + def forward(self, x, reverse=False, shape=None): + if not reverse: + # flatten + bs,c,h,w = x.shape + assert c == self.nc + x = x.reshape(bs,c,h*w,1) + C = self.C.permute(1,0) + C = C.reshape(1,c,1,self.ncluster) + a = ((x-C)**2).sum(1).argmin(-1) # bs, h*w indices + return a + else: + # flatten + bs, HW = x.shape + """ + c = self.C.reshape( 1, self.nc, 1, self.ncluster) + c = c[bs*[0],:,:,:] + c = c[:,:,HW*[0],:] + x = x.reshape(bs, 1, HW, 1) + x = x[:,3*[0],:,:] + x = torch.gather(c, dim=3, index=x) + """ + x = self.C[x] + x = x.permute(0,2,1) + shape = shape if shape is not None else self.shape + x = x.reshape(bs, *shape) + + return x diff --git a/taming/modules/transformer/permuter.py b/taming/modules/transformer/permuter.py new file mode 100644 index 0000000..0d43bb1 --- /dev/null +++ b/taming/modules/transformer/permuter.py @@ -0,0 +1,248 @@ +import torch +import torch.nn as nn +import numpy as np + + +class AbstractPermuter(nn.Module): + def __init__(self, *args, **kwargs): + super().__init__() + def forward(self, x, reverse=False): + raise NotImplementedError + + +class Identity(AbstractPermuter): + def __init__(self): + super().__init__() + + def forward(self, x, reverse=False): + return x + + +class Subsample(AbstractPermuter): + def __init__(self, H, W): + super().__init__() + C = 1 + indices = np.arange(H*W).reshape(C,H,W) + while min(H, W) > 1: + indices = indices.reshape(C,H//2,2,W//2,2) + indices = indices.transpose(0,2,4,1,3) + indices = indices.reshape(C*4,H//2, W//2) + H = H//2 + W = W//2 + C = C*4 + assert H == W == 1 + idx = torch.tensor(indices.ravel()) + self.register_buffer('forward_shuffle_idx', + nn.Parameter(idx, requires_grad=False)) + self.register_buffer('backward_shuffle_idx', + nn.Parameter(torch.argsort(idx), requires_grad=False)) + + def forward(self, x, reverse=False): + if not reverse: + return x[:, self.forward_shuffle_idx] + else: + return x[:, self.backward_shuffle_idx] + + +def mortonify(i, j): + """(i,j) index to linear morton code""" + i = np.uint64(i) + j = np.uint64(j) + + z = np.uint(0) + + for pos in range(32): + z = (z | + ((j & (np.uint64(1) << np.uint64(pos))) << np.uint64(pos)) | + ((i & (np.uint64(1) << np.uint64(pos))) << np.uint64(pos+1)) + ) + return z + + +class ZCurve(AbstractPermuter): + def __init__(self, H, W): + super().__init__() + reverseidx = [np.int64(mortonify(i,j)) for i in range(H) for j in range(W)] + idx = np.argsort(reverseidx) + idx = torch.tensor(idx) + reverseidx = torch.tensor(reverseidx) + self.register_buffer('forward_shuffle_idx', + idx) + self.register_buffer('backward_shuffle_idx', + reverseidx) + + def forward(self, x, reverse=False): + if not reverse: + return x[:, self.forward_shuffle_idx] + else: + return x[:, self.backward_shuffle_idx] + + +class SpiralOut(AbstractPermuter): + def __init__(self, H, W): + super().__init__() + assert H == W + size = W + indices = np.arange(size*size).reshape(size,size) + + i0 = size//2 + j0 = size//2-1 + + i = i0 + j = j0 + + idx = [indices[i0, j0]] + step_mult = 0 + for c in range(1, size//2+1): + step_mult += 1 + # steps left + for k in range(step_mult): + i = i - 1 + j = j + idx.append(indices[i, j]) + + # step down + for k in range(step_mult): + i = i + j = j + 1 + idx.append(indices[i, j]) + + step_mult += 1 + if c < size//2: + # step right + for k in range(step_mult): + i = i + 1 + j = j + idx.append(indices[i, j]) + + # step up + for k in range(step_mult): + i = i + j = j - 1 + idx.append(indices[i, j]) + else: + # end reached + for k in range(step_mult-1): + i = i + 1 + idx.append(indices[i, j]) + + assert len(idx) == size*size + idx = torch.tensor(idx) + self.register_buffer('forward_shuffle_idx', idx) + self.register_buffer('backward_shuffle_idx', torch.argsort(idx)) + + def forward(self, x, reverse=False): + if not reverse: + return x[:, self.forward_shuffle_idx] + else: + return x[:, self.backward_shuffle_idx] + + +class SpiralIn(AbstractPermuter): + def __init__(self, H, W): + super().__init__() + assert H == W + size = W + indices = np.arange(size*size).reshape(size,size) + + i0 = size//2 + j0 = size//2-1 + + i = i0 + j = j0 + + idx = [indices[i0, j0]] + step_mult = 0 + for c in range(1, size//2+1): + step_mult += 1 + # steps left + for k in range(step_mult): + i = i - 1 + j = j + idx.append(indices[i, j]) + + # step down + for k in range(step_mult): + i = i + j = j + 1 + idx.append(indices[i, j]) + + step_mult += 1 + if c < size//2: + # step right + for k in range(step_mult): + i = i + 1 + j = j + idx.append(indices[i, j]) + + # step up + for k in range(step_mult): + i = i + j = j - 1 + idx.append(indices[i, j]) + else: + # end reached + for k in range(step_mult-1): + i = i + 1 + idx.append(indices[i, j]) + + assert len(idx) == size*size + idx = idx[::-1] + idx = torch.tensor(idx) + self.register_buffer('forward_shuffle_idx', idx) + self.register_buffer('backward_shuffle_idx', torch.argsort(idx)) + + def forward(self, x, reverse=False): + if not reverse: + return x[:, self.forward_shuffle_idx] + else: + return x[:, self.backward_shuffle_idx] + + +class Random(nn.Module): + def __init__(self, H, W): + super().__init__() + indices = np.random.RandomState(1).permutation(H*W) + idx = torch.tensor(indices.ravel()) + self.register_buffer('forward_shuffle_idx', idx) + self.register_buffer('backward_shuffle_idx', torch.argsort(idx)) + + def forward(self, x, reverse=False): + if not reverse: + return x[:, self.forward_shuffle_idx] + else: + return x[:, self.backward_shuffle_idx] + + +class AlternateParsing(AbstractPermuter): + def __init__(self, H, W): + super().__init__() + indices = np.arange(W*H).reshape(H,W) + for i in range(1, H, 2): + indices[i, :] = indices[i, ::-1] + idx = indices.flatten() + assert len(idx) == H*W + idx = torch.tensor(idx) + self.register_buffer('forward_shuffle_idx', idx) + self.register_buffer('backward_shuffle_idx', torch.argsort(idx)) + + def forward(self, x, reverse=False): + if not reverse: + return x[:, self.forward_shuffle_idx] + else: + return x[:, self.backward_shuffle_idx] + + +if __name__ == "__main__": + p0 = AlternateParsing(16, 16) + print(p0.forward_shuffle_idx) + print(p0.backward_shuffle_idx) + + x = torch.randint(0, 768, size=(11, 256)) + y = p0(x) + xre = p0(y, reverse=True) + assert torch.equal(x, xre) + + p1 = SpiralOut(2, 2) + print(p1.forward_shuffle_idx) + print(p1.backward_shuffle_idx) diff --git a/taming/modules/util.py b/taming/modules/util.py new file mode 100644 index 0000000..9ee1638 --- /dev/null +++ b/taming/modules/util.py @@ -0,0 +1,130 @@ +import torch +import torch.nn as nn + + +def count_params(model): + total_params = sum(p.numel() for p in model.parameters()) + return total_params + + +class ActNorm(nn.Module): + def __init__(self, num_features, logdet=False, affine=True, + allow_reverse_init=False): + assert affine + super().__init__() + self.logdet = logdet + self.loc = nn.Parameter(torch.zeros(1, num_features, 1, 1)) + self.scale = nn.Parameter(torch.ones(1, num_features, 1, 1)) + self.allow_reverse_init = allow_reverse_init + + self.register_buffer('initialized', torch.tensor(0, dtype=torch.uint8)) + + def initialize(self, input): + with torch.no_grad(): + flatten = input.permute(1, 0, 2, 3).contiguous().view(input.shape[1], -1) + mean = ( + flatten.mean(1) + .unsqueeze(1) + .unsqueeze(2) + .unsqueeze(3) + .permute(1, 0, 2, 3) + ) + std = ( + flatten.std(1) + .unsqueeze(1) + .unsqueeze(2) + .unsqueeze(3) + .permute(1, 0, 2, 3) + ) + + self.loc.data.copy_(-mean) + self.scale.data.copy_(1 / (std + 1e-6)) + + def forward(self, input, reverse=False): + if reverse: + return self.reverse(input) + if len(input.shape) == 2: + input = input[:,:,None,None] + squeeze = True + else: + squeeze = False + + _, _, height, width = input.shape + + if self.training and self.initialized.item() == 0: + self.initialize(input) + self.initialized.fill_(1) + + h = self.scale * (input + self.loc) + + if squeeze: + h = h.squeeze(-1).squeeze(-1) + + if self.logdet: + log_abs = torch.log(torch.abs(self.scale)) + logdet = height*width*torch.sum(log_abs) + logdet = logdet * torch.ones(input.shape[0]).to(input) + return h, logdet + + return h + + def reverse(self, output): + if self.training and self.initialized.item() == 0: + if not self.allow_reverse_init: + raise RuntimeError( + "Initializing ActNorm in reverse direction is " + "disabled by default. Use allow_reverse_init=True to enable." + ) + else: + self.initialize(output) + self.initialized.fill_(1) + + if len(output.shape) == 2: + output = output[:,:,None,None] + squeeze = True + else: + squeeze = False + + h = output / self.scale - self.loc + + if squeeze: + h = h.squeeze(-1).squeeze(-1) + return h + + +class AbstractEncoder(nn.Module): + def __init__(self): + super().__init__() + + def encode(self, *args, **kwargs): + raise NotImplementedError + + +class Labelator(AbstractEncoder): + """Net2Net Interface for Class-Conditional Model""" + def __init__(self, n_classes, quantize_interface=True): + super().__init__() + self.n_classes = n_classes + self.quantize_interface = quantize_interface + + def encode(self, c): + c = c[:,None] + if self.quantize_interface: + return c, None, [None, None, c.long()] + return c + + +class SOSProvider(AbstractEncoder): + # for unconditional training + def __init__(self, sos_token, quantize_interface=True): + super().__init__() + self.sos_token = sos_token + self.quantize_interface = quantize_interface + + def encode(self, x): + # get batch size from data and replicate sos_token + c = torch.ones(x.shape[0], 1)*self.sos_token + c = c.long().to(x.device) + if self.quantize_interface: + return c, None, [None, None, c] + return c diff --git a/taming/modules/vqvae/quantize.py b/taming/modules/vqvae/quantize.py new file mode 100644 index 0000000..d75544e --- /dev/null +++ b/taming/modules/vqvae/quantize.py @@ -0,0 +1,445 @@ +import torch +import torch.nn as nn +import torch.nn.functional as F +import numpy as np +from torch import einsum +from einops import rearrange + + +class VectorQuantizer(nn.Module): + """ + see https://github.com/MishaLaskin/vqvae/blob/d761a999e2267766400dc646d82d3ac3657771d4/models/quantizer.py + ____________________________________________ + Discretization bottleneck part of the VQ-VAE. + Inputs: + - n_e : number of embeddings + - e_dim : dimension of embedding + - beta : commitment cost used in loss term, beta * ||z_e(x)-sg[e]||^2 + _____________________________________________ + """ + + # NOTE: this class contains a bug regarding beta; see VectorQuantizer2 for + # a fix and use legacy=False to apply that fix. VectorQuantizer2 can be + # used wherever VectorQuantizer has been used before and is additionally + # more efficient. + def __init__(self, n_e, e_dim, beta): + super(VectorQuantizer, self).__init__() + self.n_e = n_e + self.e_dim = e_dim + self.beta = beta + + self.embedding = nn.Embedding(self.n_e, self.e_dim) + self.embedding.weight.data.uniform_(-1.0 / self.n_e, 1.0 / self.n_e) + + def forward(self, z): + """ + Inputs the output of the encoder network z and maps it to a discrete + one-hot vector that is the index of the closest embedding vector e_j + z (continuous) -> z_q (discrete) + z.shape = (batch, channel, height, width) + quantization pipeline: + 1. get encoder input (B,C,H,W) + 2. flatten input to (B*H*W,C) + """ + # reshape z -> (batch, height, width, channel) and flatten + z = z.permute(0, 2, 3, 1).contiguous() + z_flattened = z.view(-1, self.e_dim) + # distances from z to embeddings e_j (z - e)^2 = z^2 + e^2 - 2 e * z + + d = torch.sum(z_flattened ** 2, dim=1, keepdim=True) + \ + torch.sum(self.embedding.weight**2, dim=1) - 2 * \ + torch.matmul(z_flattened, self.embedding.weight.t()) + + ## could possible replace this here + # #\start... + # find closest encodings + min_encoding_indices = torch.argmin(d, dim=1).unsqueeze(1) + + min_encodings = torch.zeros( + min_encoding_indices.shape[0], self.n_e).to(z) + min_encodings.scatter_(1, min_encoding_indices, 1) + + # dtype min encodings: torch.float32 + # min_encodings shape: torch.Size([2048, 512]) + # min_encoding_indices.shape: torch.Size([2048, 1]) + + # get quantized latent vectors + z_q = torch.matmul(min_encodings, self.embedding.weight).view(z.shape) + #.........\end + + # with: + # .........\start + #min_encoding_indices = torch.argmin(d, dim=1) + #z_q = self.embedding(min_encoding_indices) + # ......\end......... (TODO) + + # compute loss for embedding + loss = torch.mean((z_q.detach()-z)**2) + self.beta * \ + torch.mean((z_q - z.detach()) ** 2) + + # preserve gradients + z_q = z + (z_q - z).detach() + + # perplexity + e_mean = torch.mean(min_encodings, dim=0) + perplexity = torch.exp(-torch.sum(e_mean * torch.log(e_mean + 1e-10))) + + # reshape back to match original input shape + z_q = z_q.permute(0, 3, 1, 2).contiguous() + + return z_q, loss, (perplexity, min_encodings, min_encoding_indices) + + def get_codebook_entry(self, indices, shape): + # shape specifying (batch, height, width, channel) + # TODO: check for more easy handling with nn.Embedding + min_encodings = torch.zeros(indices.shape[0], self.n_e).to(indices) + min_encodings.scatter_(1, indices[:,None], 1) + + # get quantized latent vectors + z_q = torch.matmul(min_encodings.float(), self.embedding.weight) + + if shape is not None: + z_q = z_q.view(shape) + + # reshape back to match original input shape + z_q = z_q.permute(0, 3, 1, 2).contiguous() + + return z_q + + +class GumbelQuantize(nn.Module): + """ + credit to @karpathy: https://github.com/karpathy/deep-vector-quantization/blob/main/model.py (thanks!) + Gumbel Softmax trick quantizer + Categorical Reparameterization with Gumbel-Softmax, Jang et al. 2016 + https://arxiv.org/abs/1611.01144 + """ + def __init__(self, num_hiddens, embedding_dim, n_embed, straight_through=True, + kl_weight=5e-4, temp_init=1.0, use_vqinterface=True, + remap=None, unknown_index="random"): + super().__init__() + + self.embedding_dim = embedding_dim + self.n_embed = n_embed + + self.straight_through = straight_through + self.temperature = temp_init + self.kl_weight = kl_weight + + self.proj = nn.Conv2d(num_hiddens, n_embed, 1) + self.embed = nn.Embedding(n_embed, embedding_dim) + + self.use_vqinterface = use_vqinterface + + self.remap = remap + if self.remap is not None: + self.register_buffer("used", torch.tensor(np.load(self.remap))) + self.re_embed = self.used.shape[0] + self.unknown_index = unknown_index # "random" or "extra" or integer + if self.unknown_index == "extra": + self.unknown_index = self.re_embed + self.re_embed = self.re_embed+1 + print(f"Remapping {self.n_embed} indices to {self.re_embed} indices. " + f"Using {self.unknown_index} for unknown indices.") + else: + self.re_embed = n_embed + + def remap_to_used(self, inds): + ishape = inds.shape + assert len(ishape)>1 + inds = inds.reshape(ishape[0],-1) + used = self.used.to(inds) + match = (inds[:,:,None]==used[None,None,...]).long() + new = match.argmax(-1) + unknown = match.sum(2)<1 + if self.unknown_index == "random": + new[unknown]=torch.randint(0,self.re_embed,size=new[unknown].shape).to(device=new.device) + else: + new[unknown] = self.unknown_index + return new.reshape(ishape) + + def unmap_to_all(self, inds): + ishape = inds.shape + assert len(ishape)>1 + inds = inds.reshape(ishape[0],-1) + used = self.used.to(inds) + if self.re_embed > self.used.shape[0]: # extra token + inds[inds>=self.used.shape[0]] = 0 # simply set to zero + back=torch.gather(used[None,:][inds.shape[0]*[0],:], 1, inds) + return back.reshape(ishape) + + def forward(self, z, temp=None, return_logits=False): + # force hard = True when we are in eval mode, as we must quantize. actually, always true seems to work + hard = self.straight_through if self.training else True + temp = self.temperature if temp is None else temp + + logits = self.proj(z) + if self.remap is not None: + # continue only with used logits + full_zeros = torch.zeros_like(logits) + logits = logits[:,self.used,...] + + soft_one_hot = F.gumbel_softmax(logits, tau=temp, dim=1, hard=hard) + if self.remap is not None: + # go back to all entries but unused set to zero + full_zeros[:,self.used,...] = soft_one_hot + soft_one_hot = full_zeros + z_q = einsum('b n h w, n d -> b d h w', soft_one_hot, self.embed.weight) + + # + kl divergence to the prior loss + qy = F.softmax(logits, dim=1) + diff = self.kl_weight * torch.sum(qy * torch.log(qy * self.n_embed + 1e-10), dim=1).mean() + + ind = soft_one_hot.argmax(dim=1) + if self.remap is not None: + ind = self.remap_to_used(ind) + if self.use_vqinterface: + if return_logits: + return z_q, diff, (None, None, ind), logits + return z_q, diff, (None, None, ind) + return z_q, diff, ind + + def get_codebook_entry(self, indices, shape): + b, h, w, c = shape + assert b*h*w == indices.shape[0] + indices = rearrange(indices, '(b h w) -> b h w', b=b, h=h, w=w) + if self.remap is not None: + indices = self.unmap_to_all(indices) + one_hot = F.one_hot(indices, num_classes=self.n_embed).permute(0, 3, 1, 2).float() + z_q = einsum('b n h w, n d -> b d h w', one_hot, self.embed.weight) + return z_q + + +class VectorQuantizer2(nn.Module): + """ + Improved version over VectorQuantizer, can be used as a drop-in replacement. Mostly + avoids costly matrix multiplications and allows for post-hoc remapping of indices. + """ + # NOTE: due to a bug the beta term was applied to the wrong term. for + # backwards compatibility we use the buggy version by default, but you can + # specify legacy=False to fix it. + def __init__(self, n_e, e_dim, beta, remap=None, unknown_index="random", + sane_index_shape=False, legacy=True): + super().__init__() + self.n_e = n_e + self.e_dim = e_dim + self.beta = beta + self.legacy = legacy + + self.embedding = nn.Embedding(self.n_e, self.e_dim) + self.embedding.weight.data.uniform_(-1.0 / self.n_e, 1.0 / self.n_e) + + self.remap = remap + if self.remap is not None: + self.register_buffer("used", torch.tensor(np.load(self.remap))) + self.re_embed = self.used.shape[0] + self.unknown_index = unknown_index # "random" or "extra" or integer + if self.unknown_index == "extra": + self.unknown_index = self.re_embed + self.re_embed = self.re_embed+1 + print(f"Remapping {self.n_e} indices to {self.re_embed} indices. " + f"Using {self.unknown_index} for unknown indices.") + else: + self.re_embed = n_e + + self.sane_index_shape = sane_index_shape + + def remap_to_used(self, inds): + ishape = inds.shape + assert len(ishape)>1 + inds = inds.reshape(ishape[0],-1) + used = self.used.to(inds) + match = (inds[:,:,None]==used[None,None,...]).long() + new = match.argmax(-1) + unknown = match.sum(2)<1 + if self.unknown_index == "random": + new[unknown]=torch.randint(0,self.re_embed,size=new[unknown].shape).to(device=new.device) + else: + new[unknown] = self.unknown_index + return new.reshape(ishape) + + def unmap_to_all(self, inds): + ishape = inds.shape + assert len(ishape)>1 + inds = inds.reshape(ishape[0],-1) + used = self.used.to(inds) + if self.re_embed > self.used.shape[0]: # extra token + inds[inds>=self.used.shape[0]] = 0 # simply set to zero + back=torch.gather(used[None,:][inds.shape[0]*[0],:], 1, inds) + return back.reshape(ishape) + + def forward(self, z, temp=None, rescale_logits=False, return_logits=False): + assert temp is None or temp==1.0, "Only for interface compatible with Gumbel" + assert rescale_logits==False, "Only for interface compatible with Gumbel" + assert return_logits==False, "Only for interface compatible with Gumbel" + # reshape z -> (batch, height, width, channel) and flatten + z = rearrange(z, 'b c h w -> b h w c').contiguous() + z_flattened = z.view(-1, self.e_dim) + # distances from z to embeddings e_j (z - e)^2 = z^2 + e^2 - 2 e * z + + d = torch.sum(z_flattened ** 2, dim=1, keepdim=True) + \ + torch.sum(self.embedding.weight**2, dim=1) - 2 * \ + torch.einsum('bd,dn->bn', z_flattened, rearrange(self.embedding.weight, 'n d -> d n')) + + min_encoding_indices = torch.argmin(d, dim=1) + z_q = self.embedding(min_encoding_indices).view(z.shape) + perplexity = None + min_encodings = None + + # compute loss for embedding + if not self.legacy: + loss = self.beta * torch.mean((z_q.detach()-z)**2) + \ + torch.mean((z_q - z.detach()) ** 2) + else: + loss = torch.mean((z_q.detach()-z)**2) + self.beta * \ + torch.mean((z_q - z.detach()) ** 2) + + # preserve gradients + z_q = z + (z_q - z).detach() + + # reshape back to match original input shape + z_q = rearrange(z_q, 'b h w c -> b c h w').contiguous() + + if self.remap is not None: + min_encoding_indices = min_encoding_indices.reshape(z.shape[0],-1) # add batch axis + min_encoding_indices = self.remap_to_used(min_encoding_indices) + min_encoding_indices = min_encoding_indices.reshape(-1,1) # flatten + + if self.sane_index_shape: + min_encoding_indices = min_encoding_indices.reshape( + z_q.shape[0], z_q.shape[2], z_q.shape[3]) + + return z_q, loss, (perplexity, min_encodings, min_encoding_indices) + + def get_codebook_entry(self, indices, shape): + # shape specifying (batch, height, width, channel) + if self.remap is not None: + indices = indices.reshape(shape[0],-1) # add batch axis + indices = self.unmap_to_all(indices) + indices = indices.reshape(-1) # flatten again + + # get quantized latent vectors + z_q = self.embedding(indices) + + if shape is not None: + z_q = z_q.view(shape) + # reshape back to match original input shape + z_q = z_q.permute(0, 3, 1, 2).contiguous() + + return z_q + +class EmbeddingEMA(nn.Module): + def __init__(self, num_tokens, codebook_dim, decay=0.99, eps=1e-5): + super().__init__() + self.decay = decay + self.eps = eps + weight = torch.randn(num_tokens, codebook_dim) + self.weight = nn.Parameter(weight, requires_grad = False) + self.cluster_size = nn.Parameter(torch.zeros(num_tokens), requires_grad = False) + self.embed_avg = nn.Parameter(weight.clone(), requires_grad = False) + self.update = True + + def forward(self, embed_id): + return F.embedding(embed_id, self.weight) + + def cluster_size_ema_update(self, new_cluster_size): + self.cluster_size.data.mul_(self.decay).add_(new_cluster_size, alpha=1 - self.decay) + + def embed_avg_ema_update(self, new_embed_avg): + self.embed_avg.data.mul_(self.decay).add_(new_embed_avg, alpha=1 - self.decay) + + def weight_update(self, num_tokens): + n = self.cluster_size.sum() + smoothed_cluster_size = ( + (self.cluster_size + self.eps) / (n + num_tokens * self.eps) * n + ) + #normalize embedding average with smoothed cluster size + embed_normalized = self.embed_avg / smoothed_cluster_size.unsqueeze(1) + self.weight.data.copy_(embed_normalized) + + +class EMAVectorQuantizer(nn.Module): + def __init__(self, n_embed, embedding_dim, beta, decay=0.99, eps=1e-5, + remap=None, unknown_index="random"): + super().__init__() + self.codebook_dim = codebook_dim + self.num_tokens = num_tokens + self.beta = beta + self.embedding = EmbeddingEMA(self.num_tokens, self.codebook_dim, decay, eps) + + self.remap = remap + if self.remap is not None: + self.register_buffer("used", torch.tensor(np.load(self.remap))) + self.re_embed = self.used.shape[0] + self.unknown_index = unknown_index # "random" or "extra" or integer + if self.unknown_index == "extra": + self.unknown_index = self.re_embed + self.re_embed = self.re_embed+1 + print(f"Remapping {self.n_embed} indices to {self.re_embed} indices. " + f"Using {self.unknown_index} for unknown indices.") + else: + self.re_embed = n_embed + + def remap_to_used(self, inds): + ishape = inds.shape + assert len(ishape)>1 + inds = inds.reshape(ishape[0],-1) + used = self.used.to(inds) + match = (inds[:,:,None]==used[None,None,...]).long() + new = match.argmax(-1) + unknown = match.sum(2)<1 + if self.unknown_index == "random": + new[unknown]=torch.randint(0,self.re_embed,size=new[unknown].shape).to(device=new.device) + else: + new[unknown] = self.unknown_index + return new.reshape(ishape) + + def unmap_to_all(self, inds): + ishape = inds.shape + assert len(ishape)>1 + inds = inds.reshape(ishape[0],-1) + used = self.used.to(inds) + if self.re_embed > self.used.shape[0]: # extra token + inds[inds>=self.used.shape[0]] = 0 # simply set to zero + back=torch.gather(used[None,:][inds.shape[0]*[0],:], 1, inds) + return back.reshape(ishape) + + def forward(self, z): + # reshape z -> (batch, height, width, channel) and flatten + #z, 'b c h w -> b h w c' + z = rearrange(z, 'b c h w -> b h w c') + z_flattened = z.reshape(-1, self.codebook_dim) + + # distances from z to embeddings e_j (z - e)^2 = z^2 + e^2 - 2 e * z + d = z_flattened.pow(2).sum(dim=1, keepdim=True) + \ + self.embedding.weight.pow(2).sum(dim=1) - 2 * \ + torch.einsum('bd,nd->bn', z_flattened, self.embedding.weight) # 'n d -> d n' + + + encoding_indices = torch.argmin(d, dim=1) + + z_q = self.embedding(encoding_indices).view(z.shape) + encodings = F.one_hot(encoding_indices, self.num_tokens).type(z.dtype) + avg_probs = torch.mean(encodings, dim=0) + perplexity = torch.exp(-torch.sum(avg_probs * torch.log(avg_probs + 1e-10))) + + if self.training and self.embedding.update: + #EMA cluster size + encodings_sum = encodings.sum(0) + self.embedding.cluster_size_ema_update(encodings_sum) + #EMA embedding average + embed_sum = encodings.transpose(0,1) @ z_flattened + self.embedding.embed_avg_ema_update(embed_sum) + #normalize embed_avg and update weight + self.embedding.weight_update(self.num_tokens) + + # compute loss for embedding + loss = self.beta * F.mse_loss(z_q.detach(), z) + + # preserve gradients + z_q = z + (z_q - z).detach() + + # reshape back to match original input shape + #z_q, 'b h w c -> b c h w' + z_q = rearrange(z_q, 'b h w c -> b c h w') + return z_q, loss, (perplexity, encodings, encoding_indices) diff --git a/taming/util.py b/taming/util.py new file mode 100644 index 0000000..06053e5 --- /dev/null +++ b/taming/util.py @@ -0,0 +1,157 @@ +import os, hashlib +import requests +from tqdm import tqdm + +URL_MAP = { + "vgg_lpips": "https://heibox.uni-heidelberg.de/f/607503859c864bc1b30b/?dl=1" +} + +CKPT_MAP = { + "vgg_lpips": "vgg.pth" +} + +MD5_MAP = { + "vgg_lpips": "d507d7349b931f0638a25a48a722f98a" +} + + +def download(url, local_path, chunk_size=1024): + os.makedirs(os.path.split(local_path)[0], exist_ok=True) + with requests.get(url, stream=True) as r: + total_size = int(r.headers.get("content-length", 0)) + with tqdm(total=total_size, unit="B", unit_scale=True) as pbar: + with open(local_path, "wb") as f: + for data in r.iter_content(chunk_size=chunk_size): + if data: + f.write(data) + pbar.update(chunk_size) + + +def md5_hash(path): + with open(path, "rb") as f: + content = f.read() + return hashlib.md5(content).hexdigest() + + +def get_ckpt_path(name, root, check=False): + assert name in URL_MAP + path = os.path.join(root, CKPT_MAP[name]) + if not os.path.exists(path) or (check and not md5_hash(path) == MD5_MAP[name]): + print("Downloading {} model from {} to {}".format(name, URL_MAP[name], path)) + download(URL_MAP[name], path) + md5 = md5_hash(path) + assert md5 == MD5_MAP[name], md5 + return path + + +class KeyNotFoundError(Exception): + def __init__(self, cause, keys=None, visited=None): + self.cause = cause + self.keys = keys + self.visited = visited + messages = list() + if keys is not None: + messages.append("Key not found: {}".format(keys)) + if visited is not None: + messages.append("Visited: {}".format(visited)) + messages.append("Cause:\n{}".format(cause)) + message = "\n".join(messages) + super().__init__(message) + + +def retrieve( + list_or_dict, key, splitval="/", default=None, expand=True, pass_success=False +): + """Given a nested list or dict return the desired value at key expanding + callable nodes if necessary and :attr:`expand` is ``True``. The expansion + is done in-place. + + Parameters + ---------- + list_or_dict : list or dict + Possibly nested list or dictionary. + key : str + key/to/value, path like string describing all keys necessary to + consider to get to the desired value. List indices can also be + passed here. + splitval : str + String that defines the delimiter between keys of the + different depth levels in `key`. + default : obj + Value returned if :attr:`key` is not found. + expand : bool + Whether to expand callable nodes on the path or not. + + Returns + ------- + The desired value or if :attr:`default` is not ``None`` and the + :attr:`key` is not found returns ``default``. + + Raises + ------ + Exception if ``key`` not in ``list_or_dict`` and :attr:`default` is + ``None``. + """ + + keys = key.split(splitval) + + success = True + try: + visited = [] + parent = None + last_key = None + for key in keys: + if callable(list_or_dict): + if not expand: + raise KeyNotFoundError( + ValueError( + "Trying to get past callable node with expand=False." + ), + keys=keys, + visited=visited, + ) + list_or_dict = list_or_dict() + parent[last_key] = list_or_dict + + last_key = key + parent = list_or_dict + + try: + if isinstance(list_or_dict, dict): + list_or_dict = list_or_dict[key] + else: + list_or_dict = list_or_dict[int(key)] + except (KeyError, IndexError, ValueError) as e: + raise KeyNotFoundError(e, keys=keys, visited=visited) + + visited += [key] + # final expansion of retrieved value + if expand and callable(list_or_dict): + list_or_dict = list_or_dict() + parent[last_key] = list_or_dict + except KeyNotFoundError as e: + if default is None: + raise e + else: + list_or_dict = default + success = False + + if not pass_success: + return list_or_dict + else: + return list_or_dict, success + + +if __name__ == "__main__": + config = {"keya": "a", + "keyb": "b", + "keyc": + {"cc1": 1, + "cc2": 2, + } + } + from omegaconf import OmegaConf + config = OmegaConf.create(config) + print(config) + retrieve(config, "keya") + diff --git a/train.py b/train.py deleted file mode 100644 index f2327b8..0000000 --- a/train.py +++ /dev/null @@ -1,32 +0,0 @@ -from argparse import ArgumentParser - -import pytorch_lightning as pl -from omegaconf import OmegaConf -import torch - -from utils.common import instantiate_from_config, load_state_dict - - -def main() -> None: - parser = ArgumentParser() - parser.add_argument("--config", type=str, default='configs/train_ccsr_stage2.yaml') - args = parser.parse_args() - - config = OmegaConf.load(args.config) - pl.seed_everything(config.lightning.seed, workers=True) - - data_module = instantiate_from_config(config.data) - model = instantiate_from_config(OmegaConf.load(config.model.config)) - # TODO: resume states saved in checkpoint. - if config.model.get("resume"): - load_state_dict(model, torch.load(config.model.resume, map_location="cpu"), strict=True) - - callbacks = [] - for callback_config in config.lightning.callbacks: - callbacks.append(instantiate_from_config(callback_config)) - trainer = pl.Trainer(callbacks=callbacks, **config.lightning.trainer) - trainer.fit(model, datamodule=data_module) - - -if __name__ == "__main__": - main() \ No newline at end of file diff --git a/utils/__pycache__/common.cpython-310.pyc b/utils/__pycache__/common.cpython-310.pyc index fc78650..83e33cf 100644 Binary files a/utils/__pycache__/common.cpython-310.pyc and b/utils/__pycache__/common.cpython-310.pyc differ diff --git a/utils/__pycache__/common.cpython-37.pyc b/utils/__pycache__/common.cpython-37.pyc deleted file mode 100644 index 0961383..0000000 Binary files a/utils/__pycache__/common.cpython-37.pyc and /dev/null differ diff --git a/utils/__pycache__/degradation.cpython-310.pyc b/utils/__pycache__/degradation.cpython-310.pyc deleted file mode 100644 index 17e8adf..0000000 Binary files a/utils/__pycache__/degradation.cpython-310.pyc and /dev/null differ diff --git a/utils/__pycache__/degradation.cpython-37.pyc b/utils/__pycache__/degradation.cpython-37.pyc deleted file mode 100644 index d9b09ef..0000000 Binary files a/utils/__pycache__/degradation.cpython-37.pyc and /dev/null differ diff --git a/utils/__pycache__/file.cpython-310.pyc b/utils/__pycache__/file.cpython-310.pyc deleted file mode 100644 index efa2be6..0000000 Binary files a/utils/__pycache__/file.cpython-310.pyc and /dev/null differ diff --git a/utils/__pycache__/file.cpython-37.pyc b/utils/__pycache__/file.cpython-37.pyc deleted file mode 100644 index 4d38007..0000000 Binary files a/utils/__pycache__/file.cpython-37.pyc and /dev/null differ diff --git a/utils/__pycache__/metrics.cpython-310.pyc b/utils/__pycache__/metrics.cpython-310.pyc deleted file mode 100644 index 55f8b41..0000000 Binary files a/utils/__pycache__/metrics.cpython-310.pyc and /dev/null differ diff --git a/utils/degradation.py b/utils/degradation.py deleted file mode 100644 index aa75976..0000000 --- a/utils/degradation.py +++ /dev/null @@ -1,765 +0,0 @@ -# https://github.com/XPixelGroup/BasicSR/blob/master/basicsr/data/degradations.py -import cv2 -import math -import numpy as np -import random -import torch -from scipy import special -from scipy.stats import multivariate_normal -from torchvision.transforms.functional_tensor import rgb_to_grayscale - -# -------------------------------------------------------------------- # -# --------------------------- blur kernels --------------------------- # -# -------------------------------------------------------------------- # - - -# --------------------------- util functions --------------------------- # -def sigma_matrix2(sig_x, sig_y, theta): - """Calculate the rotated sigma matrix (two dimensional matrix). - - Args: - sig_x (float): - sig_y (float): - theta (float): Radian measurement. - - Returns: - ndarray: Rotated sigma matrix. - """ - d_matrix = np.array([[sig_x**2, 0], [0, sig_y**2]]) - u_matrix = np.array([[np.cos(theta), -np.sin(theta)], [np.sin(theta), np.cos(theta)]]) - return np.dot(u_matrix, np.dot(d_matrix, u_matrix.T)) - - -def mesh_grid(kernel_size): - """Generate the mesh grid, centering at zero. - - Args: - kernel_size (int): - - Returns: - xy (ndarray): with the shape (kernel_size, kernel_size, 2) - xx (ndarray): with the shape (kernel_size, kernel_size) - yy (ndarray): with the shape (kernel_size, kernel_size) - """ - ax = np.arange(-kernel_size // 2 + 1., kernel_size // 2 + 1.) - xx, yy = np.meshgrid(ax, ax) - xy = np.hstack((xx.reshape((kernel_size * kernel_size, 1)), yy.reshape(kernel_size * kernel_size, - 1))).reshape(kernel_size, kernel_size, 2) - return xy, xx, yy - - -def pdf2(sigma_matrix, grid): - """Calculate PDF of the bivariate Gaussian distribution. - - Args: - sigma_matrix (ndarray): with the shape (2, 2) - grid (ndarray): generated by :func:`mesh_grid`, - with the shape (K, K, 2), K is the kernel size. - - Returns: - kernel (ndarrray): un-normalized kernel. - """ - inverse_sigma = np.linalg.inv(sigma_matrix) - kernel = np.exp(-0.5 * np.sum(np.dot(grid, inverse_sigma) * grid, 2)) - return kernel - - -def cdf2(d_matrix, grid): - """Calculate the CDF of the standard bivariate Gaussian distribution. - Used in skewed Gaussian distribution. - - Args: - d_matrix (ndarrasy): skew matrix. - grid (ndarray): generated by :func:`mesh_grid`, - with the shape (K, K, 2), K is the kernel size. - - Returns: - cdf (ndarray): skewed cdf. - """ - rv = multivariate_normal([0, 0], [[1, 0], [0, 1]]) - grid = np.dot(grid, d_matrix) - cdf = rv.cdf(grid) - return cdf - - -def bivariate_Gaussian(kernel_size, sig_x, sig_y, theta, grid=None, isotropic=True): - """Generate a bivariate isotropic or anisotropic Gaussian kernel. - - In the isotropic mode, only `sig_x` is used. `sig_y` and `theta` is ignored. - - Args: - kernel_size (int): - sig_x (float): - sig_y (float): - theta (float): Radian measurement. - grid (ndarray, optional): generated by :func:`mesh_grid`, - with the shape (K, K, 2), K is the kernel size. Default: None - isotropic (bool): - - Returns: - kernel (ndarray): normalized kernel. - """ - if grid is None: - grid, _, _ = mesh_grid(kernel_size) - if isotropic: - sigma_matrix = np.array([[sig_x**2, 0], [0, sig_x**2]]) - else: - sigma_matrix = sigma_matrix2(sig_x, sig_y, theta) - kernel = pdf2(sigma_matrix, grid) - kernel = kernel / np.sum(kernel) - return kernel - - -def bivariate_generalized_Gaussian(kernel_size, sig_x, sig_y, theta, beta, grid=None, isotropic=True): - """Generate a bivariate generalized Gaussian kernel. - - ``Paper: Parameter Estimation For Multivariate Generalized Gaussian Distributions`` - - In the isotropic mode, only `sig_x` is used. `sig_y` and `theta` is ignored. - - Args: - kernel_size (int): - sig_x (float): - sig_y (float): - theta (float): Radian measurement. - beta (float): shape parameter, beta = 1 is the normal distribution. - grid (ndarray, optional): generated by :func:`mesh_grid`, - with the shape (K, K, 2), K is the kernel size. Default: None - - Returns: - kernel (ndarray): normalized kernel. - """ - if grid is None: - grid, _, _ = mesh_grid(kernel_size) - if isotropic: - sigma_matrix = np.array([[sig_x**2, 0], [0, sig_x**2]]) - else: - sigma_matrix = sigma_matrix2(sig_x, sig_y, theta) - inverse_sigma = np.linalg.inv(sigma_matrix) - kernel = np.exp(-0.5 * np.power(np.sum(np.dot(grid, inverse_sigma) * grid, 2), beta)) - kernel = kernel / np.sum(kernel) - return kernel - - -def bivariate_plateau(kernel_size, sig_x, sig_y, theta, beta, grid=None, isotropic=True): - """Generate a plateau-like anisotropic kernel. - - 1 / (1+x^(beta)) - - Reference: https://stats.stackexchange.com/questions/203629/is-there-a-plateau-shaped-distribution - - In the isotropic mode, only `sig_x` is used. `sig_y` and `theta` is ignored. - - Args: - kernel_size (int): - sig_x (float): - sig_y (float): - theta (float): Radian measurement. - beta (float): shape parameter, beta = 1 is the normal distribution. - grid (ndarray, optional): generated by :func:`mesh_grid`, - with the shape (K, K, 2), K is the kernel size. Default: None - - Returns: - kernel (ndarray): normalized kernel. - """ - if grid is None: - grid, _, _ = mesh_grid(kernel_size) - if isotropic: - sigma_matrix = np.array([[sig_x**2, 0], [0, sig_x**2]]) - else: - sigma_matrix = sigma_matrix2(sig_x, sig_y, theta) - inverse_sigma = np.linalg.inv(sigma_matrix) - kernel = np.reciprocal(np.power(np.sum(np.dot(grid, inverse_sigma) * grid, 2), beta) + 1) - kernel = kernel / np.sum(kernel) - return kernel - - -def random_bivariate_Gaussian(kernel_size, - sigma_x_range, - sigma_y_range, - rotation_range, - noise_range=None, - isotropic=True): - """Randomly generate bivariate isotropic or anisotropic Gaussian kernels. - - In the isotropic mode, only `sigma_x_range` is used. `sigma_y_range` and `rotation_range` is ignored. - - Args: - kernel_size (int): - sigma_x_range (tuple): [0.6, 5] - sigma_y_range (tuple): [0.6, 5] - rotation range (tuple): [-math.pi, math.pi] - noise_range(tuple, optional): multiplicative kernel noise, - [0.75, 1.25]. Default: None - - Returns: - kernel (ndarray): - """ - assert kernel_size % 2 == 1, 'Kernel size must be an odd number.' - assert sigma_x_range[0] < sigma_x_range[1], 'Wrong sigma_x_range.' - sigma_x = np.random.uniform(sigma_x_range[0], sigma_x_range[1]) - if isotropic is False: - assert sigma_y_range[0] < sigma_y_range[1], 'Wrong sigma_y_range.' - assert rotation_range[0] < rotation_range[1], 'Wrong rotation_range.' - sigma_y = np.random.uniform(sigma_y_range[0], sigma_y_range[1]) - rotation = np.random.uniform(rotation_range[0], rotation_range[1]) - else: - sigma_y = sigma_x - rotation = 0 - - kernel = bivariate_Gaussian(kernel_size, sigma_x, sigma_y, rotation, isotropic=isotropic) - - # add multiplicative noise - if noise_range is not None: - assert noise_range[0] < noise_range[1], 'Wrong noise range.' - noise = np.random.uniform(noise_range[0], noise_range[1], size=kernel.shape) - kernel = kernel * noise - kernel = kernel / np.sum(kernel) - return kernel - - -def random_bivariate_generalized_Gaussian(kernel_size, - sigma_x_range, - sigma_y_range, - rotation_range, - beta_range, - noise_range=None, - isotropic=True): - """Randomly generate bivariate generalized Gaussian kernels. - - In the isotropic mode, only `sigma_x_range` is used. `sigma_y_range` and `rotation_range` is ignored. - - Args: - kernel_size (int): - sigma_x_range (tuple): [0.6, 5] - sigma_y_range (tuple): [0.6, 5] - rotation range (tuple): [-math.pi, math.pi] - beta_range (tuple): [0.5, 8] - noise_range(tuple, optional): multiplicative kernel noise, - [0.75, 1.25]. Default: None - - Returns: - kernel (ndarray): - """ - assert kernel_size % 2 == 1, 'Kernel size must be an odd number.' - assert sigma_x_range[0] < sigma_x_range[1], 'Wrong sigma_x_range.' - sigma_x = np.random.uniform(sigma_x_range[0], sigma_x_range[1]) - if isotropic is False: - assert sigma_y_range[0] < sigma_y_range[1], 'Wrong sigma_y_range.' - assert rotation_range[0] < rotation_range[1], 'Wrong rotation_range.' - sigma_y = np.random.uniform(sigma_y_range[0], sigma_y_range[1]) - rotation = np.random.uniform(rotation_range[0], rotation_range[1]) - else: - sigma_y = sigma_x - rotation = 0 - - # assume beta_range[0] < 1 < beta_range[1] - if np.random.uniform() < 0.5: - beta = np.random.uniform(beta_range[0], 1) - else: - beta = np.random.uniform(1, beta_range[1]) - - kernel = bivariate_generalized_Gaussian(kernel_size, sigma_x, sigma_y, rotation, beta, isotropic=isotropic) - - # add multiplicative noise - if noise_range is not None: - assert noise_range[0] < noise_range[1], 'Wrong noise range.' - noise = np.random.uniform(noise_range[0], noise_range[1], size=kernel.shape) - kernel = kernel * noise - kernel = kernel / np.sum(kernel) - return kernel - - -def random_bivariate_plateau(kernel_size, - sigma_x_range, - sigma_y_range, - rotation_range, - beta_range, - noise_range=None, - isotropic=True): - """Randomly generate bivariate plateau kernels. - - In the isotropic mode, only `sigma_x_range` is used. `sigma_y_range` and `rotation_range` is ignored. - - Args: - kernel_size (int): - sigma_x_range (tuple): [0.6, 5] - sigma_y_range (tuple): [0.6, 5] - rotation range (tuple): [-math.pi/2, math.pi/2] - beta_range (tuple): [1, 4] - noise_range(tuple, optional): multiplicative kernel noise, - [0.75, 1.25]. Default: None - - Returns: - kernel (ndarray): - """ - assert kernel_size % 2 == 1, 'Kernel size must be an odd number.' - assert sigma_x_range[0] < sigma_x_range[1], 'Wrong sigma_x_range.' - sigma_x = np.random.uniform(sigma_x_range[0], sigma_x_range[1]) - if isotropic is False: - assert sigma_y_range[0] < sigma_y_range[1], 'Wrong sigma_y_range.' - assert rotation_range[0] < rotation_range[1], 'Wrong rotation_range.' - sigma_y = np.random.uniform(sigma_y_range[0], sigma_y_range[1]) - rotation = np.random.uniform(rotation_range[0], rotation_range[1]) - else: - sigma_y = sigma_x - rotation = 0 - - # TODO: this may be not proper - if np.random.uniform() < 0.5: - beta = np.random.uniform(beta_range[0], 1) - else: - beta = np.random.uniform(1, beta_range[1]) - - kernel = bivariate_plateau(kernel_size, sigma_x, sigma_y, rotation, beta, isotropic=isotropic) - # add multiplicative noise - if noise_range is not None: - assert noise_range[0] < noise_range[1], 'Wrong noise range.' - noise = np.random.uniform(noise_range[0], noise_range[1], size=kernel.shape) - kernel = kernel * noise - kernel = kernel / np.sum(kernel) - - return kernel - - -def random_mixed_kernels(kernel_list, - kernel_prob, - kernel_size=21, - sigma_x_range=(0.6, 5), - sigma_y_range=(0.6, 5), - rotation_range=(-math.pi, math.pi), - betag_range=(0.5, 8), - betap_range=(0.5, 8), - noise_range=None): - """Randomly generate mixed kernels. - - Args: - kernel_list (tuple): a list name of kernel types, - support ['iso', 'aniso', 'skew', 'generalized', 'plateau_iso', - 'plateau_aniso'] - kernel_prob (tuple): corresponding kernel probability for each - kernel type - kernel_size (int): - sigma_x_range (tuple): [0.6, 5] - sigma_y_range (tuple): [0.6, 5] - rotation range (tuple): [-math.pi, math.pi] - beta_range (tuple): [0.5, 8] - noise_range(tuple, optional): multiplicative kernel noise, - [0.75, 1.25]. Default: None - - Returns: - kernel (ndarray): - """ - kernel_type = random.choices(kernel_list, kernel_prob)[0] - if kernel_type == 'iso': - kernel = random_bivariate_Gaussian( - kernel_size, sigma_x_range, sigma_y_range, rotation_range, noise_range=noise_range, isotropic=True) - elif kernel_type == 'aniso': - kernel = random_bivariate_Gaussian( - kernel_size, sigma_x_range, sigma_y_range, rotation_range, noise_range=noise_range, isotropic=False) - elif kernel_type == 'generalized_iso': - kernel = random_bivariate_generalized_Gaussian( - kernel_size, - sigma_x_range, - sigma_y_range, - rotation_range, - betag_range, - noise_range=noise_range, - isotropic=True) - elif kernel_type == 'generalized_aniso': - kernel = random_bivariate_generalized_Gaussian( - kernel_size, - sigma_x_range, - sigma_y_range, - rotation_range, - betag_range, - noise_range=noise_range, - isotropic=False) - elif kernel_type == 'plateau_iso': - kernel = random_bivariate_plateau( - kernel_size, sigma_x_range, sigma_y_range, rotation_range, betap_range, noise_range=None, isotropic=True) - elif kernel_type == 'plateau_aniso': - kernel = random_bivariate_plateau( - kernel_size, sigma_x_range, sigma_y_range, rotation_range, betap_range, noise_range=None, isotropic=False) - return kernel - - -np.seterr(divide='ignore', invalid='ignore') - - -def circular_lowpass_kernel(cutoff, kernel_size, pad_to=0): - """2D sinc filter - - Reference: https://dsp.stackexchange.com/questions/58301/2-d-circularly-symmetric-low-pass-filter - - Args: - cutoff (float): cutoff frequency in radians (pi is max) - kernel_size (int): horizontal and vertical size, must be odd. - pad_to (int): pad kernel size to desired size, must be odd or zero. - """ - assert kernel_size % 2 == 1, 'Kernel size must be an odd number.' - kernel = np.fromfunction( - lambda x, y: cutoff * special.j1(cutoff * np.sqrt( - (x - (kernel_size - 1) / 2)**2 + (y - (kernel_size - 1) / 2)**2)) / (2 * np.pi * np.sqrt( - (x - (kernel_size - 1) / 2)**2 + (y - (kernel_size - 1) / 2)**2)), [kernel_size, kernel_size]) - kernel[(kernel_size - 1) // 2, (kernel_size - 1) // 2] = cutoff**2 / (4 * np.pi) - kernel = kernel / np.sum(kernel) - if pad_to > kernel_size: - pad_size = (pad_to - kernel_size) // 2 - kernel = np.pad(kernel, ((pad_size, pad_size), (pad_size, pad_size))) - return kernel - - -# ------------------------------------------------------------- # -# --------------------------- noise --------------------------- # -# ------------------------------------------------------------- # - -# ----------------------- Gaussian Noise ----------------------- # - - -def generate_gaussian_noise(img, sigma=10, gray_noise=False): - """Generate Gaussian noise. - - Args: - img (Numpy array): Input image, shape (h, w, c), range [0, 1], float32. - sigma (float): Noise scale (measured in range 255). Default: 10. - - Returns: - (Numpy array): Returned noisy image, shape (h, w, c), range[0, 1], - float32. - """ - if gray_noise: - noise = np.float32(np.random.randn(*(img.shape[0:2]))) * sigma / 255. - noise = np.expand_dims(noise, axis=2).repeat(3, axis=2) - else: - noise = np.float32(np.random.randn(*(img.shape))) * sigma / 255. - return noise - - -def add_gaussian_noise(img, sigma=10, clip=True, rounds=False, gray_noise=False): - """Add Gaussian noise. - - Args: - img (Numpy array): Input image, shape (h, w, c), range [0, 1], float32. - sigma (float): Noise scale (measured in range 255). Default: 10. - - Returns: - (Numpy array): Returned noisy image, shape (h, w, c), range[0, 1], - float32. - """ - noise = generate_gaussian_noise(img, sigma, gray_noise) - out = img + noise - if clip and rounds: - out = np.clip((out * 255.0).round(), 0, 255) / 255. - elif clip: - out = np.clip(out, 0, 1) - elif rounds: - out = (out * 255.0).round() / 255. - return out - - -def generate_gaussian_noise_pt(img, sigma=10, gray_noise=0): - """Add Gaussian noise (PyTorch version). - - Args: - img (Tensor): Shape (b, c, h, w), range[0, 1], float32. - scale (float | Tensor): Noise scale. Default: 1.0. - - Returns: - (Tensor): Returned noisy image, shape (b, c, h, w), range[0, 1], - float32. - """ - b, _, h, w = img.size() - if not isinstance(sigma, (float, int)): - sigma = sigma.view(img.size(0), 1, 1, 1) - if isinstance(gray_noise, (float, int)): - cal_gray_noise = gray_noise > 0 - else: - gray_noise = gray_noise.view(b, 1, 1, 1) - cal_gray_noise = torch.sum(gray_noise) > 0 - - if cal_gray_noise: - noise_gray = torch.randn(*img.size()[2:4], dtype=img.dtype, device=img.device) * sigma / 255. - noise_gray = noise_gray.view(b, 1, h, w) - - # always calculate color noise - noise = torch.randn(*img.size(), dtype=img.dtype, device=img.device) * sigma / 255. - - if cal_gray_noise: - noise = noise * (1 - gray_noise) + noise_gray * gray_noise - return noise - - -def add_gaussian_noise_pt(img, sigma=10, gray_noise=0, clip=True, rounds=False): - """Add Gaussian noise (PyTorch version). - - Args: - img (Tensor): Shape (b, c, h, w), range[0, 1], float32. - scale (float | Tensor): Noise scale. Default: 1.0. - - Returns: - (Tensor): Returned noisy image, shape (b, c, h, w), range[0, 1], - float32. - """ - noise = generate_gaussian_noise_pt(img, sigma, gray_noise) - out = img + noise - if clip and rounds: - out = torch.clamp((out * 255.0).round(), 0, 255) / 255. - elif clip: - out = torch.clamp(out, 0, 1) - elif rounds: - out = (out * 255.0).round() / 255. - return out - - -# ----------------------- Random Gaussian Noise ----------------------- # -def random_generate_gaussian_noise(img, sigma_range=(0, 10), gray_prob=0): - sigma = np.random.uniform(sigma_range[0], sigma_range[1]) - if np.random.uniform() < gray_prob: - gray_noise = True - else: - gray_noise = False - return generate_gaussian_noise(img, sigma, gray_noise) - - -def random_add_gaussian_noise(img, sigma_range=(0, 1.0), gray_prob=0, clip=True, rounds=False): - noise = random_generate_gaussian_noise(img, sigma_range, gray_prob) - out = img + noise - if clip and rounds: - out = np.clip((out * 255.0).round(), 0, 255) / 255. - elif clip: - out = np.clip(out, 0, 1) - elif rounds: - out = (out * 255.0).round() / 255. - return out - - -def random_generate_gaussian_noise_pt(img, sigma_range=(0, 10), gray_prob=0): - sigma = torch.rand( - img.size(0), dtype=img.dtype, device=img.device) * (sigma_range[1] - sigma_range[0]) + sigma_range[0] - gray_noise = torch.rand(img.size(0), dtype=img.dtype, device=img.device) - gray_noise = (gray_noise < gray_prob).float() - return generate_gaussian_noise_pt(img, sigma, gray_noise) - - -def random_add_gaussian_noise_pt(img, sigma_range=(0, 1.0), gray_prob=0, clip=True, rounds=False): - noise = random_generate_gaussian_noise_pt(img, sigma_range, gray_prob) - out = img + noise - if clip and rounds: - out = torch.clamp((out * 255.0).round(), 0, 255) / 255. - elif clip: - out = torch.clamp(out, 0, 1) - elif rounds: - out = (out * 255.0).round() / 255. - return out - - -# ----------------------- Poisson (Shot) Noise ----------------------- # - - -def generate_poisson_noise(img, scale=1.0, gray_noise=False): - """Generate poisson noise. - - Reference: https://github.com/scikit-image/scikit-image/blob/main/skimage/util/noise.py#L37-L219 - - Args: - img (Numpy array): Input image, shape (h, w, c), range [0, 1], float32. - scale (float): Noise scale. Default: 1.0. - gray_noise (bool): Whether generate gray noise. Default: False. - - Returns: - (Numpy array): Returned noisy image, shape (h, w, c), range[0, 1], - float32. - """ - if gray_noise: - img = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY) - # round and clip image for counting vals correctly - img = np.clip((img * 255.0).round(), 0, 255) / 255. - vals = len(np.unique(img)) - vals = 2**np.ceil(np.log2(vals)) - out = np.float32(np.random.poisson(img * vals) / float(vals)) - noise = out - img - if gray_noise: - noise = np.repeat(noise[:, :, np.newaxis], 3, axis=2) - return noise * scale - - -def add_poisson_noise(img, scale=1.0, clip=True, rounds=False, gray_noise=False): - """Add poisson noise. - - Args: - img (Numpy array): Input image, shape (h, w, c), range [0, 1], float32. - scale (float): Noise scale. Default: 1.0. - gray_noise (bool): Whether generate gray noise. Default: False. - - Returns: - (Numpy array): Returned noisy image, shape (h, w, c), range[0, 1], - float32. - """ - noise = generate_poisson_noise(img, scale, gray_noise) - out = img + noise - if clip and rounds: - out = np.clip((out * 255.0).round(), 0, 255) / 255. - elif clip: - out = np.clip(out, 0, 1) - elif rounds: - out = (out * 255.0).round() / 255. - return out - - -def generate_poisson_noise_pt(img, scale=1.0, gray_noise=0): - """Generate a batch of poisson noise (PyTorch version) - - Args: - img (Tensor): Input image, shape (b, c, h, w), range [0, 1], float32. - scale (float | Tensor): Noise scale. Number or Tensor with shape (b). - Default: 1.0. - gray_noise (float | Tensor): 0-1 number or Tensor with shape (b). - 0 for False, 1 for True. Default: 0. - - Returns: - (Tensor): Returned noisy image, shape (b, c, h, w), range[0, 1], - float32. - """ - b, _, h, w = img.size() - if isinstance(gray_noise, (float, int)): - cal_gray_noise = gray_noise > 0 - else: - gray_noise = gray_noise.view(b, 1, 1, 1) - cal_gray_noise = torch.sum(gray_noise) > 0 - if cal_gray_noise: - img_gray = rgb_to_grayscale(img, num_output_channels=1) - # round and clip image for counting vals correctly - img_gray = torch.clamp((img_gray * 255.0).round(), 0, 255) / 255. - # use for-loop to get the unique values for each sample - vals_list = [len(torch.unique(img_gray[i, :, :, :])) for i in range(b)] - vals_list = [2**np.ceil(np.log2(vals)) for vals in vals_list] - vals = img_gray.new_tensor(vals_list).view(b, 1, 1, 1) - out = torch.poisson(img_gray * vals) / vals - noise_gray = out - img_gray - noise_gray = noise_gray.expand(b, 3, h, w) - - # always calculate color noise - # round and clip image for counting vals correctly - img = torch.clamp((img * 255.0).round(), 0, 255) / 255. - # use for-loop to get the unique values for each sample - vals_list = [len(torch.unique(img[i, :, :, :])) for i in range(b)] - vals_list = [2**np.ceil(np.log2(vals)) for vals in vals_list] - vals = img.new_tensor(vals_list).view(b, 1, 1, 1) - out = torch.poisson(img * vals) / vals - noise = out - img - if cal_gray_noise: - noise = noise * (1 - gray_noise) + noise_gray * gray_noise - if not isinstance(scale, (float, int)): - scale = scale.view(b, 1, 1, 1) - return noise * scale - - -def add_poisson_noise_pt(img, scale=1.0, clip=True, rounds=False, gray_noise=0): - """Add poisson noise to a batch of images (PyTorch version). - - Args: - img (Tensor): Input image, shape (b, c, h, w), range [0, 1], float32. - scale (float | Tensor): Noise scale. Number or Tensor with shape (b). - Default: 1.0. - gray_noise (float | Tensor): 0-1 number or Tensor with shape (b). - 0 for False, 1 for True. Default: 0. - - Returns: - (Tensor): Returned noisy image, shape (b, c, h, w), range[0, 1], - float32. - """ - noise = generate_poisson_noise_pt(img, scale, gray_noise) - out = img + noise - if clip and rounds: - out = torch.clamp((out * 255.0).round(), 0, 255) / 255. - elif clip: - out = torch.clamp(out, 0, 1) - elif rounds: - out = (out * 255.0).round() / 255. - return out - - -# ----------------------- Random Poisson (Shot) Noise ----------------------- # - - -def random_generate_poisson_noise(img, scale_range=(0, 1.0), gray_prob=0): - scale = np.random.uniform(scale_range[0], scale_range[1]) - if np.random.uniform() < gray_prob: - gray_noise = True - else: - gray_noise = False - return generate_poisson_noise(img, scale, gray_noise) - - -def random_add_poisson_noise(img, scale_range=(0, 1.0), gray_prob=0, clip=True, rounds=False): - noise = random_generate_poisson_noise(img, scale_range, gray_prob) - out = img + noise - if clip and rounds: - out = np.clip((out * 255.0).round(), 0, 255) / 255. - elif clip: - out = np.clip(out, 0, 1) - elif rounds: - out = (out * 255.0).round() / 255. - return out - - -def random_generate_poisson_noise_pt(img, scale_range=(0, 1.0), gray_prob=0): - scale = torch.rand( - img.size(0), dtype=img.dtype, device=img.device) * (scale_range[1] - scale_range[0]) + scale_range[0] - gray_noise = torch.rand(img.size(0), dtype=img.dtype, device=img.device) - gray_noise = (gray_noise < gray_prob).float() - return generate_poisson_noise_pt(img, scale, gray_noise) - - -def random_add_poisson_noise_pt(img, scale_range=(0, 1.0), gray_prob=0, clip=True, rounds=False): - noise = random_generate_poisson_noise_pt(img, scale_range, gray_prob) - out = img + noise - if clip and rounds: - out = torch.clamp((out * 255.0).round(), 0, 255) / 255. - elif clip: - out = torch.clamp(out, 0, 1) - elif rounds: - out = (out * 255.0).round() / 255. - return out - - -# ------------------------------------------------------------------------ # -# --------------------------- JPEG compression --------------------------- # -# ------------------------------------------------------------------------ # - - -def add_jpg_compression(img, quality=90): - """Add JPG compression artifacts. - - Args: - img (Numpy array): Input image, shape (h, w, c), range [0, 1], float32. - quality (float): JPG compression quality. 0 for lowest quality, 100 for - best quality. Default: 90. - - Returns: - (Numpy array): Returned image after JPG, shape (h, w, c), range[0, 1], - float32. - """ - img = np.clip(img, 0, 1) - encode_param = [int(cv2.IMWRITE_JPEG_QUALITY), quality] - _, encimg = cv2.imencode('.jpg', img * 255., encode_param) - img = np.float32(cv2.imdecode(encimg, 1)) / 255. - return img - - -def random_add_jpg_compression(img, quality_range=(90, 100)): - """Randomly add JPG compression artifacts. - - Args: - img (Numpy array): Input image, shape (h, w, c), range [0, 1], float32. - quality_range (tuple[float] | list[float]): JPG compression quality - range. 0 for lowest quality, 100 for best quality. - Default: (90, 100). - - Returns: - (Numpy array): Returned image after JPG, shape (h, w, c), range[0, 1], - float32. - """ - quality = np.random.uniform(quality_range[0], quality_range[1]) - return add_jpg_compression(img, int(quality)) diff --git a/utils/face_restoration_helper.py b/utils/face_restoration_helper.py deleted file mode 100644 index 6f09a3c..0000000 --- a/utils/face_restoration_helper.py +++ /dev/null @@ -1,517 +0,0 @@ -import cv2 -import numpy as np -import os -import torch -from torchvision.transforms.functional import normalize - -from facexlib.detection import init_detection_model -from facexlib.parsing import init_parsing_model -from facexlib.utils.misc import img2tensor, imwrite - -from .file import load_file_from_url - -def get_largest_face(det_faces, h, w): - - def get_location(val, length): - if val < 0: - return 0 - elif val > length: - return length - else: - return val - - face_areas = [] - for det_face in det_faces: - left = get_location(det_face[0], w) - right = get_location(det_face[2], w) - top = get_location(det_face[1], h) - bottom = get_location(det_face[3], h) - face_area = (right - left) * (bottom - top) - face_areas.append(face_area) - largest_idx = face_areas.index(max(face_areas)) - return det_faces[largest_idx], largest_idx - - -def get_center_face(det_faces, h=0, w=0, center=None): - if center is not None: - center = np.array(center) - else: - center = np.array([w / 2, h / 2]) - center_dist = [] - for det_face in det_faces: - face_center = np.array([(det_face[0] + det_face[2]) / 2, (det_face[1] + det_face[3]) / 2]) - dist = np.linalg.norm(face_center - center) - center_dist.append(dist) - center_idx = center_dist.index(min(center_dist)) - return det_faces[center_idx], center_idx - - -class FaceRestoreHelper(object): - """Helper for the face restoration pipeline (base class).""" - - def __init__(self, - upscale_factor, - face_size=512, - crop_ratio=(1, 1), - det_model='retinaface_resnet50', - save_ext='png', - template_3points=False, - pad_blur=False, - use_parse=False, - device=None): - self.template_3points = template_3points # improve robustness - self.upscale_factor = int(upscale_factor) - # the cropped face ratio based on the square face - self.crop_ratio = crop_ratio # (h, w) - assert (self.crop_ratio[0] >= 1 and self.crop_ratio[1] >= 1), 'crop ration only supports >=1' - self.face_size = (int(face_size * self.crop_ratio[1]), int(face_size * self.crop_ratio[0])) - self.det_model = det_model - - if self.det_model == 'dlib': - # standard 5 landmarks for FFHQ faces with 1024 x 1024 - self.face_template = np.array([[686.77227723, 488.62376238], [586.77227723, 493.59405941], - [337.91089109, 488.38613861], [437.95049505, 493.51485149], - [513.58415842, 678.5049505]]) - self.face_template = self.face_template / (1024 // face_size) - elif self.template_3points: - self.face_template = np.array([[192, 240], [319, 240], [257, 371]]) - else: - # standard 5 landmarks for FFHQ faces with 512 x 512 - # facexlib - self.face_template = np.array([[192.98138, 239.94708], [318.90277, 240.1936], [256.63416, 314.01935], - [201.26117, 371.41043], [313.08905, 371.15118]]) - - # dlib: left_eye: 36:41 right_eye: 42:47 nose: 30,32,33,34 left mouth corner: 48 right mouth corner: 54 - # self.face_template = np.array([[193.65928, 242.98541], [318.32558, 243.06108], [255.67984, 328.82894], - # [198.22603, 372.82502], [313.91018, 372.75659]]) - - self.face_template = self.face_template * (face_size / 512.0) - if self.crop_ratio[0] > 1: - self.face_template[:, 1] += face_size * (self.crop_ratio[0] - 1) / 2 - if self.crop_ratio[1] > 1: - self.face_template[:, 0] += face_size * (self.crop_ratio[1] - 1) / 2 - self.save_ext = save_ext - self.pad_blur = pad_blur - if self.pad_blur is True: - self.template_3points = False - - self.all_landmarks_5 = [] - self.det_faces = [] - self.affine_matrices = [] - self.inverse_affine_matrices = [] - self.cropped_faces = [] - self.restored_faces = [] - self.pad_input_imgs = [] - - if device is None: - self.device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') - # self.device = get_device() - else: - self.device = device - - # init face detection model - self.face_detector = init_detection_model(det_model, half=False, device=self.device) - - # init face parsing model - self.use_parse = use_parse - self.face_parse = init_parsing_model(model_name='parsenet', device=self.device) - - def set_upscale_factor(self, upscale_factor): - self.upscale_factor = upscale_factor - - def read_image(self, img): - """img can be image path or cv2 loaded image.""" - # self.input_img is Numpy array, (h, w, c), BGR, uint8, [0, 255] - if isinstance(img, str): - img = cv2.imread(img) - - if np.max(img) > 256: # 16-bit image - img = img / 65535 * 255 - if len(img.shape) == 2: # gray image - img = cv2.cvtColor(img, cv2.COLOR_GRAY2BGR) - elif img.shape[2] == 4: # BGRA image with alpha channel - img = img[:, :, 0:3] - - self.input_img = img - # self.is_gray = is_gray(img, threshold=10) - # if self.is_gray: - # print('Grayscale input: True') - - if min(self.input_img.shape[:2])<512: - f = 512.0/min(self.input_img.shape[:2]) - self.input_img = cv2.resize(self.input_img, (0,0), fx=f, fy=f, interpolation=cv2.INTER_LINEAR) - - def init_dlib(self, detection_path, landmark5_path): - """Initialize the dlib detectors and predictors.""" - try: - import dlib - except ImportError: - print('Please install dlib by running:' 'conda install -c conda-forge dlib') - detection_path = load_file_from_url(url=detection_path, model_dir='weights/dlib', progress=True, file_name=None) - landmark5_path = load_file_from_url(url=landmark5_path, model_dir='weights/dlib', progress=True, file_name=None) - face_detector = dlib.cnn_face_detection_model_v1(detection_path) - shape_predictor_5 = dlib.shape_predictor(landmark5_path) - return face_detector, shape_predictor_5 - - def get_face_landmarks_5_dlib(self, - only_keep_largest=False, - scale=1): - det_faces = self.face_detector(self.input_img, scale) - - if len(det_faces) == 0: - print('No face detected. Try to increase upsample_num_times.') - return 0 - else: - if only_keep_largest: - print('Detect several faces and only keep the largest.') - face_areas = [] - for i in range(len(det_faces)): - face_area = (det_faces[i].rect.right() - det_faces[i].rect.left()) * ( - det_faces[i].rect.bottom() - det_faces[i].rect.top()) - face_areas.append(face_area) - largest_idx = face_areas.index(max(face_areas)) - self.det_faces = [det_faces[largest_idx]] - else: - self.det_faces = det_faces - - if len(self.det_faces) == 0: - return 0 - - for face in self.det_faces: - shape = self.shape_predictor_5(self.input_img, face.rect) - landmark = np.array([[part.x, part.y] for part in shape.parts()]) - self.all_landmarks_5.append(landmark) - - return len(self.all_landmarks_5) - - - def get_face_landmarks_5(self, - only_keep_largest=False, - only_center_face=False, - resize=None, - blur_ratio=0.01, - eye_dist_threshold=None): - if self.det_model == 'dlib': - return self.get_face_landmarks_5_dlib(only_keep_largest) - - if resize is None: - scale = 1 - input_img = self.input_img - else: - h, w = self.input_img.shape[0:2] - scale = resize / min(h, w) - scale = max(1, scale) # always scale up - h, w = int(h * scale), int(w * scale) - interp = cv2.INTER_AREA if scale < 1 else cv2.INTER_LINEAR - input_img = cv2.resize(self.input_img, (w, h), interpolation=interp) - - with torch.no_grad(): - bboxes = self.face_detector.detect_faces(input_img) - - if bboxes is None or bboxes.shape[0] == 0: - return 0 - else: - bboxes = bboxes / scale - - for bbox in bboxes: - # remove faces with too small eye distance: side faces or too small faces - eye_dist = np.linalg.norm([bbox[6] - bbox[8], bbox[7] - bbox[9]]) - if eye_dist_threshold is not None and (eye_dist < eye_dist_threshold): - continue - - if self.template_3points: - landmark = np.array([[bbox[i], bbox[i + 1]] for i in range(5, 11, 2)]) - else: - landmark = np.array([[bbox[i], bbox[i + 1]] for i in range(5, 15, 2)]) - self.all_landmarks_5.append(landmark) - self.det_faces.append(bbox[0:5]) - - if len(self.det_faces) == 0: - return 0 - if only_keep_largest: - h, w, _ = self.input_img.shape - self.det_faces, largest_idx = get_largest_face(self.det_faces, h, w) - self.all_landmarks_5 = [self.all_landmarks_5[largest_idx]] - elif only_center_face: - h, w, _ = self.input_img.shape - self.det_faces, center_idx = get_center_face(self.det_faces, h, w) - self.all_landmarks_5 = [self.all_landmarks_5[center_idx]] - - # pad blurry images - if self.pad_blur: - self.pad_input_imgs = [] - for landmarks in self.all_landmarks_5: - # get landmarks - eye_left = landmarks[0, :] - eye_right = landmarks[1, :] - eye_avg = (eye_left + eye_right) * 0.5 - mouth_avg = (landmarks[3, :] + landmarks[4, :]) * 0.5 - eye_to_eye = eye_right - eye_left - eye_to_mouth = mouth_avg - eye_avg - - # Get the oriented crop rectangle - # x: half width of the oriented crop rectangle - x = eye_to_eye - np.flipud(eye_to_mouth) * [-1, 1] - # - np.flipud(eye_to_mouth) * [-1, 1]: rotate 90 clockwise - # norm with the hypotenuse: get the direction - x /= np.hypot(*x) # get the hypotenuse of a right triangle - rect_scale = 1.5 - x *= max(np.hypot(*eye_to_eye) * 2.0 * rect_scale, np.hypot(*eye_to_mouth) * 1.8 * rect_scale) - # y: half height of the oriented crop rectangle - y = np.flipud(x) * [-1, 1] - - # c: center - c = eye_avg + eye_to_mouth * 0.1 - # quad: (left_top, left_bottom, right_bottom, right_top) - quad = np.stack([c - x - y, c - x + y, c + x + y, c + x - y]) - # qsize: side length of the square - qsize = np.hypot(*x) * 2 - border = max(int(np.rint(qsize * 0.1)), 3) - - # get pad - # pad: (width_left, height_top, width_right, height_bottom) - pad = (int(np.floor(min(quad[:, 0]))), int(np.floor(min(quad[:, 1]))), int(np.ceil(max(quad[:, 0]))), - int(np.ceil(max(quad[:, 1])))) - pad = [ - max(-pad[0] + border, 1), - max(-pad[1] + border, 1), - max(pad[2] - self.input_img.shape[0] + border, 1), - max(pad[3] - self.input_img.shape[1] + border, 1) - ] - - if max(pad) > 1: - # pad image - pad_img = np.pad(self.input_img, ((pad[1], pad[3]), (pad[0], pad[2]), (0, 0)), 'reflect') - # modify landmark coords - landmarks[:, 0] += pad[0] - landmarks[:, 1] += pad[1] - # blur pad images - h, w, _ = pad_img.shape - y, x, _ = np.ogrid[:h, :w, :1] - mask = np.maximum(1.0 - np.minimum(np.float32(x) / pad[0], - np.float32(w - 1 - x) / pad[2]), - 1.0 - np.minimum(np.float32(y) / pad[1], - np.float32(h - 1 - y) / pad[3])) - blur = int(qsize * blur_ratio) - if blur % 2 == 0: - blur += 1 - blur_img = cv2.boxFilter(pad_img, 0, ksize=(blur, blur)) - # blur_img = cv2.GaussianBlur(pad_img, (blur, blur), 0) - - pad_img = pad_img.astype('float32') - pad_img += (blur_img - pad_img) * np.clip(mask * 3.0 + 1.0, 0.0, 1.0) - pad_img += (np.median(pad_img, axis=(0, 1)) - pad_img) * np.clip(mask, 0.0, 1.0) - pad_img = np.clip(pad_img, 0, 255) # float32, [0, 255] - self.pad_input_imgs.append(pad_img) - else: - self.pad_input_imgs.append(np.copy(self.input_img)) - - return len(self.all_landmarks_5) - - def align_warp_face(self, save_cropped_path=None, border_mode='constant'): - """Align and warp faces with face template. - """ - if self.pad_blur: - assert len(self.pad_input_imgs) == len( - self.all_landmarks_5), f'Mismatched samples: {len(self.pad_input_imgs)} and {len(self.all_landmarks_5)}' - for idx, landmark in enumerate(self.all_landmarks_5): - # use 5 landmarks to get affine matrix - # use cv2.LMEDS method for the equivalence to skimage transform - # ref: https://blog.csdn.net/yichxi/article/details/115827338 - affine_matrix = cv2.estimateAffinePartial2D(landmark, self.face_template, method=cv2.LMEDS)[0] - self.affine_matrices.append(affine_matrix) - # warp and crop faces - if border_mode == 'constant': - border_mode = cv2.BORDER_CONSTANT - elif border_mode == 'reflect101': - border_mode = cv2.BORDER_REFLECT101 - elif border_mode == 'reflect': - border_mode = cv2.BORDER_REFLECT - if self.pad_blur: - input_img = self.pad_input_imgs[idx] - else: - input_img = self.input_img - cropped_face = cv2.warpAffine( - input_img, affine_matrix, self.face_size, borderMode=border_mode, borderValue=(135, 133, 132)) # gray - self.cropped_faces.append(cropped_face) - # save the cropped face - if save_cropped_path is not None: - path = os.path.splitext(save_cropped_path)[0] - save_path = f'{path}_{idx:02d}.{self.save_ext}' - imwrite(cropped_face, save_path) - - def get_inverse_affine(self, save_inverse_affine_path=None): - """Get inverse affine matrix.""" - for idx, affine_matrix in enumerate(self.affine_matrices): - inverse_affine = cv2.invertAffineTransform(affine_matrix) - inverse_affine *= self.upscale_factor - self.inverse_affine_matrices.append(inverse_affine) - # save inverse affine matrices - if save_inverse_affine_path is not None: - path, _ = os.path.splitext(save_inverse_affine_path) - save_path = f'{path}_{idx:02d}.pth' - torch.save(inverse_affine, save_path) - - - def add_restored_face(self, restored_face, input_face=None): - # if self.is_gray: - # restored_face = bgr2gray(restored_face) # convert img into grayscale - # if input_face is not None: - # restored_face = adain_npy(restored_face, input_face) # transfer the color - self.restored_faces.append(restored_face) - - - def paste_faces_to_input_image(self, save_path=None, upsample_img=None, draw_box=False, face_upsampler=None): - h, w, _ = self.input_img.shape - h_up, w_up = int(h * self.upscale_factor), int(w * self.upscale_factor) - - if upsample_img is None: - # simply resize the background - # upsample_img = cv2.resize(self.input_img, (w_up, h_up), interpolation=cv2.INTER_LANCZOS4) - upsample_img = cv2.resize(self.input_img, (w_up, h_up), interpolation=cv2.INTER_LINEAR) - else: - upsample_img = cv2.resize(upsample_img, (w_up, h_up), interpolation=cv2.INTER_LANCZOS4) - - assert len(self.restored_faces) == len( - self.inverse_affine_matrices), ('length of restored_faces and affine_matrices are different.') - - inv_mask_borders = [] - for restored_face, inverse_affine in zip(self.restored_faces, self.inverse_affine_matrices): - if face_upsampler is not None: - restored_face = face_upsampler.enhance(restored_face, outscale=self.upscale_factor)[0] - inverse_affine /= self.upscale_factor - inverse_affine[:, 2] *= self.upscale_factor - face_size = (self.face_size[0]*self.upscale_factor, self.face_size[1]*self.upscale_factor) - else: - # Add an offset to inverse affine matrix, for more precise back alignment - if self.upscale_factor > 1: - extra_offset = 0.5 * self.upscale_factor - else: - extra_offset = 0 - inverse_affine[:, 2] += extra_offset - face_size = self.face_size - inv_restored = cv2.warpAffine(restored_face, inverse_affine, (w_up, h_up)) - - # if draw_box or not self.use_parse: # use square parse maps - # mask = np.ones(face_size, dtype=np.float32) - # inv_mask = cv2.warpAffine(mask, inverse_affine, (w_up, h_up)) - # # remove the black borders - # inv_mask_erosion = cv2.erode( - # inv_mask, np.ones((int(2 * self.upscale_factor), int(2 * self.upscale_factor)), np.uint8)) - # pasted_face = inv_mask_erosion[:, :, None] * inv_restored - # total_face_area = np.sum(inv_mask_erosion) # // 3 - # # add border - # if draw_box: - # h, w = face_size - # mask_border = np.ones((h, w, 3), dtype=np.float32) - # border = int(1400/np.sqrt(total_face_area)) - # mask_border[border:h-border, border:w-border,:] = 0 - # inv_mask_border = cv2.warpAffine(mask_border, inverse_affine, (w_up, h_up)) - # inv_mask_borders.append(inv_mask_border) - # if not self.use_parse: - # # compute the fusion edge based on the area of face - # w_edge = int(total_face_area**0.5) // 20 - # erosion_radius = w_edge * 2 - # inv_mask_center = cv2.erode(inv_mask_erosion, np.ones((erosion_radius, erosion_radius), np.uint8)) - # blur_size = w_edge * 2 - # inv_soft_mask = cv2.GaussianBlur(inv_mask_center, (blur_size + 1, blur_size + 1), 0) - # if len(upsample_img.shape) == 2: # upsample_img is gray image - # upsample_img = upsample_img[:, :, None] - # inv_soft_mask = inv_soft_mask[:, :, None] - - # always use square mask - mask = np.ones(face_size, dtype=np.float32) - inv_mask = cv2.warpAffine(mask, inverse_affine, (w_up, h_up)) - # remove the black borders - inv_mask_erosion = cv2.erode( - inv_mask, np.ones((int(2 * self.upscale_factor), int(2 * self.upscale_factor)), np.uint8)) - pasted_face = inv_mask_erosion[:, :, None] * inv_restored - total_face_area = np.sum(inv_mask_erosion) # // 3 - # add border - if draw_box: - h, w = face_size - mask_border = np.ones((h, w, 3), dtype=np.float32) - border = int(1400/np.sqrt(total_face_area)) - mask_border[border:h-border, border:w-border,:] = 0 - inv_mask_border = cv2.warpAffine(mask_border, inverse_affine, (w_up, h_up)) - inv_mask_borders.append(inv_mask_border) - # compute the fusion edge based on the area of face - w_edge = int(total_face_area**0.5) // 20 - erosion_radius = w_edge * 2 - inv_mask_center = cv2.erode(inv_mask_erosion, np.ones((erosion_radius, erosion_radius), np.uint8)) - blur_size = w_edge * 2 - inv_soft_mask = cv2.GaussianBlur(inv_mask_center, (blur_size + 1, blur_size + 1), 0) - if len(upsample_img.shape) == 2: # upsample_img is gray image - upsample_img = upsample_img[:, :, None] - inv_soft_mask = inv_soft_mask[:, :, None] - - # parse mask - if self.use_parse: - # inference - face_input = cv2.resize(restored_face, (512, 512), interpolation=cv2.INTER_LINEAR) - face_input = img2tensor(face_input.astype('float32') / 255., bgr2rgb=True, float32=True) - normalize(face_input, (0.5, 0.5, 0.5), (0.5, 0.5, 0.5), inplace=True) - face_input = torch.unsqueeze(face_input, 0).to(self.device) - with torch.no_grad(): - out = self.face_parse(face_input)[0] - out = out.argmax(dim=1).squeeze().cpu().numpy() - - parse_mask = np.zeros(out.shape) - MASK_COLORMAP = [0, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255, 255, 0, 255, 0, 0, 0] - for idx, color in enumerate(MASK_COLORMAP): - parse_mask[out == idx] = color - # blur the mask - parse_mask = cv2.GaussianBlur(parse_mask, (101, 101), 11) - parse_mask = cv2.GaussianBlur(parse_mask, (101, 101), 11) - # remove the black borders - thres = 10 - parse_mask[:thres, :] = 0 - parse_mask[-thres:, :] = 0 - parse_mask[:, :thres] = 0 - parse_mask[:, -thres:] = 0 - parse_mask = parse_mask / 255. - - parse_mask = cv2.resize(parse_mask, face_size) - parse_mask = cv2.warpAffine(parse_mask, inverse_affine, (w_up, h_up), flags=3) - inv_soft_parse_mask = parse_mask[:, :, None] - # pasted_face = inv_restored - fuse_mask = (inv_soft_parse_mask 256: # 16-bit image - upsample_img = upsample_img.astype(np.uint16) - else: - upsample_img = upsample_img.astype(np.uint8) - - # draw bounding box - if draw_box: - # upsample_input_img = cv2.resize(input_img, (w_up, h_up)) - img_color = np.ones([*upsample_img.shape], dtype=np.float32) - img_color[:,:,0] = 0 - img_color[:,:,1] = 255 - img_color[:,:,2] = 0 - for inv_mask_border in inv_mask_borders: - upsample_img = inv_mask_border * img_color + (1 - inv_mask_border) * upsample_img - # upsample_input_img = inv_mask_border * img_color + (1 - inv_mask_border) * upsample_input_img - - if save_path is not None: - path = os.path.splitext(save_path)[0] - save_path = f'{path}.{self.save_ext}' - imwrite(upsample_img, save_path) - return upsample_img - - def clean_all(self): - self.all_landmarks_5 = [] - self.restored_faces = [] - self.affine_matrices = [] - self.cropped_faces = [] - self.inverse_affine_matrices = [] - self.det_faces = [] - self.pad_input_imgs = [] \ No newline at end of file diff --git a/utils/file.py b/utils/file.py deleted file mode 100644 index 9cf5d11..0000000 --- a/utils/file.py +++ /dev/null @@ -1,103 +0,0 @@ -import os -from typing import List, Tuple - -from urllib.parse import urlparse -from torch.hub import download_url_to_file, get_dir - - -def load_file_list(file_list_path: str) -> List[str]: - files = [] - # each line in file list contains a path of an image - with open(file_list_path, "r") as fin: - for line in fin: - path = line.strip() - if path: - files.append(path) - return files - - -# def list_image_files( -# img_dir: str, -# exts: Tuple[str]=(".jpg", ".png", ".jpeg"), -# follow_links: bool=False, -# log_progress: bool=False, -# log_every_n_files: int=10000, -# max_size: int=-1 -# ) -> List[str]: -# files = [] -# for dir_path, _, file_names in os.walk(img_dir, followlinks=follow_links): -# early_stop = False -# for file_name in file_names: -# if os.path.splitext(file_name)[1].lower() in exts: -# if max_size >= 0 and len(files) >= max_size: -# early_stop = True -# break -# files.append(os.path.join(dir_path, file_name)) -# if log_progress and len(files) % log_every_n_files == 0: -# print(f"find {len(files)} images in {img_dir}") -# if early_stop: -# break -# return files - -def list_image_files( - img_dir: List[str], - exts: Tuple[str]=(".jpg", ".png", ".jpeg"), - follow_links: bool=False, - log_progress: bool=False, - log_every_n_files: int=10000, - max_size: int=-1 -) -> List[str]: - files = [] - for img_d in img_dir: - for dir_path, _, file_names in os.walk(img_d, followlinks=follow_links): - early_stop = False - for file_name in file_names: - if os.path.splitext(file_name)[1].lower() in exts: - if max_size >= 0 and len(files) >= max_size: - early_stop = True - break - files.append(os.path.join(dir_path, file_name)) - if log_progress and len(files) % log_every_n_files == 0: - print(f"find {len(files)} images in {img_dir}") - if early_stop: - break - return files - - -def get_file_name_parts(file_path: str) -> Tuple[str, str, str]: - parent_path, file_name = os.path.split(file_path) - stem, ext = os.path.splitext(file_name) - return parent_path, stem, ext - - -# https://github.com/XPixelGroup/BasicSR/blob/master/basicsr/utils/download_util.py/ -def load_file_from_url(url, model_dir=None, progress=True, file_name=None): - """Load file form http url, will download models if necessary. - - Ref:https://github.com/1adrianb/face-alignment/blob/master/face_alignment/utils.py - - Args: - url (str): URL to be downloaded. - model_dir (str): The path to save the downloaded model. Should be a full path. If None, use pytorch hub_dir. - Default: None. - progress (bool): Whether to show the download progress. Default: True. - file_name (str): The downloaded file name. If None, use the file name in the url. Default: None. - - Returns: - str: The path to the downloaded file. - """ - if model_dir is None: # use the pytorch hub_dir - hub_dir = get_dir() - model_dir = os.path.join(hub_dir, 'checkpoints') - - os.makedirs(model_dir, exist_ok=True) - - parts = urlparse(url) - filename = os.path.basename(parts.path) - if file_name is not None: - filename = file_name - cached_file = os.path.abspath(os.path.join(model_dir, filename)) - if not os.path.exists(cached_file): - print(f'Downloading: "{url}" to {cached_file}\n') - download_url_to_file(url, cached_file, hash_prefix=None, progress=progress) - return cached_file diff --git a/utils/image/__pycache__/__init__.cpython-310.pyc b/utils/image/__pycache__/__init__.cpython-310.pyc index ba710e5..4d3d508 100644 Binary files a/utils/image/__pycache__/__init__.cpython-310.pyc and b/utils/image/__pycache__/__init__.cpython-310.pyc differ diff --git a/utils/image/__pycache__/align_color.cpython-310.pyc b/utils/image/__pycache__/align_color.cpython-310.pyc index 975b7fd..8579942 100644 Binary files a/utils/image/__pycache__/align_color.cpython-310.pyc and b/utils/image/__pycache__/align_color.cpython-310.pyc differ diff --git a/utils/image/__pycache__/common.cpython-310.pyc b/utils/image/__pycache__/common.cpython-310.pyc index c4ce7bb..696c1eb 100644 Binary files a/utils/image/__pycache__/common.cpython-310.pyc and b/utils/image/__pycache__/common.cpython-310.pyc differ diff --git a/utils/image/__pycache__/diffjpeg.cpython-310.pyc b/utils/image/__pycache__/diffjpeg.cpython-310.pyc index 2a2a50d..553cd18 100644 Binary files a/utils/image/__pycache__/diffjpeg.cpython-310.pyc and b/utils/image/__pycache__/diffjpeg.cpython-310.pyc differ diff --git a/utils/image/__pycache__/usm_sharp.cpython-310.pyc b/utils/image/__pycache__/usm_sharp.cpython-310.pyc index 906e0bb..8291699 100644 Binary files a/utils/image/__pycache__/usm_sharp.cpython-310.pyc and b/utils/image/__pycache__/usm_sharp.cpython-310.pyc differ diff --git a/utils/metrics.py b/utils/metrics.py deleted file mode 100644 index e1c3e49..0000000 --- a/utils/metrics.py +++ /dev/null @@ -1,66 +0,0 @@ -import torch -import lpips - -from .image import rgb2ycbcr_pt -from .common import frozen_module - - -# https://github.com/XPixelGroup/BasicSR/blob/033cd6896d898fdd3dcda32e3102a792efa1b8f4/basicsr/metrics/psnr_ssim.py#L52 -def calculate_psnr_pt(img, img2, crop_border, test_y_channel=False): - """Calculate PSNR (Peak Signal-to-Noise Ratio) (PyTorch version). - - Reference: https://en.wikipedia.org/wiki/Peak_signal-to-noise_ratio - - Args: - img (Tensor): Images with range [0, 1], shape (n, 3/1, h, w). - img2 (Tensor): Images with range [0, 1], shape (n, 3/1, h, w). - crop_border (int): Cropped pixels in each edge of an image. These pixels are not involved in the calculation. - test_y_channel (bool): Test on Y channel of YCbCr. Default: False. - - Returns: - float: PSNR result. - """ - - assert img.shape == img2.shape, (f'Image shapes are different: {img.shape}, {img2.shape}.') - - if crop_border != 0: - img = img[:, :, crop_border:-crop_border, crop_border:-crop_border] - img2 = img2[:, :, crop_border:-crop_border, crop_border:-crop_border] - - if test_y_channel: - img = rgb2ycbcr_pt(img, y_only=True) - img2 = rgb2ycbcr_pt(img2, y_only=True) - - img = img.to(torch.float64) - img2 = img2.to(torch.float64) - - mse = torch.mean((img - img2)**2, dim=[1, 2, 3]) - return 10. * torch.log10(1. / (mse + 1e-8)) - - -class LPIPS: - - def __init__(self, net: str) -> None: - self.model = lpips.LPIPS(net=net) - frozen_module(self.model) - - @torch.no_grad() - def __call__(self, img1: torch.Tensor, img2: torch.Tensor, normalize: bool) -> torch.Tensor: - """ - Compute LPIPS. - - Args: - img1 (torch.Tensor): The first image (NCHW, RGB, [-1, 1]). Specify `normalize` if input - image is range in [0, 1]. - img2 (torch.Tensor): The second image (NCHW, RGB, [-1, 1]). Specify `normalize` if input - image is range in [0, 1]. - normalize (bool): If specified, the input images will be normalized from [0, 1] to [-1, 1]. - - Returns: - lpips_values (torch.Tensor): The lpips scores of this batch. - """ - return self.model(img1, img2, normalize=normalize) - - def to(self, device: str) -> "LPIPS": - self.model.to(device) - return self diff --git a/utils/realesrgan/realesrganer.py b/utils/realesrgan/realesrganer.py deleted file mode 100644 index ce10fa9..0000000 --- a/utils/realesrgan/realesrganer.py +++ /dev/null @@ -1,339 +0,0 @@ -import cv2 -import math -import numpy as np -import os -import queue -import threading -import torch -from torch.nn import functional as F - -from utils.file import load_file_from_url -from utils.realesrgan.rrdbnet import RRDBNet - -# ROOT_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) - -class RealESRGANer(): - """A helper class for upsampling images with RealESRGAN. - - Args: - scale (int): Upsampling scale factor used in the networks. It is usually 2 or 4. - model_path (str): The path to the pretrained model. It can be urls (will first download it automatically). - model (nn.Module): The defined network. Default: None. - tile (int): As too large images result in the out of GPU memory issue, so this tile option will first crop - input images into tiles, and then process each of them. Finally, they will be merged into one image. - 0 denotes for do not use tile. Default: 0. - tile_pad (int): The pad size for each tile, to remove border artifacts. Default: 10. - pre_pad (int): Pad the input images to avoid border artifacts. Default: 10. - half (float): Whether to use half precision during inference. Default: False. - """ - - def __init__(self, - scale, - model_path, - model=None, - tile=0, - tile_pad=10, - pre_pad=10, - half=False, - device=None): - self.scale = scale - self.tile_size = tile - self.tile_pad = tile_pad - self.pre_pad = pre_pad - self.mod_scale = None - self.half = half - - # initialize model - # if gpu_id: - # self.device = torch.device( - # f'cuda:{gpu_id}' if torch.cuda.is_available() else 'cpu') if device is None else device - # else: - # self.device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') if device is None else device - - self.device = device - - # if the model_path starts with https, it will first download models to the folder: realesrgan/weights - if model_path.startswith('https://'): - model_path = load_file_from_url( - url=model_path, model_dir=os.path.join('weights/realesrgan'), progress=True, file_name=None) - loadnet = torch.load(model_path, map_location=torch.device('cpu')) - # prefer to use params_ema - if 'params_ema' in loadnet: - keyname = 'params_ema' - else: - keyname = 'params' - model.load_state_dict(loadnet[keyname], strict=True) - model.eval() - self.model = model.to(self.device) - if self.half: - self.model = self.model.half() - - def pre_process(self, img): - """Pre-process, such as pre-pad and mod pad, so that the images can be divisible - """ - img = torch.from_numpy(np.transpose(img, (2, 0, 1))).float() - self.img = img.unsqueeze(0).to(self.device) - if self.half: - self.img = self.img.half() - - # pre_pad - if self.pre_pad != 0: - self.img = F.pad(self.img, (0, self.pre_pad, 0, self.pre_pad), 'reflect') - # mod pad for divisible borders - if self.scale == 2: - self.mod_scale = 2 - elif self.scale == 1: - self.mod_scale = 4 - if self.mod_scale is not None: - self.mod_pad_h, self.mod_pad_w = 0, 0 - _, _, h, w = self.img.size() - if (h % self.mod_scale != 0): - self.mod_pad_h = (self.mod_scale - h % self.mod_scale) - if (w % self.mod_scale != 0): - self.mod_pad_w = (self.mod_scale - w % self.mod_scale) - self.img = F.pad(self.img, (0, self.mod_pad_w, 0, self.mod_pad_h), 'reflect') - - def process(self): - # model inference - self.output = self.model(self.img) - - def tile_process(self): - """It will first crop input images to tiles, and then process each tile. - Finally, all the processed tiles are merged into one images. - - Modified from: https://github.com/ata4/esrgan-launcher - """ - batch, channel, height, width = self.img.shape - output_height = height * self.scale - output_width = width * self.scale - output_shape = (batch, channel, output_height, output_width) - - # start with black image - self.output = self.img.new_zeros(output_shape) - tiles_x = math.ceil(width / self.tile_size) - tiles_y = math.ceil(height / self.tile_size) - - # loop over all tiles - for y in range(tiles_y): - for x in range(tiles_x): - # extract tile from input image - ofs_x = x * self.tile_size - ofs_y = y * self.tile_size - # input tile area on total image - input_start_x = ofs_x - input_end_x = min(ofs_x + self.tile_size, width) - input_start_y = ofs_y - input_end_y = min(ofs_y + self.tile_size, height) - - # input tile area on total image with padding - input_start_x_pad = max(input_start_x - self.tile_pad, 0) - input_end_x_pad = min(input_end_x + self.tile_pad, width) - input_start_y_pad = max(input_start_y - self.tile_pad, 0) - input_end_y_pad = min(input_end_y + self.tile_pad, height) - - # input tile dimensions - input_tile_width = input_end_x - input_start_x - input_tile_height = input_end_y - input_start_y - tile_idx = y * tiles_x + x + 1 - input_tile = self.img[:, :, input_start_y_pad:input_end_y_pad, input_start_x_pad:input_end_x_pad] - - # upscale tile - try: - with torch.no_grad(): - output_tile = self.model(input_tile) - except RuntimeError as error: - print('Error', error) - # print(f'\tTile {tile_idx}/{tiles_x * tiles_y}') - - # output tile area on total image - output_start_x = input_start_x * self.scale - output_end_x = input_end_x * self.scale - output_start_y = input_start_y * self.scale - output_end_y = input_end_y * self.scale - - # output tile area without padding - output_start_x_tile = (input_start_x - input_start_x_pad) * self.scale - output_end_x_tile = output_start_x_tile + input_tile_width * self.scale - output_start_y_tile = (input_start_y - input_start_y_pad) * self.scale - output_end_y_tile = output_start_y_tile + input_tile_height * self.scale - - # put tile into output image - self.output[:, :, output_start_y:output_end_y, - output_start_x:output_end_x] = output_tile[:, :, output_start_y_tile:output_end_y_tile, - output_start_x_tile:output_end_x_tile] - - def post_process(self): - # remove extra pad - if self.mod_scale is not None: - _, _, h, w = self.output.size() - self.output = self.output[:, :, 0:h - self.mod_pad_h * self.scale, 0:w - self.mod_pad_w * self.scale] - # remove prepad - if self.pre_pad != 0: - _, _, h, w = self.output.size() - self.output = self.output[:, :, 0:h - self.pre_pad * self.scale, 0:w - self.pre_pad * self.scale] - return self.output - - @torch.no_grad() - def enhance(self, img, outscale=None, alpha_upsampler='realesrgan'): - h_input, w_input = img.shape[0:2] - # img: numpy - img = img.astype(np.float32) - if np.max(img) > 256: # 16-bit image - max_range = 65535 - print('\tInput is a 16-bit image') - else: - max_range = 255 - img = img / max_range - if len(img.shape) == 2: # gray image - img_mode = 'L' - img = cv2.cvtColor(img, cv2.COLOR_GRAY2RGB) - elif img.shape[2] == 4: # RGBA image with alpha channel - img_mode = 'RGBA' - alpha = img[:, :, 3] - img = img[:, :, 0:3] - img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) - if alpha_upsampler == 'realesrgan': - alpha = cv2.cvtColor(alpha, cv2.COLOR_GRAY2RGB) - else: - img_mode = 'RGB' - img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) - - # ------------------- process image (without the alpha channel) ------------------- # - try: - with torch.no_grad(): - self.pre_process(img) - if self.tile_size > 0: - self.tile_process() - else: - self.process() - output_img_t = self.post_process() - output_img = output_img_t.data.squeeze().float().cpu().clamp_(0, 1).numpy() - output_img = np.transpose(output_img[[2, 1, 0], :, :], (1, 2, 0)) - if img_mode == 'L': - output_img = cv2.cvtColor(output_img, cv2.COLOR_BGR2GRAY) - del output_img_t - torch.cuda.empty_cache() - except RuntimeError as error: - print(f"Failed inference for RealESRGAN: {error}") - - # ------------------- process the alpha channel if necessary ------------------- # - if img_mode == 'RGBA': - if alpha_upsampler == 'realesrgan': - self.pre_process(alpha) - if self.tile_size > 0: - self.tile_process() - else: - self.process() - output_alpha = self.post_process() - output_alpha = output_alpha.data.squeeze().float().cpu().clamp_(0, 1).numpy() - output_alpha = np.transpose(output_alpha[[2, 1, 0], :, :], (1, 2, 0)) - output_alpha = cv2.cvtColor(output_alpha, cv2.COLOR_BGR2GRAY) - else: # use the cv2 resize for alpha channel - h, w = alpha.shape[0:2] - output_alpha = cv2.resize(alpha, (w * self.scale, h * self.scale), interpolation=cv2.INTER_LINEAR) - - # merge the alpha channel - output_img = cv2.cvtColor(output_img, cv2.COLOR_BGR2BGRA) - output_img[:, :, 3] = output_alpha - - # ------------------------------ return ------------------------------ # - if max_range == 65535: # 16-bit image - output = (output_img * 65535.0).round().astype(np.uint16) - else: - output = (output_img * 255.0).round().astype(np.uint8) - - if outscale is not None and outscale != float(self.scale): - output = cv2.resize( - output, ( - int(w_input * outscale), - int(h_input * outscale), - ), interpolation=cv2.INTER_LANCZOS4) - - return output, img_mode - - -class PrefetchReader(threading.Thread): - """Prefetch images. - - Args: - img_list (list[str]): A image list of image paths to be read. - num_prefetch_queue (int): Number of prefetch queue. - """ - - def __init__(self, img_list, num_prefetch_queue): - super().__init__() - self.que = queue.Queue(num_prefetch_queue) - self.img_list = img_list - - def run(self): - for img_path in self.img_list: - img = cv2.imread(img_path, cv2.IMREAD_UNCHANGED) - self.que.put(img) - - self.que.put(None) - - def __next__(self): - next_item = self.que.get() - if next_item is None: - raise StopIteration - return next_item - - def __iter__(self): - return self - - -class IOConsumer(threading.Thread): - - def __init__(self, opt, que, qid): - super().__init__() - self._queue = que - self.qid = qid - self.opt = opt - - def run(self): - while True: - msg = self._queue.get() - if isinstance(msg, str) and msg == 'quit': - break - - output = msg['output'] - save_path = msg['save_path'] - cv2.imwrite(save_path, output) - print(f'IO worker {self.qid} is done.') - -def set_realesrgan(bg_tile, device, scale=2): - ''' - scale: options: 2, 4. Default: 2. RealESRGAN official models only support x2 and x4 upsampling. - ''' - assert isinstance(scale, int), 'Expected param scale to be an integer!' - - use_half = False - if 'cuda' in str(device): # set False in CPU/MPS mode - no_half_gpu_list = ['1650', '1660'] # set False for GPUs that don't support f16 - if not True in [gpu in torch.cuda.get_device_name(0) for gpu in no_half_gpu_list]: - use_half = True - - model_url = { - 2: 'https://github.com/xinntao/Real-ESRGAN/releases/download/v0.2.1/RealESRGAN_x2plus.pth', - 4: 'https://github.com/xinntao/Real-ESRGAN/releases/download/v0.1.0/RealESRGAN_x4plus.pth' - } - - model = RRDBNet( - num_in_ch=3, - num_out_ch=3, - num_feat=64, - num_block=23, - num_grow_ch=32, - scale=scale, - ) - upsampler = RealESRGANer( - scale=scale, - model_path=model_url[scale], - model=model, - tile=bg_tile, - tile_pad=40, - pre_pad=0, - device=device, - half=use_half - ) - return upsampler \ No newline at end of file diff --git a/utils/realesrgan/rrdbnet.py b/utils/realesrgan/rrdbnet.py deleted file mode 100644 index 11b23a3..0000000 --- a/utils/realesrgan/rrdbnet.py +++ /dev/null @@ -1,183 +0,0 @@ -import torch -from torch import nn as nn -from torch.nn import functional as F -from torch.nn import init as init -from torch.nn.modules.batchnorm import _BatchNorm - - -def default_init_weights(module_list, scale=1, bias_fill=0, **kwargs): - """Initialize network weights. - - Args: - module_list (list[nn.Module] | nn.Module): Modules to be initialized. - scale (float): Scale initialized weights, especially for residual - blocks. Default: 1. - bias_fill (float): The value to fill bias. Default: 0 - kwargs (dict): Other arguments for initialization function. - """ - if not isinstance(module_list, list): - module_list = [module_list] - for module in module_list: - for m in module.modules(): - if isinstance(m, nn.Conv2d): - init.kaiming_normal_(m.weight, **kwargs) - m.weight.data *= scale - if m.bias is not None: - m.bias.data.fill_(bias_fill) - elif isinstance(m, nn.Linear): - init.kaiming_normal_(m.weight, **kwargs) - m.weight.data *= scale - if m.bias is not None: - m.bias.data.fill_(bias_fill) - elif isinstance(m, _BatchNorm): - init.constant_(m.weight, 1) - if m.bias is not None: - m.bias.data.fill_(bias_fill) - - -def make_layer(basic_block, num_basic_block, **kwarg): - """Make layers by stacking the same blocks. - - Args: - basic_block (nn.module): nn.module class for basic block. - num_basic_block (int): number of blocks. - - Returns: - nn.Sequential: Stacked blocks in nn.Sequential. - """ - layers = [] - for _ in range(num_basic_block): - layers.append(basic_block(**kwarg)) - return nn.Sequential(*layers) - - -# TODO: may write a cpp file -def pixel_unshuffle(x, scale): - """ Pixel unshuffle. - - Args: - x (Tensor): Input feature with shape (b, c, hh, hw). - scale (int): Downsample ratio. - - Returns: - Tensor: the pixel unshuffled feature. - """ - b, c, hh, hw = x.size() - out_channel = c * (scale**2) - assert hh % scale == 0 and hw % scale == 0 - h = hh // scale - w = hw // scale - x_view = x.view(b, c, h, scale, w, scale) - return x_view.permute(0, 1, 3, 5, 2, 4).reshape(b, out_channel, h, w) - - -class ResidualDenseBlock(nn.Module): - """Residual Dense Block. - - Used in RRDB block in ESRGAN. - - Args: - num_feat (int): Channel number of intermediate features. - num_grow_ch (int): Channels for each growth. - """ - - def __init__(self, num_feat=64, num_grow_ch=32): - super(ResidualDenseBlock, self).__init__() - self.conv1 = nn.Conv2d(num_feat, num_grow_ch, 3, 1, 1) - self.conv2 = nn.Conv2d(num_feat + num_grow_ch, num_grow_ch, 3, 1, 1) - self.conv3 = nn.Conv2d(num_feat + 2 * num_grow_ch, num_grow_ch, 3, 1, 1) - self.conv4 = nn.Conv2d(num_feat + 3 * num_grow_ch, num_grow_ch, 3, 1, 1) - self.conv5 = nn.Conv2d(num_feat + 4 * num_grow_ch, num_feat, 3, 1, 1) - - self.lrelu = nn.LeakyReLU(negative_slope=0.2, inplace=True) - - # initialization - default_init_weights([self.conv1, self.conv2, self.conv3, self.conv4, self.conv5], 0.1) - - def forward(self, x): - x1 = self.lrelu(self.conv1(x)) - x2 = self.lrelu(self.conv2(torch.cat((x, x1), 1))) - x3 = self.lrelu(self.conv3(torch.cat((x, x1, x2), 1))) - x4 = self.lrelu(self.conv4(torch.cat((x, x1, x2, x3), 1))) - x5 = self.conv5(torch.cat((x, x1, x2, x3, x4), 1)) - # Empirically, we use 0.2 to scale the residual for better performance - return x5 * 0.2 + x - - -class RRDB(nn.Module): - """Residual in Residual Dense Block. - - Used in RRDB-Net in ESRGAN. - - Args: - num_feat (int): Channel number of intermediate features. - num_grow_ch (int): Channels for each growth. - """ - - def __init__(self, num_feat, num_grow_ch=32): - super(RRDB, self).__init__() - self.rdb1 = ResidualDenseBlock(num_feat, num_grow_ch) - self.rdb2 = ResidualDenseBlock(num_feat, num_grow_ch) - self.rdb3 = ResidualDenseBlock(num_feat, num_grow_ch) - - def forward(self, x): - out = self.rdb1(x) - out = self.rdb2(out) - out = self.rdb3(out) - # Empirically, we use 0.2 to scale the residual for better performance - return out * 0.2 + x - - -class RRDBNet(nn.Module): - """Networks consisting of Residual in Residual Dense Block, which is used - in ESRGAN. - - ESRGAN: Enhanced Super-Resolution Generative Adversarial Networks. - - We extend ESRGAN for scale x2 and scale x1. - Note: This is one option for scale 1, scale 2 in RRDBNet. - We first employ the pixel-unshuffle (an inverse operation of pixelshuffle to reduce the spatial size - and enlarge the channel size before feeding inputs into the main ESRGAN architecture. - - Args: - num_in_ch (int): Channel number of inputs. - num_out_ch (int): Channel number of outputs. - num_feat (int): Channel number of intermediate features. - Default: 64 - num_block (int): Block number in the trunk network. Defaults: 23 - num_grow_ch (int): Channels for each growth. Default: 32. - """ - - def __init__(self, num_in_ch, num_out_ch, scale=4, num_feat=64, num_block=23, num_grow_ch=32): - super(RRDBNet, self).__init__() - self.scale = scale - if scale == 2: - num_in_ch = num_in_ch * 4 - elif scale == 1: - num_in_ch = num_in_ch * 16 - self.conv_first = nn.Conv2d(num_in_ch, num_feat, 3, 1, 1) - self.body = make_layer(RRDB, num_block, num_feat=num_feat, num_grow_ch=num_grow_ch) - self.conv_body = nn.Conv2d(num_feat, num_feat, 3, 1, 1) - # upsample - self.conv_up1 = nn.Conv2d(num_feat, num_feat, 3, 1, 1) - self.conv_up2 = nn.Conv2d(num_feat, num_feat, 3, 1, 1) - self.conv_hr = nn.Conv2d(num_feat, num_feat, 3, 1, 1) - self.conv_last = nn.Conv2d(num_feat, num_out_ch, 3, 1, 1) - - self.lrelu = nn.LeakyReLU(negative_slope=0.2, inplace=True) - - def forward(self, x): - if self.scale == 2: - feat = pixel_unshuffle(x, scale=2) - elif self.scale == 1: - feat = pixel_unshuffle(x, scale=4) - else: - feat = x - feat = self.conv_first(feat) - body_feat = self.conv_body(self.body(feat)) - feat = feat + body_feat - # upsample - feat = self.lrelu(self.conv_up1(F.interpolate(feat, scale_factor=2, mode='nearest'))) - feat = self.lrelu(self.conv_up2(F.interpolate(feat, scale_factor=2, mode='nearest'))) - out = self.conv_last(self.lrelu(self.conv_hr(feat))) - return out \ No newline at end of file