Initial commit
@@ -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
|
||||

|
||||
|
||||
## 😍 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.
|
||||
|
||||

|
||||
|
||||
### Comparisons on Bicubic SR
|
||||

|
||||
For more comparisons, please refer to our paper for details.
|
||||
|
||||
## 📝 Quantitative comparisons
|
||||
We propose new stability metrics, namely global standard deviation (G-STD) and local standard deviation (L-STD), to respectively measure the image-level and pixel-level variations of the SR results of diffusion-based methods.
|
||||
|
||||
More details about G-STD and L-STD can be found in our paper.
|
||||
|
||||

|
||||
## ⚙ Dependencies and Installation
|
||||
```shell
|
||||
## git clone this repository
|
||||
git clone https://github.com/csslc/CCSR.git
|
||||
cd CCSR
|
||||
|
||||
# create an environment with python >= 3.9
|
||||
conda create -n ccsr python=3.9
|
||||
conda activate ccsr
|
||||
pip install -r requirements.txt
|
||||
pip install -e git+https://github.com/CompVis/taming-transformers.git@master#egg=taming-transformers
|
||||
```
|
||||
## 🍭 Quick Inference
|
||||
#### Step 1: Download the pretrained models
|
||||
- Download the CCSR models from:
|
||||
|
||||
| Model Name | Description | GoogleDrive | OneDive |
|
||||
|:---------------------|:---------------------------------------------|:--------------------------------------------------------------------------------------|:------------------------------------------------------------------------------|
|
||||
| real-world_ccsr.ckpt | CCSR model for real-world image restoration. | [download](https://drive.google.com/drive/folders/1jM1mxDryPk9CTuFTvYcraP2XIVzbPiw_?usp=drive_link) | download |
|
||||
| bicubic_ccsr.ckpt | CCSR model for bicubic image restoration. | download | download |
|
||||
|
||||
|
||||
#### Step 2: Prepare testing data
|
||||
You can put the testing images in the `preset/test_datasets`.
|
||||
|
||||
#### Step 3: Running testing command
|
||||
```
|
||||
python inference_ccsr.py \
|
||||
--input preset/test_datasets \
|
||||
--config configs/model/ccsr_stage2.yaml \
|
||||
--ckpt weights/real-world_ccsr.ckpt \
|
||||
--steps 45 \
|
||||
--sr_scale 4 \
|
||||
--t_max 0.6667 \
|
||||
--t_min 0.3333 \
|
||||
--color_fix_type adain \
|
||||
--output experiments/test \
|
||||
--device cuda \
|
||||
--repeat_times 1
|
||||
```
|
||||
You can obtain `N` different SR results by setting `repeat_time` as `N` to test the stability of CCSR. The data folder should be like this:
|
||||
|
||||
```
|
||||
experiments/test
|
||||
├── sample0 # the first group of SR results
|
||||
└── sample1 # the second group of SR results
|
||||
...
|
||||
└── sampleN # the N-th group of SR results
|
||||
```
|
||||
|
||||
## 📏 Evaluation
|
||||
1. Calculate the Image Quality Assessment for each restored group.
|
||||
|
||||
Fill in the required information in [cal_iqa.py](cal_iqa/cal_iqa.py) and run, then you can obtain the evaluation results in the folder like this:
|
||||
```
|
||||
log_path
|
||||
├── log_name_npy # save the IQA values of each restored group as the npy files
|
||||
└── log_name.log # log recode
|
||||
```
|
||||
|
||||
2. Calculate the G-STD value for the diffusion-based SR method.
|
||||
|
||||
Fill in the required information in [iqa_G-STD.py](cal_iqa/iqa_G-STD.py) and run, then you can obtain the mean IQA values of N restored groups and G-STD value.
|
||||
|
||||
3. Calculate the L-STD value for the diffusion-based SR method.
|
||||
|
||||
Fill in the required information in [iqa_L-STD.py](cal_iqa/iqa_L-STD.py) and run, then you can obtain the L-STD value.
|
||||
|
||||
|
||||
## 🚋 Train
|
||||
|
||||
#### Step1: Prepare training data
|
||||
|
||||
1. Generate file list of training set and validation set.
|
||||
|
||||
```shell
|
||||
python scripts/make_file_list.py \
|
||||
--img_folder [hq_dir_path] \
|
||||
--val_size [validation_set_size] \
|
||||
--save_folder [save_dir_path] \
|
||||
--follow_links
|
||||
```
|
||||
|
||||
This script will collect all image files in `img_folder` and split them into training set and validation set automatically. You will get two file lists in `save_folder`, each line in a file list contains an absolute path of an image file:
|
||||
|
||||
```
|
||||
save_dir_path
|
||||
├── train.list # training file list
|
||||
└── val.list # validation file list
|
||||
```
|
||||
|
||||
2. Configure training set and validation set.
|
||||
|
||||
For real-world image restoration, fill in the following configuration files with appropriate values.
|
||||
|
||||
- [training set](configs/dataset/general_deg_stablesr_realesrgan_train.yaml) and [validation set](configs/dataset/general_deg_stablesr_realesrgan_val.yaml) for **Real-ESRGAN** degradation.
|
||||
|
||||
#### Step2: Train Stage1 Model
|
||||
1. Download pretrained [Stable Diffusion v2.1](https://huggingface.co/stabilityai/stable-diffusion-2-1-base) to provide generative capabilities.
|
||||
|
||||
```shell
|
||||
wget https://huggingface.co/stabilityai/stable-diffusion-2-1-base/resolve/main/v2-1_512-ema-pruned.ckpt --no-check-certificate
|
||||
```
|
||||
|
||||
2. Create the initial model weights.
|
||||
|
||||
```shell
|
||||
python scripts/make_stage2_init_weight.py \
|
||||
--cldm_config configs/model/ccsr_stage1.yaml \
|
||||
--sd_weight [sd_v2.1_ckpt_path] \
|
||||
--output weights/init_weight_ccsr.ckpt
|
||||
```
|
||||
|
||||
3. Configure training-related information.
|
||||
|
||||
Fill in the configuration file of [training of stage1](configs/train_ccsr_stage1.yaml) with appropriate settings.
|
||||
|
||||
4. Start training.
|
||||
|
||||
```shell
|
||||
python train.py --config configs/train_ccsr_stage1.yaml
|
||||
```
|
||||
|
||||
#### Step3: Train Stage2 Model
|
||||
1. Configure training-related information.
|
||||
|
||||
Fill in the configuration file of [training of stage2](configs/train_ccsr_stage2.yaml) with appropriate settings.
|
||||
|
||||
2. Start training.
|
||||
```shell
|
||||
python train.py --config configs/train_ccsr_stage2.yaml
|
||||
```
|
||||
|
||||
### Citations
|
||||
If our code helps your research or work, please consider citing our paper.
|
||||
The following are BibTeX references:
|
||||
|
||||
```
|
||||
@article{sun2023ccsr,
|
||||
title={Improving the Stability of Diffusion Models for Content Consistent Super-Resolution},
|
||||
author={Sun, Lingchen and Wu, Rongyuan and Zhang, Zhengqiang and Yong, Hongwei and Zhang, Lei},
|
||||
journal={arXiv preprint arXiv:2401.00877},
|
||||
year={2024}
|
||||
}
|
||||
```
|
||||
|
||||
### License
|
||||
This project is released under the [Apache 2.0 license](LICENSE).
|
||||
|
||||
### Acknowledgement
|
||||
This project is based on [ControlNet](https://github.com/lllyasviel/ControlNet), [BasicSR](https://github.com/XPixelGroup/BasicSR) and [DiffBIR](https://github.com/XPixelGroup/DiffBIR). Some codes are brought from [StableSR](https://github.com/IceClear/StableSR). Thanks for their awesome works.
|
||||
|
||||
### Contact
|
||||
If you have any questions, please contact: ling-chen.sun@connect.polyu.hk
|
||||
|
||||
|
||||
<details>
|
||||
<summary>statistics</summary>
|
||||
|
||||

|
||||
|
||||
</details>
|
||||
|
||||
|
||||
@@ -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"]
|
||||
@@ -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()
|
||||
@@ -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))
|
||||
@@ -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]
|
||||
@@ -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
|
||||
|
||||
@@ -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}"
|
||||
@@ -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}"
|
||||
@@ -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"])
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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}"
|
||||
)
|
||||
@@ -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}
|
||||
@@ -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)
|
||||
|
Before Width: | Height: | Size: 638 KiB |
|
Before Width: | Height: | Size: 3.1 MiB |
|
Before Width: | Height: | Size: 2.2 MiB |
|
Before Width: | Height: | Size: 3.2 MiB |
|
Before Width: | Height: | Size: 4.6 MiB |
|
Before Width: | Height: | Size: 2.7 MiB |
|
Before Width: | Height: | Size: 3.0 MiB |
|
Before Width: | Height: | Size: 791 KiB |
|
Before Width: | Height: | Size: 30 KiB |
|
Before Width: | Height: | Size: 1.6 MiB |
|
Before Width: | Height: | Size: 1.4 MiB |
@@ -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()
|
||||
@@ -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',
|
||||
|
||||
@@ -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',
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 @@
|
||||
from ldm.modules.losses.contperceptual import LPIPSWithDiscriminator
|
||||
from ....ldm.modules.losses.contperceptual import LPIPSWithDiscriminator
|
||||
@@ -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,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):
|
||||
|
||||
@@ -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,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)
|
||||
|
||||
@@ -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",
|
||||
}
|
||||
|
Before Width: | Height: | Size: 13 KiB |
|
Before Width: | Height: | Size: 2.7 KiB |
@@ -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
|
||||
@@ -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
|
||||
@@ -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")
|
||||
@@ -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.")
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -1,10 +0,0 @@
|
||||
|
||||
|
||||
|
||||
|
||||
python scripts/make_file_list.py \
|
||||
--img_folder ****\
|
||||
**** \
|
||||
--val_size 10 \
|
||||
--save_folder 'preset/train_datasets' \
|
||||
--follow_links
|
||||
@@ -1,3 +0,0 @@
|
||||
|
||||
|
||||
python train.py --config configs/train_ccsr_stage1.yaml
|
||||
@@ -1,3 +0,0 @@
|
||||
|
||||
|
||||
python train.py --config configs/train_ccsr_stage2.yaml
|
||||
@@ -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])
|
||||
@@ -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()
|
||||
@@ -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}
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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()
|
||||