毕设


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)