(self)
| 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( |
no test coverage detected