CoolFace
Apppublic

RabbitRUI/ruispace

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
train.py142 linesDownload Raw Back to arcface_torch
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