Initial commit

This commit is contained in:
Kijai
2024-01-12 17:01:00 +02:00
parent b5e7d03b45
commit 6597ff79d7
139 changed files with 6334 additions and 4398 deletions
-214
View File
@@ -1,214 +0,0 @@
<p align="center">
<img src="figs/logo.png" width="400">
</p>
## Improving the Stability of Diffusion Models for Content Consistent Super-Resolution
<a href='https://arxiv.org/pdf/2401.00877.pdf'><img src='https://img.shields.io/badge/Paper-Arxiv-red'></a> <a href='https://csslc.github.io/project-CCSR'><img src='https://img.shields.io/badge/Project page-Github-blue'></a> <a href='https://github.com/csslc/CCSR'><img src='https://img.shields.io/badge/Code-Github-green'></a>
[Lingchen Sun](https://scholar.google.com/citations?hl=zh-CN&tzom=-480&user=ZCDjTn8AAAAJ)<sup>1,2</sup>
| [Rongyuan Wu](https://scholar.google.com/citations?user=A-U8zE8AAAAJ&hl=zh-CN)<sup>1,2</sup> |
[Zhengqiang Zhang](https://scholar.google.com/citations?hl=zh-CN&user=UX26wSMAAAAJ&view_op=list_works&sortby=pubdate)<sup>1,2</sup> |
[Hongwei Yong](https://scholar.google.com.hk/citations?user=Xii74qQAAAAJ&hl=zh-CN)<sup>1</sup> |
[Lei Zhang](https://www4.comp.polyu.edu.hk/~cslzhang)<sup>1,2</sup>
<sup>1</sup>The Hong Kong Polytechnic University, <sup>2</sup>OPPO Research Institute
## ⏰ Update
- **2024.1.4**: Code and the model for real-world SR are released.
- **2024.1.3**: Paper is released.
- **2023.12.23**: Repo is released.
:star: If CCSR is helpful to your images or projects, please help star this repo. Thanks! :hugs:
## 🌟 Overview Framework
![ccsr](figs/framework.png)
## 😍 Visual Results
### Demo on Real-World SR
[<img src="figs/compare_1.png" height="223px"/>](https://imgsli.com/MjMxMzA0) [<img src="figs/compare_2.png" height="223px"/>](https://imgsli.com/MjMxMzEx) [<img src="figs/compare_4.png" height="223px"/>](https://imgsli.com/MjMxMzE1) [<img src="figs/compare_6.png" height="223px"/>](https://imgsli.com/MjMxMzI3)
[<img src="figs/compare_3.png" height="223px"/>](https://imgsli.com/MjMxMzEy) [<img src="figs/compare_5.png" height="223px"/>](https://imgsli.com/MjMxMzE5)
### Comparisons on Real-World SR
For the diffusion model-based method, two restored images that have the best and worst PSNR values over 10 runs are shown for a more comprehensive and fair comparison.
![ccsr](figs/realworld.png)
### Comparisons on Bicubic SR
![ccsr](figs/bicubic.png)
For more comparisons, please refer to our paper for details.
## 📝 Quantitative comparisons
We propose new stability metrics, namely global standard deviation (G-STD) and local standard deviation (L-STD), to respectively measure the image-level and pixel-level variations of the SR results of diffusion-based methods.
More details about G-STD and L-STD can be found in our paper.
![ccsr](figs/table.png)
## ⚙ Dependencies and Installation
```shell
## git clone this repository
git clone https://github.com/csslc/CCSR.git
cd CCSR
# create an environment with python >= 3.9
conda create -n ccsr python=3.9
conda activate ccsr
pip install -r requirements.txt
pip install -e git+https://github.com/CompVis/taming-transformers.git@master#egg=taming-transformers
```
## 🍭 Quick Inference
#### Step 1: Download the pretrained models
- Download the CCSR models from:
| Model Name | Description | GoogleDrive | OneDive |
|:---------------------|:---------------------------------------------|:--------------------------------------------------------------------------------------|:------------------------------------------------------------------------------|
| real-world_ccsr.ckpt | CCSR model for real-world image restoration. | [download](https://drive.google.com/drive/folders/1jM1mxDryPk9CTuFTvYcraP2XIVzbPiw_?usp=drive_link) | download |
| bicubic_ccsr.ckpt | CCSR model for bicubic image restoration. | download | download |
#### Step 2: Prepare testing data
You can put the testing images in the `preset/test_datasets`.
#### Step 3: Running testing command
```
python inference_ccsr.py \
--input preset/test_datasets \
--config configs/model/ccsr_stage2.yaml \
--ckpt weights/real-world_ccsr.ckpt \
--steps 45 \
--sr_scale 4 \
--t_max 0.6667 \
--t_min 0.3333 \
--color_fix_type adain \
--output experiments/test \
--device cuda \
--repeat_times 1
```
You can obtain `N` different SR results by setting `repeat_time` as `N` to test the stability of CCSR. The data folder should be like this:
```
experiments/test
├── sample0 # the first group of SR results
└── sample1 # the second group of SR results
...
└── sampleN # the N-th group of SR results
```
## 📏 Evaluation
1. Calculate the Image Quality Assessment for each restored group.
Fill in the required information in [cal_iqa.py](cal_iqa/cal_iqa.py) and run, then you can obtain the evaluation results in the folder like this:
```
log_path
├── log_name_npy # save the IQA values of each restored group as the npy files
└── log_name.log # log recode
```
2. Calculate the G-STD value for the diffusion-based SR method.
Fill in the required information in [iqa_G-STD.py](cal_iqa/iqa_G-STD.py) and run, then you can obtain the mean IQA values of N restored groups and G-STD value.
3. Calculate the L-STD value for the diffusion-based SR method.
Fill in the required information in [iqa_L-STD.py](cal_iqa/iqa_L-STD.py) and run, then you can obtain the L-STD value.
## 🚋 Train
#### Step1: Prepare training data
1. Generate file list of training set and validation set.
```shell
python scripts/make_file_list.py \
--img_folder [hq_dir_path] \
--val_size [validation_set_size] \
--save_folder [save_dir_path] \
--follow_links
```
This script will collect all image files in `img_folder` and split them into training set and validation set automatically. You will get two file lists in `save_folder`, each line in a file list contains an absolute path of an image file:
```
save_dir_path
├── train.list # training file list
└── val.list # validation file list
```
2. Configure training set and validation set.
For real-world image restoration, fill in the following configuration files with appropriate values.
- [training set](configs/dataset/general_deg_stablesr_realesrgan_train.yaml) and [validation set](configs/dataset/general_deg_stablesr_realesrgan_val.yaml) for **Real-ESRGAN** degradation.
#### Step2: Train Stage1 Model
1. Download pretrained [Stable Diffusion v2.1](https://huggingface.co/stabilityai/stable-diffusion-2-1-base) to provide generative capabilities.
```shell
wget https://huggingface.co/stabilityai/stable-diffusion-2-1-base/resolve/main/v2-1_512-ema-pruned.ckpt --no-check-certificate
```
2. Create the initial model weights.
```shell
python scripts/make_stage2_init_weight.py \
--cldm_config configs/model/ccsr_stage1.yaml \
--sd_weight [sd_v2.1_ckpt_path] \
--output weights/init_weight_ccsr.ckpt
```
3. Configure training-related information.
Fill in the configuration file of [training of stage1](configs/train_ccsr_stage1.yaml) with appropriate settings.
4. Start training.
```shell
python train.py --config configs/train_ccsr_stage1.yaml
```
#### Step3: Train Stage2 Model
1. Configure training-related information.
Fill in the configuration file of [training of stage2](configs/train_ccsr_stage2.yaml) with appropriate settings.
2. Start training.
```shell
python train.py --config configs/train_ccsr_stage2.yaml
```
### Citations
If our code helps your research or work, please consider citing our paper.
The following are BibTeX references:
```
@article{sun2023ccsr,
title={Improving the Stability of Diffusion Models for Content Consistent Super-Resolution},
author={Sun, Lingchen and Wu, Rongyuan and Zhang, Zhengqiang and Yong, Hongwei and Zhang, Lei},
journal={arXiv preprint arXiv:2401.00877},
year={2024}
}
```
### License
This project is released under the [Apache 2.0 license](LICENSE).
### Acknowledgement
This project is based on [ControlNet](https://github.com/lllyasviel/ControlNet), [BasicSR](https://github.com/XPixelGroup/BasicSR) and [DiffBIR](https://github.com/XPixelGroup/DiffBIR). Some codes are brought from [StableSR](https://github.com/IceClear/StableSR). Thanks for their awesome works.
### Contact
If you have any questions, please contact: ling-chen.sun@connect.polyu.hk
<details>
<summary>statistics</summary>
![visitors](https://visitor-badge.laobi.icu/badge?page_id=csslc/CCSR)
</details>
+4
View File
@@ -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"]
Binary file not shown.
Binary file not shown.
-236
View File
@@ -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()
-111
View File
@@ -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))
-190
View File
@@ -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()
@@ -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]
@@ -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]
+6 -6
View File
@@ -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
-47
View File
@@ -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}"
-49
View File
@@ -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}"
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
-374
View File
@@ -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"])
-107
View File
@@ -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)
-109
View File
@@ -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)
-68
View File
@@ -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}"
)
-121
View File
@@ -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}
-183
View File
@@ -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)
BIN
View File
Binary file not shown.

