(self, model, save_dir="", *, save_to_disk=None, **checkpointables)
| 36 | |
| 37 | class SGTCheckPointer(Checkpointer): |
| 38 | def __init__(self, model, save_dir="", *, save_to_disk=None, **checkpointables): |
| 39 | if isinstance(model, DistributedDataParallel): |
| 40 | model = model.module |
| 41 | # super().__init__(model.detector, save_dir, **kwargs) |
| 42 | # tracker = model.tracker |
| 43 | is_main_process = comm.is_main_process() |
| 44 | super().__init__( |
| 45 | model, |
| 46 | save_dir, |
| 47 | save_to_disk=is_main_process if save_to_disk is None else save_to_disk, |
| 48 | **checkpointables, |
| 49 | ) |
| 50 | self.whole_weight_flag = False |
| 51 | self.backbone_name = self.get_backbone_name() # default by dla |
| 52 | backbone_convert_fn_dict = { |
| 53 | 'resnet': self._convert_weight_name_resnet, |
| 54 | 'dla': self._convert_weight_name_dla, |
| 55 | 'hourglass': self._convert_weight_name_hourglass, |
| 56 | } |
| 57 | self.convert_weight_name_fn = backbone_convert_fn_dict[self.backbone_name] |
| 58 | self.dcnv2_flag = False |
| 59 | |
| 60 | def get_backbone_name(self): |
| 61 | backbone_name = 'dla' |
nothing calls this directly
no test coverage detected