CoolFace
Apppublic

pengsida/NeuralBody

sourceHugging Faceupdated 2y agoView on Hugging Face
1likes
train_net.py109 linesDownload Raw Back to root
1from lib.config import cfg, args2from lib.networks import make_network3from lib.train import make_trainer, make_optimizer, make_lr_scheduler, make_recorder, set_lr_scheduler4from lib.datasets import make_data_loader5from lib.utils.net_utils import load_model, save_model, load_network6from lib.evaluators import make_evaluator7import torch.multiprocessing8import torch9import torch.distributed as dist10import os11 12if cfg.fix_random:13    torch.manual_seed(0)14    torch.backends.cudnn.deterministic = True15    torch.backends.cudnn.benchmark = False16 17 18def train(cfg, network):19    trainer = make_trainer(cfg, network)20    optimizer = make_optimizer(cfg, network)21    scheduler = make_lr_scheduler(cfg, optimizer)22    recorder = make_recorder(cfg)23    evaluator = make_evaluator(cfg)24 25    begin_epoch = load_model(network,26                             optimizer,27                             scheduler,28                             recorder,29                             cfg.trained_model_dir,30                             resume=cfg.resume)31    set_lr_scheduler(cfg, scheduler)32 33    train_loader = make_data_loader(cfg,34                                    is_train=True,35                                    is_distributed=cfg.distributed,36                                    max_iter=cfg.ep_iter)37    val_loader = make_data_loader(cfg, is_train=False)38 39    for epoch in range(begin_epoch, cfg.train.epoch):40        recorder.epoch = epoch41        if cfg.distributed:42            train_loader.batch_sampler.sampler.set_epoch(epoch)43 44        trainer.train(epoch, train_loader, optimizer, recorder)45        scheduler.step()46 47        if (epoch + 1) % cfg.save_ep == 0 and cfg.local_rank == 0:48            save_model(network, optimizer, scheduler, recorder,49                       cfg.trained_model_dir, epoch)50 51        if (epoch + 1) % cfg.save_latest_ep == 0 and cfg.local_rank == 0:52            save_model(network,53                       optimizer,54                       scheduler,55                       recorder,56                       cfg.trained_model_dir,57                       epoch,58                       last=True)59 60        if (epoch + 1) % cfg.eval_ep == 0:61            trainer.val(epoch, val_loader, evaluator, recorder)62 63    return network64 65 66def test(cfg, network):67    trainer = make_trainer(cfg, network)68    val_loader = make_data_loader(cfg, is_train=False)69    evaluator = make_evaluator(cfg)70    epoch = load_network(network,71                         cfg.trained_model_dir,72                         resume=cfg.resume,73                         epoch=cfg.test.epoch)74    trainer.val(epoch, val_loader, evaluator)75 76 77def synchronize():78    """79    Helper function to synchronize (barrier) among all processes when80    using distributed training81    """82    if not dist.is_available():83        return84    if not dist.is_initialized():85        return86    world_size = dist.get_world_size()87    if world_size == 1:88        return89    dist.barrier()90 91 92def main():93    if cfg.distributed:94        cfg.local_rank = int(os.environ['RANK']) % torch.cuda.device_count()95        torch.cuda.set_device(cfg.local_rank)96        torch.distributed.init_process_group(backend="nccl",97                                             init_method="env://")98        synchronize()99 100    network = make_network(cfg)101    if args.test:102        test(cfg, network)103    else:104        train(cfg, network)105 106 107if __name__ == "__main__":108    main()109