Before

Width:  |  Height:  |  Size: 638 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 3.1 MiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 2.2 MiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 3.2 MiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 4.6 MiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 2.7 MiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 3.0 MiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 791 KiB

BIN
View File
Binary file not shown.

Before

Width:  |  Height:  |  Size: 30 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 1.6 MiB

BIN
View File
Binary file not shown.

Before

Width:  |  Height:  |  Size: 1.4 MiB

-230
View File
@@ -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()
Binary file not shown.
Binary file not shown.
Binary file not shown.
+8 -8
View File
@@ -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',
+9 -9
View File
@@ -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',
Binary file not shown.
Binary file not shown.
+2 -2
View File
@@ -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
+3 -3
View File
@@ -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
+1 -1
View File
@@ -1 +1 @@
from ldm.modules.losses.contperceptual import LPIPSWithDiscriminator
from ....ldm.modules.losses.contperceptual import LPIPSWithDiscriminator
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
+7 -7
View File
@@ -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):
+7 -7
View File
@@ -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):
+3 -3
View File
@@ -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
)
+7 -3
View File
@@ -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)
+155
View File
@@ -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",
}
Binary file not shown.

Before

Width:  |  Height:  |  Size: 13 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 2.7 KiB

+9
View File
@@ -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
+2 -18
View File
@@ -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
einops
-36
View File
@@ -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")
-80
View File
@@ -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.")
View File
-3
View File
@@ -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
-14
View File
@@ -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
-11
View File
@@ -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
-10
View File
@@ -1,10 +0,0 @@
python scripts/make_file_list.py \
--img_folder ****\
**** \
--val_size 10 \
--save_folder 'preset/train_datasets' \
--follow_links
-3
View File
@@ -1,3 +0,0 @@
python train.py --config configs/train_ccsr_stage1.yaml
-3
View File
@@ -1,3 +0,0 @@
python train.py --config configs/train_ccsr_stage2.yaml
+124
View File
@@ -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])
+139
View File
@@ -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()
+218
View File
@@ -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
@@ -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}
+70
View File
@@ -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
+176
View File
@@ -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"
@@ -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.
@@ -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)
+105
View File
@@ -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)
+38
View File
@@ -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)
+134
View File
@@ -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
+49
View File
@@ -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
+132
View File
@@ -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
+558
View File
@@ -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()

Some files were not shown because too many files have changed in this diff Show More