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
-
-
-## 😍 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.
-
-
-
-### Comparisons on Bicubic SR
-
-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.
-
-
-## ⚙ 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
-
-
-
-
-
-
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