深度学习模型训练代码注释详解
导入必要的库
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='模型名称') parser.add_argument('--num_workers', default=16, type=int, help='工作进程数量') parser.add_argument('--use_mp', action='store_true', default=False, help='是否使用混合精度训练') parser.add_argument('--use_ddp', action='store_true', default=False, help='是否使用分布式数据并行') parser.add_argument('--save_dir', default='./saved_models/', type=str, help='模型保存路径') parser.add_argument('--data_dir', default='./data/', type=str, help='数据集路径') parser.add_argument('--log_dir', default='./logs/', type=str, help='日志保存路径') parser.add_argument('--train_set', default='SOTS-OUT', type=str, help='训练数据集名称') parser.add_argument('--val_set', default='SOTS-IN/SOTS-IN', type=str, help='验证数据集名称') parser.add_argument('--exp', default='reside-in', type=str, help='实验设置') 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() # 设置当前GPU设备 torch.cuda.set_device(local_rank) # 打印信息 if local_rank == 0: print('==> 使用DDP.') 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' # default.json 作为基线的配置文件 with open(os.path.join('configs', args.exp, config_name), 'r') as f: # 加载模型配置文件 m_setup = json.load(f) print(m_setup)
定义函数:reduce_mean
功能:对张量进行平均归约操作,用于分布式训练中计算全局平均损失
def reduce_mean(tensor, nprocs): # 创建张量的副本 rt = tensor.clone() # 进行全局求和操作 dist.all_reduce(rt, op=dist.ReduceOp.SUM) # 计算平均值 rt /= nprocs # 返回结果 return rt
定义模型训练函数:train
功能:训练模型
def train(train_loader, network, criterion, optimizer, scaler, frozen_bn=False): # 初始化损失平均器 losses = AverageMeter() # 清空GPU缓存 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
定义模型验证函数:valid
功能:验证模型性能
def valid(val_loader, network): # 初始化PSNR平均器 PSNR = AverageMeter() # 清空GPU缓存 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:]
# 对源图像进行填充,以适应模型的patch size
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
psnr = 10 * torch.log10(1 / mse_loss).mean()
# 更新PSNR平均器
PSNR.update(psnr.item(), source_img.size(0))
# 返回平均PSNR
return PSNR.avg
定义主函数
def main(): # 定义模型 network = eval(args.model)() # 使用eval函数动态创建模型实例 # 将模型移动到GPU network.cuda()
# 使用DDP
if args.use_ddp:
# 使用分布式数据并行封装模型
network = DistributedDataParallel(network, device_ids=[local_rank], output_device=local_rank)
# 检查批次大小是否过小,如果过小则使用同步BN
if m_setup['batch_size'] // world_size < 16:
if local_rank == 0: print('==> 使用同步BN,因为规范化批次大小过小.')
# 将模型中的BN层转换为同步BN层
nn.SyncBatchNorm.convert_sync_batchnorm(network)
else:
# 使用数据并行封装模型
network = DataParallel(network)
# 检查批次大小是否过小,如果过小则使用同步BN
if m_setup['batch_size'] // torch.cuda.device_count() < 16:
print('==> 使用同步BN,因为规范化批次大小过小.')
# 将模型中的BN层转换为同步BN层
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['const_epochs'])
# 定义权重衰减调度器
wd_scheduler = CosineScheduler(optimizer, param_name='weight_decay', t_max=b_setup['epochs']) # 似乎不起作用
# 定义梯度缩放器
scaler = GradScaler()
# 加载已训练的模型
save_dir = os.path.join(args.save_dir, args.exp)
# 创建模型保存目录
os.makedirs(save_dir, exist_ok=True)
# 检查模型文件是否存在
if not os.path.exists(os.path.join(save_dir, args.model+'.pth')):
# 初始化最佳PSNR和当前训练轮数
best_psnr = 0
cur_epoch = 0
else:
# 打印信息
if not args.use_ddp or local_rank == 0: print('==> 加载已存在的训练模型.')
# 加载模型信息
model_info = torch.load(os.path.join(save_dir, args.model+'.pth'), map_location='cpu')
# 加载模型参数
network.load_state_dict(model_info['state_dict'])
# 加载优化器参数
optimizer.load_state_dict(model_info['optimizer'])
# 加载学习率调度器参数
lr_scheduler.load_state_dict(model_info['lr_scheduler'])
# 加载权重衰减调度器参数
wd_scheduler.load_state_dict(model_info['wd_scheduler'])
# 加载梯度缩放器参数
scaler.load_state_dict(model_info['scaler'])
# 设置当前训练轮数和最佳PSNR
cur_epoch = model_info['cur_epoch']
best_psnr = model_info['best_psnr']
# 定义数据集
# 训练数据集
train_dataset = PairLoader(os.path.join(args.data_dir, args.train_set), 'train',
b_setup['t_patch_size'],
b_setup['edge_decay'],
b_setup['data_augment'],
b_setup['cache_memory'])
# 创建训练数据加载器
train_loader = DataLoader(train_dataset,
batch_size=m_setup['batch_size'] // world_size,
sampler=RandomSampler(train_dataset, num_samples=b_setup['num_iter'] // world_size),
num_workers=args.num_workers // world_size,
pin_memory=True,
drop_last=True,
persistent_workers=True)
# 验证数据集
val_dataset = PairLoader(os.path.join(args.data_dir, args.val_set), b_setup['valid_mode'],
b_setup['v_patch_size'])
# 创建验证数据加载器
val_loader = DataLoader(val_dataset,
batch_size=max(int(m_setup['batch_size'] * b_setup['v_batch_ratio'] // world_size), 1),
num_workers=args.num_workers // world_size,
pin_memory=True)
# 开始训练
if not args.use_ddp or local_rank == 0:
# 打印信息
print('==> 开始训练,当前模型名称:' + args.model)
# 创建TensorBoard记录器
writer = SummaryWriter(log_dir=os.path.join(args.log_dir, args.exp, args.model))
# 训练循环
for epoch in tqdm(range(cur_epoch, b_setup['epochs'] + 1)):
# 设置是否冻结BN层
frozen_bn = epoch > (b_setup['epochs'] - b_setup['frozen_epochs'])
# 训练模型
loss = train(train_loader, network, criterion, optimizer, scaler, frozen_bn)
# 更新学习率调度器
lr_scheduler.step(epoch + 1)
# 更新权重衰减调度器
wd_scheduler.step(epoch + 1)
# 记录训练损失
if not args.use_ddp or local_rank == 0:
writer.add_scalar('train_loss', loss, epoch)
# 定期进行模型验证
if epoch % b_setup['eval_freq'] == 0:
# 验证模型性能
avg_psnr = valid(val_loader, network)
# 打印验证结果并保存模型
if not args.use_ddp or local_rank == 0:
# 更新最佳PSNR
if avg_psnr > best_psnr:
best_psnr = avg_psnr
# 保存模型
torch.save({'cur_epoch': epoch + 1,
'best_psnr': best_psnr,
'state_dict': network.state_dict(),
'optimizer' : optimizer.state_dict(),
'lr_scheduler' : lr_scheduler.state_dict(),
'wd_scheduler' : wd_scheduler.state_dict(),
'scaler' : scaler.state_dict()}, os.path.join(save_dir, args.model+'.pth'))
# 记录验证结果
writer.add_scalar('valid_psnr', avg_psnr, epoch)
writer.add_scalar('best_psnr', best_psnr, epoch)
# 使用DDP进行同步
if args.use_ddp: dist.barrier()
程序入口
if name == 'main': # 运行主函数 main()
原文地址: https://www.cveoy.top/t/topic/nwRk 著作权归作者所有。请勿转载和采集!