导入必要的库

import os import argparse import json import torch import torch.nn as nn import torch.nn.functional as F import torch.distributed as dist from torch.cuda.amp import autocast, GradScaler from torch.nn.parallel import DataParallel, DistributedDataParallel from torch.utils.data import DataLoader, RandomSampler, DistributedSampler from torch.utils.tensorboard import SummaryWriter from tqdm import tqdm

自定义库

from utils import AverageMeter, CosineScheduler, pad_img from datasets import PairLoader from models import *

定义命令行参数

parser = argparse.ArgumentParser() parser.add_argument('--model', default='gunet_t', type=str, help='model name') # 模型名称 parser.add_argument('--num_workers', default=16, type=int, help='number of workers') # 工作进程数 parser.add_argument('--use_mp', action='store_true', default=False, help='use Mixed Precision') # 是否使用混合精度训练 parser.add_argument('--use_ddp', action='store_true', default=False, help='use Distributed Data Parallel') # 是否使用分布式训练 parser.add_argument('--save_dir', default='./saved_models/', type=str, help='path to models saving') # 模型保存路径 parser.add_argument('--data_dir', default='./data/', type=str, help='path to dataset') # 数据集路径 parser.add_argument('--log_dir', default='./logs/', type=str, help='path to logs') # 训练日志路径 parser.add_argument('--train_set', default='SOTS-OUT', type=str, help='train dataset name') # 训练集名称 parser.add_argument('--val_set', default='SOTS-IN/SOTS-IN', type=str, help='valid dataset name') # 验证集名称 parser.add_argument('--exp', default='reside-in', type=str, help='experiment setting') # 实验设置 args = parser.parse_args()

训练环境

if args.use_ddp: torch.distributed.init_process_group(backend='nccl', init_method='env://') # 初始化分布式训练 world_size = 16 # 分布式训练时设备总数 local_rank = dist.get_rank() # 获取当前进程的rank torch.cuda.set_device(local_rank) # 设置当前进程使用的GPU设备 if local_rank == 0: print('==> Using DDP.') # 仅进程0输出信息 else: world_size = 16 # 设备总数

训练配置

with open(os.path.join('configs', args.exp, 'base.json'), 'r') as f: b_setup = json.load(f) # 读取基本配置 print(b_setup)

print('----------------------------------------------------------------------------')

variant = args.model.split('')[-1] config_name = 'model'+variant+'.json' if variant in ['t', 's', 'b', 'd'] else 'default.json' # 模型配置文件名称 with open(os.path.join('configs', args.exp, config_name), 'r') as f: m_setup = json.load(f) # 读取模型配置 print(m_setup)

求平均值

def reduce_mean(tensor, nprocs): rt = tensor.clone() # 返回tensor的拷贝 dist.all_reduce(rt, op=dist.ReduceOp.SUM) # 执行求和操作 rt /= nprocs # 求平均 return rt

训练函数

def train(train_loader, network, criterion, optimizer, scaler, frozen_bn=False): losses = AverageMeter() # 损失函数值的平均值

torch.cuda.empty_cache() # 清空缓存

network.eval() if frozen_bn else network.train() # 设置网络模型的训练状态

for batch in train_loader: # 遍历训练集
    source_img = batch['source'].cuda() # 输入图像
    target_img = batch['target'].cuda() # 目标图像

    with autocast(args.use_mp): # 自动混合精度训练
        output = network(source_img) # 模型输出
        loss = criterion(output, target_img) # 计算损失函数值

    optimizer.zero_grad() # 优化器梯度归零
    scaler.scale(loss).backward() # 反向传播
    scaler.step(optimizer) # 参数更新
    scaler.update() # 更新梯度缩放器

    if args.use_ddp: loss = reduce_mean(loss, dist.get_world_size()) # 求平均损失函数值
    losses.update(loss.item()) # 更新平均损失函数值

return losses.avg # 返回平均损失函数值

验证函数

def valid(val_loader, network): PSNR = AverageMeter() # PSNR的平均值

torch.cuda.empty_cache() # 清空缓存

network.eval() # 设置网络模型的验证状态

for batch in val_loader: # 遍历验证集
    source_img = batch['source'].cuda() # 输入图像
    target_img = batch['target'].cuda() # 目标图像

    with torch.no_grad(): # 不进行梯度计算
        H, W = source_img.shape[2:]
        source_img = pad_img(source_img, network.module.patch_size if hasattr(network.module, 'patch_size') else 16)
        output = network(source_img).clamp_(-1, 1)
        output = output[:, :, :H, :W]

    mse_loss = F.mse_loss(output * 0.5 + 0.5, target_img * 0.5 + 0.5, reduction='none').mean((1, 2, 3))
    psnr = 10 * torch.log10(1 / mse_loss).mean() # 计算PSNR

    PSNR.update(psnr.item(), source_img.size(0)) # 更新平均PSNR值

return PSNR.avg # 返回平均PSNR值

主函数

def main(): # 定义网络模型 network = eval(args.model)() # 根据模型名称创建模型 network.cuda() # 将模型移动到GPU上

if args.use_ddp:
    network = DistributedDataParallel(network, device_ids=[local_rank], output_device=local_rank) # 分布式训练
    if m_setup['batch_size'] // world_size < 16:
        if local_rank == 0: print('==> Using SyncBN because of too small norm-batch-size.')
        nn.SyncBatchNorm.convert_sync_batchnorm(network) # 将批量归一化(BN)转换为同步批量归一化(SyncBN)
else:
    network = DataParallel(network) # 多GPU训练
    if m_setup['batch_size'] // torch.cuda.device_count() < 16:
        print('==> Using SyncBN because of too small norm-batch-size.')
        convert_model(network) # 转换模型

# 定义损失函数
criterion = nn.L1Loss()

# 定义优化器
optimizer = torch.optim.AdamW(network.parameters(), lr=m_setup['lr'], weight_decay=b_setup['weight_decay'])
lr_scheduler = CosineScheduler(optimizer, param_name='lr', t_max=b_setup['epochs'], value_min=m_setup['lr'] * 1e-2, 
                               warmup_t=b_setup['warmup_epochs'], const_t=b_setup['cons
代码注释:图像超分辨率模型训练代码详解

原文地址: http://www.cveoy.top/t/topic/nwS0 著作权归作者所有。请勿转载和采集!

免费AI点我,无需注册和登录