Main api for training model.
(model,
dataset,
cfg,
distributed=False,
validate=False,
timestamp=None,
device='cuda',
meta=None)
| 37 | |
| 38 | |
| 39 | def train_model(model, |
| 40 | dataset, |
| 41 | cfg, |
| 42 | distributed=False, |
| 43 | validate=False, |
| 44 | timestamp=None, |
| 45 | device='cuda', |
| 46 | meta=None): |
| 47 | """Main api for training model.""" |
| 48 | logger = get_root_logger(cfg.log_level) |
| 49 | |
| 50 | # prepare data loaders |
| 51 | dataset = dataset if isinstance(dataset, (list, tuple)) else [dataset] |
| 52 | |
| 53 | data_loaders = [ |
| 54 | build_dataloader( |
| 55 | ds, |
| 56 | cfg.data.samples_per_gpu, |
| 57 | cfg.data.workers_per_gpu, |
| 58 | # cfg.gpus will be ignored if distributed |
| 59 | num_gpus=len(cfg.gpu_ids), |
| 60 | dist=distributed, |
| 61 | round_up=True, |
| 62 | seed=cfg.seed) for ds in dataset |
| 63 | ] |
| 64 | |
| 65 | # determine whether use adversarial training precess or not |
| 66 | use_adverserial_train = cfg.get('use_adversarial_train', False) |
| 67 | |
| 68 | # put model on gpus |
| 69 | if distributed: |
| 70 | find_unused_parameters = cfg.get('find_unused_parameters', True) |
| 71 | # Sets the `find_unused_parameters` parameter in |
| 72 | # torch.nn.parallel.DistributedDataParallel |
| 73 | if use_adverserial_train: |
| 74 | # Use DistributedDataParallelWrapper for adversarial training |
| 75 | model = DistributedDataParallelWrapper( |
| 76 | model, |
| 77 | device_ids=[torch.cuda.current_device()], |
| 78 | broadcast_buffers=False, |
| 79 | find_unused_parameters=find_unused_parameters) |
| 80 | else: |
| 81 | model = MMDistributedDataParallel( |
| 82 | model.cuda(), |
| 83 | device_ids=[torch.cuda.current_device()], |
| 84 | broadcast_buffers=False, |
| 85 | find_unused_parameters=find_unused_parameters) |
| 86 | else: |
| 87 | if device == 'cuda': |
| 88 | model = MMDataParallel( |
| 89 | model.cuda(cfg.gpu_ids[0]), device_ids=cfg.gpu_ids) |
| 90 | elif device == 'cpu': |
| 91 | model = model.cpu() |
| 92 | else: |
| 93 | raise ValueError(F'unsupported device name {device}.') |
| 94 | |
| 95 | # build runner |
| 96 | optimizer = build_optimizers(model, cfg.optimizer) |
no test coverage detected