| 13 | from projects.CenterNet.centernet.checkpoint.centernet_checkpoint import CenterNetCheckpointer |
| 14 | |
| 15 | class Trainer(DefaultEpochTrainer): |
| 16 | |
| 17 | @classmethod |
| 18 | def build_checkpointer(cls, model, output_dir, optimizer, scheduler): |
| 19 | checkpointer = CenterNetCheckpointer( |
| 20 | # Assume you want to save checkpoints together with logs/statistics |
| 21 | model, |
| 22 | output_dir, |
| 23 | optimizer=optimizer, |
| 24 | scheduler=scheduler, |
| 25 | ) |
| 26 | return checkpointer |
| 27 | |
| 28 | @classmethod |
| 29 | def build_train_loader(cls, cfg): |
| 30 | if "mix" in cfg.DATASETS.TRAIN[0]: |
| 31 | assert all(["mix" in t for t in cfg.DATASETS.TRAIN]), "Joint dataset only accepts only mix format" |
| 32 | mapper = MOTCenterNetDatasetMapper(cfg, is_train=True) |
| 33 | return build_mix_train_loader(cfg, mapper=mapper) |
| 34 | elif "mot" in cfg.DATASETS.TRAIN[0]: |
| 35 | mapper = MOTCenterNetDatasetMapper(cfg, is_train=True) |
| 36 | return build_mot_train_loader(cfg, mapper=mapper) |
| 37 | else: |
| 38 | mapper = COCOCenterNetDatasetMapper(cfg, is_train=True) |
| 39 | return build_detection_train_loader(cfg, mapper=mapper) |
| 40 | |
| 41 | @classmethod |
| 42 | def build_test_loader(cls, cfg, dataset_name): |
| 43 | if "mot" in cfg.DATASETS.TEST[0]: |
| 44 | mapper = MOTCenterNetDatasetMapper(cfg, is_train=False) |
| 45 | return build_mot_test_loader(cfg, dataset_name, mapper=mapper) |
| 46 | else: |
| 47 | mapper = COCOCenterNetDatasetMapper(cfg, is_train=False) |
| 48 | return build_detection_test_loader(cfg, dataset_name, mapper=mapper) |
| 49 | |
| 50 | @classmethod |
| 51 | def build_evaluator(cls, cfg, dataset_name, output_folder=None): |
| 52 | if output_folder is None: |
| 53 | output_folder = os.path.join(cfg.OUTPUT_DIR, "inference") |
| 54 | evaluator_list = [] |
| 55 | evaluator_type = MetadataCatalog.get(dataset_name).evaluator_type |
| 56 | if evaluator_type in ["coco", "coco_panoptic_seg"]: |
| 57 | evaluator_list.append(COCOEvaluator(dataset_name, cfg, True, output_folder)) |
| 58 | if evaluator_type in ["mot"]: |
| 59 | evaluator_list.append(MotEvaluator(dataset_name, cfg, True, output_folder)) |
| 60 | return DatasetEvaluators(evaluator_list) |
| 61 | |
| 62 | def run_step(self): |
| 63 | start = time.perf_counter() |
| 64 | data = next(self._trainer._data_loader_iter) |
| 65 | data_time = time.perf_counter() - start |
| 66 | |
| 67 | loss_dict = self.model(data) |
| 68 | loss = loss_dict['total_loss'] |
| 69 | self._write_metrics(loss_dict, data_time, prefix="train") |
| 70 | |
| 71 | self.optimizer.zero_grad() |
| 72 | loss.backward() |