MCPcopy Create free account
hub / github.com/PeizeSun/SparseR-CNN / do_train

Function do_train

tools/plain_train_net.py:126–189  ·  view source on GitHub ↗
(cfg, model, resume=False)

Source from the content-addressed store, hash-verified

124
125
126def do_train(cfg, model, resume=False):
127 model.train()
128 optimizer = build_optimizer(cfg, model)
129 scheduler = build_lr_scheduler(cfg, optimizer)
130
131 checkpointer = DetectionCheckpointer(
132 model, cfg.OUTPUT_DIR, optimizer=optimizer, scheduler=scheduler
133 )
134 start_iter = (
135 checkpointer.resume_or_load(cfg.MODEL.WEIGHTS, resume=resume).get("iteration", -1) + 1
136 )
137 max_iter = cfg.SOLVER.MAX_ITER
138
139 periodic_checkpointer = PeriodicCheckpointer(
140 checkpointer, cfg.SOLVER.CHECKPOINT_PERIOD, max_iter=max_iter
141 )
142
143 writers = (
144 [
145 CommonMetricPrinter(max_iter),
146 JSONWriter(os.path.join(cfg.OUTPUT_DIR, "metrics.json")),
147 TensorboardXWriter(cfg.OUTPUT_DIR),
148 ]
149 if comm.is_main_process()
150 else []
151 )
152
153 # compared to "train_net.py", we do not support accurate timing and
154 # precise BN here, because they are not trivial to implement in a small training loop
155 data_loader = build_detection_train_loader(cfg)
156 logger.info("Starting training from iteration {}".format(start_iter))
157 with EventStorage(start_iter) as storage:
158 for data, iteration in zip(data_loader, range(start_iter, max_iter)):
159 iteration = iteration + 1
160 storage.step()
161
162 loss_dict = model(data)
163 losses = sum(loss_dict.values())
164 assert torch.isfinite(losses).all(), loss_dict
165
166 loss_dict_reduced = {k: v.item() for k, v in comm.reduce_dict(loss_dict).items()}
167 losses_reduced = sum(loss for loss in loss_dict_reduced.values())
168 if comm.is_main_process():
169 storage.put_scalars(total_loss=losses_reduced, **loss_dict_reduced)
170
171 optimizer.zero_grad()
172 losses.backward()
173 optimizer.step()
174 storage.put_scalar("lr", optimizer.param_groups[0]["lr"], smoothing_hint=False)
175 scheduler.step()
176
177 if (
178 cfg.TEST.EVAL_PERIOD > 0
179 and iteration % cfg.TEST.EVAL_PERIOD == 0
180 and iteration != max_iter
181 ):
182 do_test(cfg, model)
183 # Compared to "train_net.py", the test results are not dumped to EventStorage

Callers 1

plain_train_net.pyFile · 0.85

Calls 15

build_optimizerFunction · 0.90
build_lr_schedulerFunction · 0.90
CommonMetricPrinterClass · 0.90
JSONWriterClass · 0.90
TensorboardXWriterClass · 0.90
EventStorageClass · 0.90
do_testFunction · 0.85
resume_or_loadMethod · 0.80
put_scalarsMethod · 0.80

Tested by

no test coverage detected