(self)
| 285 | self.create_lr_scheduler() |
| 286 | |
| 287 | def pre_run(self): |
| 288 | tmp = self.tmp |
| 289 | tmp.vbatch_time = AverageMeter(10) |
| 290 | tmp.vdata_time = AverageMeter(10) |
| 291 | tmp.vloss = AverageMeter(10) |
| 292 | tmp.vtop1 = AverageMeter(10) |
| 293 | |
| 294 | tmp.loss_list = [torch.Tensor(1).cuda() for _ in range(self.C.world_size)] |
| 295 | tmp.top1_list = [torch.Tensor(1).cuda() for _ in range(self.C.world_size)] |
| 296 | |
| 297 | tmp.vbackbone_grad_norm = AverageMeter(10) |
| 298 | tmp.backbone_grad_norm_list = [torch.Tensor(1).cuda() for _ in range(self.C.world_size)] |
| 299 | tmp.vneck_grad_norm = AverageMeter(10) |
| 300 | tmp.neck_grad_norm_list = [torch.Tensor(1).cuda() for _ in range(self.C.world_size)] |
| 301 | tmp.vdecoder_grad_norm = AverageMeter(10) |
| 302 | tmp.decoder_grad_norm_list = [torch.Tensor(1).cuda() for _ in range(self.C.world_size)] |
| 303 | |
| 304 | self.model.train() |
| 305 | # if self.fix_bn: |
| 306 | # names = freeze_bn(self.model) |
| 307 | # if self.C.rank == 0: |
| 308 | # for name in names: |
| 309 | # self.logger.info('fixing BN [{}]'.format(name)) |
| 310 | |
| 311 | def prepare_data(self): |
| 312 | ginfo = self.ginfo |
no test coverage detected