import matplotlib.pyplot as plt
from mmedit.apis import init_model
from mmedit.apis import restoration_inference
import setuptools
from mmcv.runner import set_random_seed
from mmcv import Config
import torch
import os.path as osp
from mmedit.datasets import build_dataset
from mmedit.models import build_model
from mmedit.apis import train_model
from mmcv.runner import init_dist
import os
import mmedit
import mmcv
def show_two():
img_LR = mmcv.imread('./data/Set5/LR/butterfly.png', channel_order='rgb')
img_HR = mmcv.imread('./data/Set5/GT/butterfly.png', channel_order='rgb')
plt.figure(figsize=(12,8))
plt.subplot(1,2,1)
plt.imshow(img_LR)
plt.subplot(1,2,2)
plt.imshow(img_HR)
plt.show()
def compare():
config_file = 'configs/restorers/basicvsr_plusplus/basicvsr_plusplus_c64n7_4x2_300k_vimeo90k_bd.py'
checkpint_file = 'checkpoint/spynet_20210409-c6c1bd09.pth'
model = init_model(config_file, checkpint_file, device='cuda:0')
result = restoration_inference(model, 'data/Set5/LR/butterfly.png')
result = torch.clamp(result, 0, 1)
img_SR = result.squeeze(0).permute(1, 2, 0).numpy()
img_LR = mmcv.imread('./data/Set5/LR/butterfly.png', channel_order='rgb')
img_HR = mmcv.imread('./data/Set5/GT/butterfly.png', channel_order='rgb')
fig = plt.figure(figsize=(15, 12))
ax1 = fig.add_subplot(1, 3, 1)
plt.title('LR', fontsize=16)
ax1.axis('off')
ax2 = fig.add_subplot(1, 3, 2)
plt.title('SR output', fontsize=16)
ax2.axis('off')
ax3 = fig.add_subplot(1, 3, 3)
plt.title('HR', fontsize=16)
ax3.axis('off')
ax1.imshow(img_LR)
ax2.imshow(img_SR)
ax3.imshow(img_HR)
plt.show()
def train(cfg):
# 构建数据集
datasets = {build_dataset(cfg.data.train)}
# 构建模型
model = build_model(cfg.model, train_cfg=cfg.train_cfg, test_cfg=cfg.test_cfg)
# 创建工作路径
mmcv.mkdir_or_exist(osp.abspath(cfg.work_dir))
# 额外信息
meta = dict()
if cfg.get('exo_name', None) is None:
cfg['exp_name'] = osp.splitext(osp.basename(cfg.work_dir))[0]
meta['exp_name'] = cfg.exp_name
meta['mmedit Version'] = mmedit.__version__
meta['seed'] = 0
# 启动训练
train_model(model, datasets, cfg, distributed=True, validate=True, meta=meta)
if __name__ == '__main__':
cfg = Config.fromfile('configs/restorers/basicvsr_plusplus/basicvsr_plusplus_c64n7_4x2_300k_vimeo90k_bd.py')
# 指定训练集的目录和标注文件
cfg.data.train.dataset.lq_folder = './data/DIV2K/LR'
cfg.data.train.dataset.gt_folder = './data/DIV2K/GT'
cfg.data.train.dataset.ann_file = './data/training_ann.txt'
# 指定验证集的目录
cfg.data.val.lq_folder = './data/Set5/LR'
cfg.data.val.gt_folder = './data/Set5/GT'
# 指定预训练模型
cfg.load_from = './checkpiont/spynet_20210409-c6c1bd09.pth'
# 设置工作目录
cfg.work_dir = './tutprial_exps/basic_plusplus'
# 配置batch_size
cfg.data.samples_per_gpu = 4
cfg.data.workers_per_gpu = 0
cfg.data.val_workers_per_gpu = 0
# 设置总迭代次数
cfg.total_iters = 200
# 在100次迭代时降低学习率
cfg.lr_config = {}
cfg.lr_config.policy = 'Step'
cfg.lr_config.by_epoch = False
cfg.lr_config.step = {100}
cfg.lr_config.gamma = 0.5
cfg.evaluation.interval = 200
cfg.checkpoint_config.interval = 200
cfg.log_config.interval = 40
cfg.seed = 0
set_random_seed(0, deterministic=False)
cfg.gpus = 1
# print(f'Config:\n{cfg.pretty_text}')
train(cfg)