pengsida/NeuralBody
1
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 