MCPcopy Create free account
hub / github.com/chenhaoxing/DiffusionInst / __init__

Method __init__

train_net.py:40–79  ·  view source on GitHub ↗

Args: cfg (CfgNode):

(self, cfg)

Source from the content-addressed store, hash-verified

38 """ Extension of the Trainer class adapted to DiffusionInst. """
39
40 def __init__(self, cfg):
41 """
42 Args:
43 cfg (CfgNode):
44 """
45 super(DefaultTrainer, self).__init__() # call grandfather's `__init__` while avoid father's `__init()`
46 logger = logging.getLogger("detectron2")
47 if not logger.isEnabledFor(logging.INFO): # setup_logger is not called for d2
48 setup_logger()
49 cfg = DefaultTrainer.auto_scale_workers(cfg, comm.get_world_size())
50
51 # Assume these objects must be constructed in this order.
52 model = self.build_model(cfg)
53 optimizer = self.build_optimizer(cfg, model)
54 data_loader = self.build_train_loader(cfg)
55
56 model = create_ddp_model(model, broadcast_buffers=False)
57 self._trainer = (AMPTrainer if cfg.SOLVER.AMP.ENABLED else SimpleTrainer)(
58 model, data_loader, optimizer
59 )
60
61 self.scheduler = self.build_lr_scheduler(cfg, optimizer)
62
63 ########## EMA ############
64 kwargs = {
65 'trainer': weakref.proxy(self),
66 }
67 kwargs.update(may_get_ema_checkpointer(cfg, model))
68 self.checkpointer = DetectionCheckpointer(
69 # Assume you want to save checkpoints together with logs/statistics
70 model,
71 cfg.OUTPUT_DIR,
72 **kwargs,
73 # trainer=weakref.proxy(self),
74 )
75 self.start_iter = 0
76 self.max_iter = cfg.SOLVER.MAX_ITER
77 self.cfg = cfg
78
79 self.register_hooks(self.build_hooks())
80
81 @classmethod
82 def build_model(cls, cfg):

Callers

nothing calls this directly

Calls 6

build_modelMethod · 0.95
build_optimizerMethod · 0.95
build_train_loaderMethod · 0.95
build_hooksMethod · 0.95
may_get_ema_checkpointerFunction · 0.90
updateMethod · 0.45

Tested by

no test coverage detected