RabbitRUI/ruispace
0
1import argparse2import logging3import os4 5import torch6import torch.distributed as dist7import torch.nn.functional as F8import torch.utils.data.distributed9from torch.nn.utils import clip_grad_norm_10 11import losses12from backbones import get_model13from dataset import MXFaceDataset, SyntheticDataset, DataLoaderX14from partial_fc import PartialFC15from utils.utils_amp import MaxClipGradScaler16from utils.utils_callbacks import CallBackVerification, CallBackLogging, CallBackModelCheckpoint17from utils.utils_config import get_config18from utils.utils_logging import AverageMeter, init_logging19 20 21def main(args):22 cfg = get_config(args.config)23 try:24 world_size = int(os.environ['WORLD_SIZE'])25 rank = int(os.environ['RANK'])26 dist.init_process_group('nccl')27 except KeyError:28 world_size = 129 rank = 030 dist.init_process_group(backend='nccl', init_method="tcp://127.0.0.1:12584", rank=rank, world_size=world_size)31 32 local_rank = args.local_rank33 torch.cuda.set_device(local_rank)34 os.makedirs(cfg.output, exist_ok=True)35 init_logging(rank, cfg.output)36 37 if cfg.rec == "synthetic":38 train_set = SyntheticDataset(local_rank=local_rank)39 else:40 train_set = MXFaceDataset(root_dir=cfg.rec, local_rank=local_rank)41 42 train_sampler = torch.utils.data.distributed.DistributedSampler(train_set, shuffle=True)43 train_loader = DataLoaderX(44 local_rank=local_rank, dataset=train_set, batch_size=cfg.batch_size,45 sampler=train_sampler, num_workers=2, pin_memory=True, drop_last=True)46 backbone = get_model(cfg.network, dropout=0.0, fp16=cfg.fp16, num_features=cfg.embedding_size).to(local_rank)47 48 if cfg.resume:49 try:50 backbone_pth = os.path.join(cfg.output, "backbone.pth")51 backbone.load_state_dict(torch.load(backbone_pth, map_location=torch.device(local_rank)))52 if rank == 0:53 logging.info("backbone resume successfully!")54 except (FileNotFoundError, KeyError, IndexError, RuntimeError):55 if rank == 0:56 logging.info("resume fail, backbone init successfully!")57 58 backbone = torch.nn.parallel.DistributedDataParallel(59 module=backbone, broadcast_buffers=False, device_ids=[local_rank])60 backbone.train()61 margin_softmax = losses.get_loss(cfg.loss)62 module_partial_fc = PartialFC(63 rank=rank, local_rank=local_rank, world_size=world_size, resume=cfg.resume,64 batch_size=cfg.batch_size, margin_softmax=margin_softmax, num_classes=cfg.num_classes,65 sample_rate=cfg.sample_rate, embedding_size=cfg.embedding_size, prefix=cfg.output)66 67 opt_backbone = torch.optim.SGD(68 params=[{'params': backbone.parameters()}],69 lr=cfg.lr / 512 * cfg.batch_size * world_size,70 momentum=0.9, weight_decay=cfg.weight_decay)71 opt_pfc = torch.optim.SGD(72 params=[{'params': module_partial_fc.parameters()}],73 lr=cfg.lr / 512 * cfg.batch_size * world_size,74 momentum=0.9, weight_decay=cfg.weight_decay)75 76 num_image = len(train_set)77 total_batch_size = cfg.batch_size * world_size78 cfg.warmup_step = num_image // total_batch_size * cfg.warmup_epoch79 cfg.total_step = num_image // total_batch_size * cfg.num_epoch80 81 def lr_step_func(current_step):82 cfg.decay_step = [x * num_image // total_batch_size for x in cfg.decay_epoch]83 if current_step < cfg.warmup_step:84 return current_step / cfg.warmup_step85 else:86 return 0.1 ** len([m for m in cfg.decay_step if m <= current_step])87 88 scheduler_backbone = torch.optim.lr_scheduler.LambdaLR(89 optimizer=opt_backbone, lr_lambda=lr_step_func)90 scheduler_pfc = torch.optim.lr_scheduler.LambdaLR(91 optimizer=opt_pfc, lr_lambda=lr_step_func)92 93 for key, value in cfg.items():94 num_space = 25 - len(key)95 logging.info(": " + key + " " * num_space + str(value))96 97 val_target = cfg.val_targets98 callback_verification = CallBackVerification(2000, rank, val_target, cfg.rec)99 callback_logging = CallBackLogging(50, rank, cfg.total_step, cfg.batch_size, world_size, None)100 callback_checkpoint = CallBackModelCheckpoint(rank, cfg.output)101 102 loss = AverageMeter()103 start_epoch = 0104 global_step = 0105 grad_amp = MaxClipGradScaler(cfg.batch_size, 128 * cfg.batch_size, growth_interval=100) if cfg.fp16 else None106 for epoch in range(start_epoch, cfg.num_epoch):107 train_sampler.set_epoch(epoch)108 for step, (img, label) in enumerate(train_loader):109 global_step += 1110 features = F.normalize(backbone(img))111 x_grad, loss_v = module_partial_fc.forward_backward(label, features, opt_pfc)112 if cfg.fp16:113 features.backward(grad_amp.scale(x_grad))114 grad_amp.unscale_(opt_backbone)115 clip_grad_norm_(backbone.parameters(), max_norm=5, norm_type=2)116 grad_amp.step(opt_backbone)117 grad_amp.update()118 else:119 features.backward(x_grad)120 clip_grad_norm_(backbone.parameters(), max_norm=5, norm_type=2)121 opt_backbone.step()122 123 opt_pfc.step()124 module_partial_fc.update()125 opt_backbone.zero_grad()126 opt_pfc.zero_grad()127 loss.update(loss_v, 1)128 callback_logging(global_step, loss, epoch, cfg.fp16, scheduler_backbone.get_last_lr()[0], grad_amp)129 callback_verification(global_step, backbone)130 scheduler_backbone.step()131 scheduler_pfc.step()132 callback_checkpoint(global_step, backbone, module_partial_fc)133 dist.destroy_process_group()134 135 136if __name__ == "__main__":137 torch.backends.cudnn.benchmark = True138 parser = argparse.ArgumentParser(description='PyTorch ArcFace Training')139 parser.add_argument('config', type=str, help='py config file')140 parser.add_argument('--local_rank', type=int, default=0, help='local_rank')141 main(parser.parse_args())142 