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

Function train_model

detrsmpl/apis/train.py:40–163  ·  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

38
39
40def train_model(model,
41 dataset,
42 cfg,
43 distributed=False,
44 validate=False,
45 timestamp=None,
46 device='cuda',
47 meta=None):
48 """Main api for training model."""
49 logger = get_root_logger(cfg.log_level)
50
51 # prepare data loaders
52 dataset = dataset if isinstance(dataset, (list, tuple)) else [dataset]
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', False)
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(model.cuda(cfg.gpu_ids[0]),
89 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)
97 if cfg.get('runner') is None:

Callers

nothing calls this directly

Calls 6

get_root_loggerFunction · 0.90
build_dataloaderFunction · 0.90
build_optimizersFunction · 0.90
build_datasetFunction · 0.90
getMethod · 0.45

Tested by

no test coverage detected