MCPcopy Create free account
hub / github.com/FoundationVision/ByteTrack / before_train

Method before_train

yolox/core/trainer.py:125–179  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

123 )
124
125 def before_train(self):
126 logger.info("args: {}".format(self.args))
127 logger.info("exp value:\n{}".format(self.exp))
128
129 # model related init
130 torch.cuda.set_device(self.local_rank)
131 model = self.exp.get_model()
132 logger.info(
133 "Model Summary: {}".format(get_model_info(model, self.exp.test_size))
134 )
135 model.to(self.device)
136
137 # solver related init
138 self.optimizer = self.exp.get_optimizer(self.args.batch_size)
139
140 # value of epoch will be set in `resume_train`
141 model = self.resume_train(model)
142
143 # data related init
144 self.no_aug = self.start_epoch >= self.max_epoch - self.exp.no_aug_epochs
145 self.train_loader = self.exp.get_data_loader(
146 batch_size=self.args.batch_size,
147 is_distributed=self.is_distributed,
148 no_aug=self.no_aug,
149 )
150 logger.info("init prefetcher, this might take one minute or less...")
151 self.prefetcher = DataPrefetcher(self.train_loader)
152 # max_iter means iters per epoch
153 self.max_iter = len(self.train_loader)
154
155 self.lr_scheduler = self.exp.get_lr_scheduler(
156 self.exp.basic_lr_per_img * self.args.batch_size, self.max_iter
157 )
158 if self.args.occupy:
159 occupy_mem(self.local_rank)
160
161 if self.is_distributed:
162 model = DDP(model, device_ids=[self.local_rank], broadcast_buffers=False)
163
164 if self.use_model_ema:
165 self.ema_model = ModelEMA(model, 0.9998)
166 self.ema_model.updates = self.max_iter * self.start_epoch
167
168 self.model = model
169 self.model.train()
170
171 self.evaluator = self.exp.get_evaluator(
172 batch_size=self.args.batch_size, is_distributed=self.is_distributed
173 )
174 # Tensorboard logger
175 if self.rank == 0:
176 self.tblogger = SummaryWriter(self.file_name)
177
178 logger.info("Training start...")
179 #logger.info("\n{}".format(model))
180
181 def after_train(self):
182 logger.info(

Callers 1

trainMethod · 0.95

Calls 11

resume_trainMethod · 0.95
get_model_infoFunction · 0.90
DataPrefetcherClass · 0.90
occupy_memFunction · 0.90
ModelEMAClass · 0.90
trainMethod · 0.80
get_modelMethod · 0.45
get_optimizerMethod · 0.45
get_data_loaderMethod · 0.45
get_lr_schedulerMethod · 0.45
get_evaluatorMethod · 0.45

Tested by

no test coverage detected