Args: cfg (CfgNode):
(self, cfg)
| 268 | """ |
| 269 | |
| 270 | def __init__(self, cfg): |
| 271 | """ |
| 272 | Args: |
| 273 | cfg (CfgNode): |
| 274 | """ |
| 275 | super().__init__() |
| 276 | logger = logging.getLogger("detectron2") |
| 277 | if not logger.isEnabledFor(logging.INFO): # setup_logger is not called for d2 |
| 278 | setup_logger() |
| 279 | cfg = DefaultTrainer.auto_scale_workers(cfg, comm.get_world_size()) |
| 280 | |
| 281 | # Assume these objects must be constructed in this order. |
| 282 | model = self.build_model(cfg) |
| 283 | optimizer = self.build_optimizer(cfg, model) |
| 284 | data_loader = self.build_train_loader(cfg) |
| 285 | |
| 286 | # For training, wrap with DDP. But don't need this for inference. |
| 287 | if comm.get_world_size() > 1: |
| 288 | model = DistributedDataParallel( |
| 289 | model, device_ids=[comm.get_local_rank()], broadcast_buffers=False |
| 290 | ) |
| 291 | self._trainer = (AMPTrainer if cfg.SOLVER.AMP.ENABLED else SimpleTrainer)( |
| 292 | model, data_loader, optimizer |
| 293 | ) |
| 294 | |
| 295 | self.scheduler = self.build_lr_scheduler(cfg, optimizer) |
| 296 | # Assume no other objects need to be checkpointed. |
| 297 | # We can later make it checkpoint the stateful hooks |
| 298 | self.checkpointer = DetectionCheckpointer( |
| 299 | # Assume you want to save checkpoints together with logs/statistics |
| 300 | model, |
| 301 | cfg.OUTPUT_DIR, |
| 302 | optimizer=optimizer, |
| 303 | scheduler=self.scheduler, |
| 304 | ) |
| 305 | self.start_iter = 0 |
| 306 | self.max_iter = cfg.SOLVER.MAX_ITER |
| 307 | self.cfg = cfg |
| 308 | |
| 309 | self.register_hooks(self.build_hooks()) |
| 310 | |
| 311 | def resume_or_load(self, resume=True): |
| 312 | """ |
nothing calls this directly
no test coverage detected