MCPcopy Create free account
hub / github.com/HYUNJS/SGT / Trainer

Class Trainer

projects/CenterNet/centernet/trainer.py:15–88  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

13from projects.CenterNet.centernet.checkpoint.centernet_checkpoint import CenterNetCheckpointer
14
15class 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()

Callers 1

mainFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected