MCPcopy Create free account
hub / github.com/MotrixLab/insactor / train_model

Function train_model

diffplanner/apis/train.py:39–165  ·  view source on GitHub ↗

Main api for training model.

(model,
                dataset,
                cfg,
                distributed=False,
                validate=False,
                timestamp=None,
                device='cuda',
                meta=None)

Source from the content-addressed store, hash-verified

37
38
39def 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)

Callers 1

mainFunction · 0.90

Calls 5

get_root_loggerFunction · 0.90
build_dataloaderFunction · 0.90
build_optimizersFunction · 0.90
build_datasetFunction · 0.90

Tested by

no test coverage detected