MCPcopy Create free account
hub / github.com/ChenhongyiYang/QueryDet-PyTorch / __init__

Method __init__

train_tools/coco_train.py:60–102  ·  view source on GitHub ↗

Args: cfg (CfgNode):

(self, cfg, resume=False, reuse_ckpt=False)

Source from the content-addressed store, hash-verified

58
59class Trainer(DefaultTrainer):
60 def __init__(self, cfg, resume=False, reuse_ckpt=False):
61 """
62 Args:
63 cfg (CfgNode):
64 """
65 super(DefaultTrainer, self).__init__()
66
67 logger = logging.getLogger("detectron2")
68 if not logger.isEnabledFor(logging.INFO): # setup_logger is not called for d2
69 setup_logger()
70 cfg = DefaultTrainer.auto_scale_workers(cfg, comm.get_world_size())
71
72 # Assume these objects must be constructed in this order.
73 model = self.build_model(cfg)
74
75 ckpt = DetectionCheckpointer(model)
76 self.start_iter = 0
77 self.start_iter = ckpt.resume_or_load(cfg.MODEL.WEIGHTS, resume=resume).get("iteration", -1) + 1
78 self.iter =self.start_iter
79
80 optimizer = self.build_optimizer(cfg, model)
81 data_loader = self.build_train_loader(cfg)
82
83 # For training, wrap with DDP. But don't need this for inference.
84 if comm.get_world_size() > 1:
85 model = DistributedDataParallel(
86 model, device_ids=[comm.get_local_rank()], broadcast_buffers=False
87 )
88 self._trainer = (AMPTrainer if cfg.SOLVER.AMP.ENABLED else SimpleTrainer)(
89 model, data_loader, optimizer
90 )
91
92 self.scheduler = self.build_lr_scheduler(cfg, optimizer)
93 self.checkpointer = DetectionCheckpointer(
94 model,
95 cfg.OUTPUT_DIR,
96 optimizer=optimizer,
97 scheduler=self.scheduler,
98 )
99 self.start_iter = 0
100 self.max_iter = cfg.SOLVER.MAX_ITER
101 self.cfg = cfg
102 self.register_hooks(self.build_hooks())
103
104 @classmethod
105 def build_evaluator(cls, cfg, dataset_name, output_folder=None):

Callers

nothing calls this directly

Calls 2

build_train_loaderMethod · 0.95
resume_or_loadMethod · 0.80

Tested by

no test coverage detected