Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model
0
1# -*- coding: utf-8 -*-2# Copyright (c) Facebook, Inc. and its affiliates.3 4import datetime5import itertools6import logging7import math8import operator9import os10import tempfile11import time12import warnings13from collections import Counter14import torch15from fvcore.common.checkpoint import Checkpointer16from fvcore.common.checkpoint import PeriodicCheckpointer as _PeriodicCheckpointer17from fvcore.common.param_scheduler import ParamScheduler18from fvcore.common.timer import Timer19from fvcore.nn.precise_bn import get_bn_modules, update_bn_stats20 21import detectron2.utils.comm as comm22from detectron2.evaluation.testing import flatten_results_dict23from detectron2.solver import LRMultiplier24from detectron2.solver import LRScheduler as _LRScheduler25from detectron2.utils.events import EventStorage, EventWriter26from detectron2.utils.file_io import PathManager27 28from .train_loop import HookBase29 30__all__ = [31 "CallbackHook",32 "IterationTimer",33 "PeriodicWriter",34 "PeriodicCheckpointer",35 "BestCheckpointer",36 "LRScheduler",37 "AutogradProfiler",38 "EvalHook",39 "PreciseBN",40 "TorchProfiler",41 "TorchMemoryStats",42]43 44 45"""46Implement some common hooks.47"""48 49 50class CallbackHook(HookBase):51 """52 Create a hook using callback functions provided by the user.53 """54 55 def __init__(self, *, before_train=None, after_train=None, before_step=None, after_step=None):56 """57 Each argument is a function that takes one argument: the trainer.58 """59 self._before_train = before_train60 self._before_step = before_step61 self._after_step = after_step62 self._after_train = after_train63 64 def before_train(self):65 if self._before_train:66 self._before_train(self.trainer)67 68 def after_train(self):69 if self._after_train:70 self._after_train(self.trainer)71 # The functions may be closures that hold reference to the trainer72 # Therefore, delete them to avoid circular reference.73 del self._before_train, self._after_train74 del self._before_step, self._after_step75 76 def before_step(self):77 if self._before_step:78 self._before_step(self.trainer)79 80 def after_step(self):81 if self._after_step:82 self._after_step(self.trainer)83 84 85class IterationTimer(HookBase):86 """87 Track the time spent for each iteration (each run_step call in the trainer).88 Print a summary in the end of training.89 90 This hook uses the time between the call to its :meth:`before_step`91 and :meth:`after_step` methods.92 Under the convention that :meth:`before_step` of all hooks should only93 take negligible amount of time, the :class:`IterationTimer` hook should be94 placed at the beginning of the list of hooks to obtain accurate timing.95 """96 97 def __init__(self, warmup_iter=3):98 """99 Args:100 warmup_iter (int): the number of iterations at the beginning to exclude101 from timing.102 """103 self._warmup_iter = warmup_iter104 self._step_timer = Timer()105 self._start_time = time.perf_counter()106 self._total_timer = Timer()107 108 def before_train(self):109 self._start_time = time.perf_counter()110 self._total_timer.reset()111 self._total_timer.pause()112 113 def after_train(self):114 logger = logging.getLogger(__name__)115 total_time = time.perf_counter() - self._start_time116 total_time_minus_hooks = self._total_timer.seconds()117 hook_time = total_time - total_time_minus_hooks118 119 num_iter = self.trainer.storage.iter + 1 - self.trainer.start_iter - self._warmup_iter120 121 if num_iter > 0 and total_time_minus_hooks > 0:122 # Speed is meaningful only after warmup123 # NOTE this format is parsed by grep in some scripts124 logger.info(125 "Overall training speed: {} iterations in {} ({:.4f} s / it)".format(126 num_iter,127 str(datetime.timedelta(seconds=int(total_time_minus_hooks))),128 total_time_minus_hooks / num_iter,129 )130 )131 132 logger.info(133 "Total training time: {} ({} on hooks)".format(134 str(datetime.timedelta(seconds=int(total_time))),135 str(datetime.timedelta(seconds=int(hook_time))),136 )137 )138 139 def before_step(self):140 self._step_timer.reset()141 self._total_timer.resume()142 143 def after_step(self):144 # +1 because we're in after_step, the current step is done145 # but not yet counted146 iter_done = self.trainer.storage.iter - self.trainer.start_iter + 1147 if iter_done >= self._warmup_iter:148 sec = self._step_timer.seconds()149 self.trainer.storage.put_scalars(time=sec)150 else:151 self._start_time = time.perf_counter()152 self._total_timer.reset()153 154 self._total_timer.pause()155 156 157class PeriodicWriter(HookBase):158 """159 Write events to EventStorage (by calling ``writer.write()``) periodically.160 161 It is executed every ``period`` iterations and after the last iteration.162 Note that ``period`` does not affect how data is smoothed by each writer.163 """164 165 def __init__(self, writers, period=20):166 """167 Args:168 writers (list[EventWriter]): a list of EventWriter objects169 period (int):170 """171 self._writers = writers172 for w in writers:173 assert isinstance(w, EventWriter), w174 self._period = period175 176 def after_step(self):177 if (self.trainer.iter + 1) % self._period == 0 or (178 self.trainer.iter == self.trainer.max_iter - 1179 ):180 for writer in self._writers:181 writer.write()182 183 def after_train(self):184 for writer in self._writers:185 # If any new data is found (e.g. produced by other after_train),186 # write them before closing187 writer.write()188 writer.close()189 190 191class PeriodicCheckpointer(_PeriodicCheckpointer, HookBase):192 """193 Same as :class:`detectron2.checkpoint.PeriodicCheckpointer`, but as a hook.194 195 Note that when used as a hook,196 it is unable to save additional data other than what's defined197 by the given `checkpointer`.198 199 It is executed every ``period`` iterations and after the last iteration.200 """201 202 def before_train(self):203 self.max_iter = self.trainer.max_iter204 205 def after_step(self):206 # No way to use **kwargs207 self.step(self.trainer.iter)208 209 210class BestCheckpointer(HookBase):211 """212 Checkpoints best weights based off given metric.213 214 This hook should be used in conjunction to and executed after the hook215 that produces the metric, e.g. `EvalHook`.216 """217 218 def __init__(219 self,220 eval_period: int,221 checkpointer: Checkpointer,222 val_metric: str,223 mode: str = "max",224 file_prefix: str = "model_best",225 ) -> None:226 """227 Args:228 eval_period (int): the period `EvalHook` is set to run.229 checkpointer: the checkpointer object used to save checkpoints.230 val_metric (str): validation metric to track for best checkpoint, e.g. "bbox/AP50"231 mode (str): one of {'max', 'min'}. controls whether the chosen val metric should be232 maximized or minimized, e.g. for "bbox/AP50" it should be "max"233 file_prefix (str): the prefix of checkpoint's filename, defaults to "model_best"234 """235 self._logger = logging.getLogger(__name__)236 self._period = eval_period237 self._val_metric = val_metric238 assert mode in [239 "max",240 "min",241 ], f'Mode "{mode}" to `BestCheckpointer` is unknown. It should be one of {"max", "min"}.'242 if mode == "max":243 self._compare = operator.gt244 else:245 self._compare = operator.lt246 self._checkpointer = checkpointer247 self._file_prefix = file_prefix248 self.best_metric = None249 self.best_iter = None250 251 def _update_best(self, val, iteration):252 if math.isnan(val) or math.isinf(val):253 return False254 self.best_metric = val255 self.best_iter = iteration256 return True257 258 def _best_checking(self):259 metric_tuple = self.trainer.storage.latest().get(self._val_metric)260 if metric_tuple is None:261 self._logger.warning(262 f"Given val metric {self._val_metric} does not seem to be computed/stored."263 "Will not be checkpointing based on it."264 )265 return266 else:267 latest_metric, metric_iter = metric_tuple268 269 if self.best_metric is None:270 if self._update_best(latest_metric, metric_iter):271 additional_state = {"iteration": metric_iter}272 self._checkpointer.save(f"{self._file_prefix}", **additional_state)273 self._logger.info(274 f"Saved first model at {self.best_metric:0.5f} @ {self.best_iter} steps"275 )276 elif self._compare(latest_metric, self.best_metric):277 additional_state = {"iteration": metric_iter}278 self._checkpointer.save(f"{self._file_prefix}", **additional_state)279 self._logger.info(280 f"Saved best model as latest eval score for {self._val_metric} is "281 f"{latest_metric:0.5f}, better than last best score "282 f"{self.best_metric:0.5f} @ iteration {self.best_iter}."283 )284 self._update_best(latest_metric, metric_iter)285 else:286 self._logger.info(287 f"Not saving as latest eval score for {self._val_metric} is {latest_metric:0.5f}, "288 f"not better than best score {self.best_metric:0.5f} @ iteration {self.best_iter}."289 )290 291 def after_step(self):292 # same conditions as `EvalHook`293 next_iter = self.trainer.iter + 1294 if (295 self._period > 0296 and next_iter % self._period == 0297 and next_iter != self.trainer.max_iter298 ):299 self._best_checking()300 301 def after_train(self):302 # same conditions as `EvalHook`303 if self.trainer.iter + 1 >= self.trainer.max_iter:304 self._best_checking()305 306 307class LRScheduler(HookBase):308 """309 A hook which executes a torch builtin LR scheduler and summarizes the LR.310 It is executed after every iteration.311 """312 313 def __init__(self, optimizer=None, scheduler=None):314 """315 Args:316 optimizer (torch.optim.Optimizer):317 scheduler (torch.optim.LRScheduler or fvcore.common.param_scheduler.ParamScheduler):318 if a :class:`ParamScheduler` object, it defines the multiplier over the base LR319 in the optimizer.320 321 If any argument is not given, will try to obtain it from the trainer.322 """323 self._optimizer = optimizer324 self._scheduler = scheduler325 326 def before_train(self):327 self._optimizer = self._optimizer or self.trainer.optimizer328 if isinstance(self.scheduler, ParamScheduler):329 self._scheduler = LRMultiplier(330 self._optimizer,331 self.scheduler,332 self.trainer.max_iter,333 last_iter=self.trainer.iter - 1,334 )335 self._best_param_group_id = LRScheduler.get_best_param_group_id(self._optimizer)336 337 @staticmethod338 def get_best_param_group_id(optimizer):339 # NOTE: some heuristics on what LR to summarize340 # summarize the param group with most parameters341 largest_group = max(len(g["params"]) for g in optimizer.param_groups)342 343 if largest_group == 1:344 # If all groups have one parameter,345 # then find the most common initial LR, and use it for summary346 lr_count = Counter([g["lr"] for g in optimizer.param_groups])347 lr = lr_count.most_common()[0][0]348 for i, g in enumerate(optimizer.param_groups):349 if g["lr"] == lr:350 return i351 else:352 for i, g in enumerate(optimizer.param_groups):353 if len(g["params"]) == largest_group:354 return i355 356 def after_step(self):357 lr = self._optimizer.param_groups[self._best_param_group_id]["lr"]358 self.trainer.storage.put_scalar("lr", lr, smoothing_hint=False)359 self.scheduler.step()360 361 @property362 def scheduler(self):363 return self._scheduler or self.trainer.scheduler364 365 def state_dict(self):366 if isinstance(self.scheduler, _LRScheduler):367 return self.scheduler.state_dict()368 return {}369 370 def load_state_dict(self, state_dict):371 if isinstance(self.scheduler, _LRScheduler):372 logger = logging.getLogger(__name__)373 logger.info("Loading scheduler from state_dict ...")374 self.scheduler.load_state_dict(state_dict)375 376 377class TorchProfiler(HookBase):378 """379 A hook which runs `torch.profiler.profile`.380 381 Examples:382 ::383 hooks.TorchProfiler(384 lambda trainer: 10 < trainer.iter < 20, self.cfg.OUTPUT_DIR385 )386 387 The above example will run the profiler for iteration 10~20 and dump388 results to ``OUTPUT_DIR``. We did not profile the first few iterations389 because they are typically slower than the rest.390 The result files can be loaded in the ``chrome://tracing`` page in chrome browser,391 and the tensorboard visualizations can be visualized using392 ``tensorboard --logdir OUTPUT_DIR/log``393 """394 395 def __init__(self, enable_predicate, output_dir, *, activities=None, save_tensorboard=True):396 """397 Args:398 enable_predicate (callable[trainer -> bool]): a function which takes a trainer,399 and returns whether to enable the profiler.400 It will be called once every step, and can be used to select which steps to profile.401 output_dir (str): the output directory to dump tracing files.402 activities (iterable): same as in `torch.profiler.profile`.403 save_tensorboard (bool): whether to save tensorboard visualizations at (output_dir)/log/404 """405 self._enable_predicate = enable_predicate406 self._activities = activities407 self._output_dir = output_dir408 self._save_tensorboard = save_tensorboard409 410 def before_step(self):411 if self._enable_predicate(self.trainer):412 if self._save_tensorboard:413 on_trace_ready = torch.profiler.tensorboard_trace_handler(414 os.path.join(415 self._output_dir,416 "log",417 "profiler-tensorboard-iter{}".format(self.trainer.iter),418 ),419 f"worker{comm.get_rank()}",420 )421 else:422 on_trace_ready = None423 self._profiler = torch.profiler.profile(424 activities=self._activities,425 on_trace_ready=on_trace_ready,426 record_shapes=True,427 profile_memory=True,428 with_stack=True,429 with_flops=True,430 )431 self._profiler.__enter__()432 else:433 self._profiler = None434 435 def after_step(self):436 if self._profiler is None:437 return438 self._profiler.__exit__(None, None, None)439 if not self._save_tensorboard:440 PathManager.mkdirs(self._output_dir)441 out_file = os.path.join(442 self._output_dir, "profiler-trace-iter{}.json".format(self.trainer.iter)443 )444 if "://" not in out_file:445 self._profiler.export_chrome_trace(out_file)446 else:447 # Support non-posix filesystems448 with tempfile.TemporaryDirectory(prefix="detectron2_profiler") as d:449 tmp_file = os.path.join(d, "tmp.json")450 self._profiler.export_chrome_trace(tmp_file)451 with open(tmp_file) as f:452 content = f.read()453 with PathManager.open(out_file, "w") as f:454 f.write(content)455 456 457class AutogradProfiler(TorchProfiler):458 """459 A hook which runs `torch.autograd.profiler.profile`.460 461 Examples:462 ::463 hooks.AutogradProfiler(464 lambda trainer: 10 < trainer.iter < 20, self.cfg.OUTPUT_DIR465 )466 467 The above example will run the profiler for iteration 10~20 and dump468 results to ``OUTPUT_DIR``. We did not profile the first few iterations469 because they are typically slower than the rest.470 The result files can be loaded in the ``chrome://tracing`` page in chrome browser.471 472 Note:473 When used together with NCCL on older version of GPUs,474 autograd profiler may cause deadlock because it unnecessarily allocates475 memory on every device it sees. The memory management calls, if476 interleaved with NCCL calls, lead to deadlock on GPUs that do not477 support ``cudaLaunchCooperativeKernelMultiDevice``.478 """479 480 def __init__(self, enable_predicate, output_dir, *, use_cuda=True):481 """482 Args:483 enable_predicate (callable[trainer -> bool]): a function which takes a trainer,484 and returns whether to enable the profiler.485 It will be called once every step, and can be used to select which steps to profile.486 output_dir (str): the output directory to dump tracing files.487 use_cuda (bool): same as in `torch.autograd.profiler.profile`.488 """489 warnings.warn("AutogradProfiler has been deprecated in favor of TorchProfiler.")490 self._enable_predicate = enable_predicate491 self._use_cuda = use_cuda492 self._output_dir = output_dir493 494 def before_step(self):495 if self._enable_predicate(self.trainer):496 self._profiler = torch.autograd.profiler.profile(use_cuda=self._use_cuda)497 self._profiler.__enter__()498 else:499 self._profiler = None500 501 502class EvalHook(HookBase):503 """504 Run an evaluation function periodically, and at the end of training.505 506 It is executed every ``eval_period`` iterations and after the last iteration.507 """508 509 def __init__(self, eval_period, eval_function, eval_after_train=True):510 """511 Args:512 eval_period (int): the period to run `eval_function`. Set to 0 to513 not evaluate periodically (but still evaluate after the last iteration514 if `eval_after_train` is True).515 eval_function (callable): a function which takes no arguments, and516 returns a nested dict of evaluation metrics.517 eval_after_train (bool): whether to evaluate after the last iteration518 519 Note:520 This hook must be enabled in all or none workers.521 If you would like only certain workers to perform evaluation,522 give other workers a no-op function (`eval_function=lambda: None`).523 """524 self._period = eval_period525 self._func = eval_function526 self._eval_after_train = eval_after_train527 528 def _do_eval(self):529 results = self._func()530 531 if results:532 assert isinstance(533 results, dict534 ), "Eval function must return a dict. Got {} instead.".format(results)535 536 flattened_results = flatten_results_dict(results)537 for k, v in flattened_results.items():538 try:539 v = float(v)540 except Exception as e:541 raise ValueError(542 "[EvalHook] eval_function should return a nested dict of float. "543 "Got '{}: {}' instead.".format(k, v)544 ) from e545 self.trainer.storage.put_scalars(**flattened_results, smoothing_hint=False)546 547 # Evaluation may take different time among workers.548 # A barrier make them start the next iteration together.549 comm.synchronize()550 551 def after_step(self):552 next_iter = self.trainer.iter + 1553 if self._period > 0 and next_iter % self._period == 0:554 # do the last eval in after_train555 if next_iter != self.trainer.max_iter:556 self._do_eval()557 558 def after_train(self):559 # This condition is to prevent the eval from running after a failed training560 if self._eval_after_train and self.trainer.iter + 1 >= self.trainer.max_iter:561 self._do_eval()562 # func is likely a closure that holds reference to the trainer563 # therefore we clean it to avoid circular reference in the end564 del self._func565 566 567class PreciseBN(HookBase):568 """569 The standard implementation of BatchNorm uses EMA in inference, which is570 sometimes suboptimal.571 This class computes the true average of statistics rather than the moving average,572 and put true averages to every BN layer in the given model.573 574 It is executed every ``period`` iterations and after the last iteration.575 """576 577 def __init__(self, period, model, data_loader, num_iter):578 """579 Args:580 period (int): the period this hook is run, or 0 to not run during training.581 The hook will always run in the end of training.582 model (nn.Module): a module whose all BN layers in training mode will be583 updated by precise BN.584 Note that user is responsible for ensuring the BN layers to be585 updated are in training mode when this hook is triggered.586 data_loader (iterable): it will produce data to be run by `model(data)`.587 num_iter (int): number of iterations used to compute the precise588 statistics.589 """590 self._logger = logging.getLogger(__name__)591 if len(get_bn_modules(model)) == 0:592 self._logger.info(593 "PreciseBN is disabled because model does not contain BN layers in training mode."594 )595 self._disabled = True596 return597 598 self._model = model599 self._data_loader = data_loader600 self._num_iter = num_iter601 self._period = period602 self._disabled = False603 604 self._data_iter = None605 606 def after_step(self):607 next_iter = self.trainer.iter + 1608 is_final = next_iter == self.trainer.max_iter609 if is_final or (self._period > 0 and next_iter % self._period == 0):610 self.update_stats()611 612 def update_stats(self):613 """614 Update the model with precise statistics. Users can manually call this method.615 """616 if self._disabled:617 return618 619 if self._data_iter is None:620 self._data_iter = iter(self._data_loader)621 622 def data_loader():623 for num_iter in itertools.count(1):624 if num_iter % 100 == 0:625 self._logger.info(626 "Running precise-BN ... {}/{} iterations.".format(num_iter, self._num_iter)627 )628 # This way we can reuse the same iterator629 yield next(self._data_iter)630 631 with EventStorage(): # capture events in a new storage to discard them632 self._logger.info(633 "Running precise-BN for {} iterations... ".format(self._num_iter)634 + "Note that this could produce different statistics every time."635 )636 update_bn_stats(self._model, data_loader(), self._num_iter)637 638 639class TorchMemoryStats(HookBase):640 """641 Writes pytorch's cuda memory statistics periodically.642 """643 644 def __init__(self, period=20, max_runs=10):645 """646 Args:647 period (int): Output stats each 'period' iterations648 max_runs (int): Stop the logging after 'max_runs'649 """650 651 self._logger = logging.getLogger(__name__)652 self._period = period653 self._max_runs = max_runs654 self._runs = 0655 656 def after_step(self):657 if self._runs > self._max_runs:658 return659 660 if (self.trainer.iter + 1) % self._period == 0 or (661 self.trainer.iter == self.trainer.max_iter - 1662 ):663 if torch.cuda.is_available():664 max_reserved_mb = torch.cuda.max_memory_reserved() / 1024.0 / 1024.0665 reserved_mb = torch.cuda.memory_reserved() / 1024.0 / 1024.0666 max_allocated_mb = torch.cuda.max_memory_allocated() / 1024.0 / 1024.0667 allocated_mb = torch.cuda.memory_allocated() / 1024.0 / 1024.0668 669 self._logger.info(670 (671 " iter: {} "672 " max_reserved_mem: {:.0f}MB "673 " reserved_mem: {:.0f}MB "674 " max_allocated_mem: {:.0f}MB "675 " allocated_mem: {:.0f}MB "676 ).format(677 self.trainer.iter,678 max_reserved_mb,679 reserved_mb,680 max_allocated_mb,681 allocated_mb,682 )683 )684 685 self._runs += 1686 if self._runs == self._max_runs:687 mem_summary = torch.cuda.memory_summary()688 self._logger.info("\n" + mem_summary)689 690 torch.cuda.reset_peak_memory_stats()691 