Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model
0
1# -*- coding: utf-8 -*-2# Copyright (c) Facebook, Inc. and its affiliates.3import concurrent.futures4import logging5import numpy as np6import time7import weakref8from typing import List, Mapping, Optional9import torch10from torch.nn.parallel import DataParallel, DistributedDataParallel11 12import detectron2.utils.comm as comm13from detectron2.utils.events import EventStorage, get_event_storage14from detectron2.utils.logger import _log_api_usage15 16__all__ = ["HookBase", "TrainerBase", "SimpleTrainer", "AMPTrainer"]17 18 19class HookBase:20 """21 Base class for hooks that can be registered with :class:`TrainerBase`.22 23 Each hook can implement 4 methods. The way they are called is demonstrated24 in the following snippet:25 ::26 hook.before_train()27 for iter in range(start_iter, max_iter):28 hook.before_step()29 trainer.run_step()30 hook.after_step()31 iter += 132 hook.after_train()33 34 Notes:35 1. In the hook method, users can access ``self.trainer`` to access more36 properties about the context (e.g., model, current iteration, or config37 if using :class:`DefaultTrainer`).38 39 2. A hook that does something in :meth:`before_step` can often be40 implemented equivalently in :meth:`after_step`.41 If the hook takes non-trivial time, it is strongly recommended to42 implement the hook in :meth:`after_step` instead of :meth:`before_step`.43 The convention is that :meth:`before_step` should only take negligible time.44 45 Following this convention will allow hooks that do care about the difference46 between :meth:`before_step` and :meth:`after_step` (e.g., timer) to47 function properly.48 49 """50 51 trainer: "TrainerBase" = None52 """53 A weak reference to the trainer object. Set by the trainer when the hook is registered.54 """55 56 def before_train(self):57 """58 Called before the first iteration.59 """60 pass61 62 def after_train(self):63 """64 Called after the last iteration.65 """66 pass67 68 def before_step(self):69 """70 Called before each iteration.71 """72 pass73 74 def after_backward(self):75 """76 Called after the backward pass of each iteration.77 """78 pass79 80 def after_step(self):81 """82 Called after each iteration.83 """84 pass85 86 def state_dict(self):87 """88 Hooks are stateless by default, but can be made checkpointable by89 implementing `state_dict` and `load_state_dict`.90 """91 return {}92 93 94class TrainerBase:95 """96 Base class for iterative trainer with hooks.97 98 The only assumption we made here is: the training runs in a loop.99 A subclass can implement what the loop is.100 We made no assumptions about the existence of dataloader, optimizer, model, etc.101 102 Attributes:103 iter(int): the current iteration.104 105 start_iter(int): The iteration to start with.106 By convention the minimum possible value is 0.107 108 max_iter(int): The iteration to end training.109 110 storage(EventStorage): An EventStorage that's opened during the course of training.111 """112 113 def __init__(self) -> None:114 self._hooks: List[HookBase] = []115 self.iter: int = 0116 self.start_iter: int = 0117 self.max_iter: int118 self.storage: EventStorage119 _log_api_usage("trainer." + self.__class__.__name__)120 121 def register_hooks(self, hooks: List[Optional[HookBase]]) -> None:122 """123 Register hooks to the trainer. The hooks are executed in the order124 they are registered.125 126 Args:127 hooks (list[Optional[HookBase]]): list of hooks128 """129 hooks = [h for h in hooks if h is not None]130 for h in hooks:131 assert isinstance(h, HookBase)132 # To avoid circular reference, hooks and trainer cannot own each other.133 # This normally does not matter, but will cause memory leak if the134 # involved objects contain __del__:135 # See http://engineering.hearsaysocial.com/2013/06/16/circular-references-in-python/136 h.trainer = weakref.proxy(self)137 self._hooks.extend(hooks)138 139 def train(self, start_iter: int, max_iter: int):140 """141 Args:142 start_iter, max_iter (int): See docs above143 """144 logger = logging.getLogger(__name__)145 logger.info("Starting training from iteration {}".format(start_iter))146 147 self.iter = self.start_iter = start_iter148 self.max_iter = max_iter149 150 with EventStorage(start_iter) as self.storage:151 try:152 self.before_train()153 for self.iter in range(start_iter, max_iter):154 self.before_step()155 self.run_step()156 self.after_step()157 # self.iter == max_iter can be used by `after_train` to158 # tell whether the training successfully finished or failed159 # due to exceptions.160 self.iter += 1161 except Exception:162 logger.exception("Exception during training:")163 raise164 finally:165 self.after_train()166 167 def before_train(self):168 for h in self._hooks:169 h.before_train()170 171 def after_train(self):172 self.storage.iter = self.iter173 for h in self._hooks:174 h.after_train()175 176 def before_step(self):177 # Maintain the invariant that storage.iter == trainer.iter178 # for the entire execution of each step179 self.storage.iter = self.iter180 181 for h in self._hooks:182 h.before_step()183 184 def after_backward(self):185 for h in self._hooks:186 h.after_backward()187 188 def after_step(self):189 for h in self._hooks:190 h.after_step()191 192 def run_step(self):193 raise NotImplementedError194 195 def state_dict(self):196 ret = {"iteration": self.iter}197 hooks_state = {}198 for h in self._hooks:199 sd = h.state_dict()200 if sd:201 name = type(h).__qualname__202 if name in hooks_state:203 # TODO handle repetitive stateful hooks204 continue205 hooks_state[name] = sd206 if hooks_state:207 ret["hooks"] = hooks_state208 return ret209 210 def load_state_dict(self, state_dict):211 logger = logging.getLogger(__name__)212 self.iter = state_dict["iteration"]213 for key, value in state_dict.get("hooks", {}).items():214 for h in self._hooks:215 try:216 name = type(h).__qualname__217 except AttributeError:218 continue219 if name == key:220 h.load_state_dict(value)221 break222 else:223 logger.warning(f"Cannot find the hook '{key}', its state_dict is ignored.")224 225 226class SimpleTrainer(TrainerBase):227 """228 A simple trainer for the most common type of task:229 single-cost single-optimizer single-data-source iterative optimization,230 optionally using data-parallelism.231 It assumes that every step, you:232 233 1. Compute the loss with a data from the data_loader.234 2. Compute the gradients with the above loss.235 3. Update the model with the optimizer.236 237 All other tasks during training (checkpointing, logging, evaluation, LR schedule)238 are maintained by hooks, which can be registered by :meth:`TrainerBase.register_hooks`.239 240 If you want to do anything fancier than this,241 either subclass TrainerBase and implement your own `run_step`,242 or write your own training loop.243 """244 245 def __init__(246 self,247 model,248 data_loader,249 optimizer,250 gather_metric_period=1,251 zero_grad_before_forward=False,252 async_write_metrics=False,253 ):254 """255 Args:256 model: a torch Module. Takes a data from data_loader and returns a257 dict of losses.258 data_loader: an iterable. Contains data to be used to call model.259 optimizer: a torch optimizer.260 gather_metric_period: an int. Every gather_metric_period iterations261 the metrics are gathered from all the ranks to rank 0 and logged.262 zero_grad_before_forward: whether to zero the gradients before the forward.263 async_write_metrics: bool. If True, then write metrics asynchronously to improve264 training speed265 """266 super().__init__()267 268 """269 We set the model to training mode in the trainer.270 However it's valid to train a model that's in eval mode.271 If you want your model (or a submodule of it) to behave272 like evaluation during training, you can overwrite its train() method.273 """274 model.train()275 276 self.model = model277 self.data_loader = data_loader278 # to access the data loader iterator, call `self._data_loader_iter`279 self._data_loader_iter_obj = None280 self.optimizer = optimizer281 self.gather_metric_period = gather_metric_period282 self.zero_grad_before_forward = zero_grad_before_forward283 self.async_write_metrics = async_write_metrics284 # create a thread pool that can execute non critical logic in run_step asynchronically285 # use only 1 worker so tasks will be executred in order of submitting.286 self.concurrent_executor = concurrent.futures.ThreadPoolExecutor(max_workers=1)287 288 def run_step(self):289 """290 Implement the standard training logic described above.291 """292 assert self.model.training, "[SimpleTrainer] model was changed to eval mode!"293 start = time.perf_counter()294 """295 If you want to do something with the data, you can wrap the dataloader.296 """297 data = next(self._data_loader_iter)298 data_time = time.perf_counter() - start299 300 if self.zero_grad_before_forward:301 """302 If you need to accumulate gradients or do something similar, you can303 wrap the optimizer with your custom `zero_grad()` method.304 """305 self.optimizer.zero_grad()306 307 """308 If you want to do something with the losses, you can wrap the model.309 """310 loss_dict = self.model(data)311 if isinstance(loss_dict, torch.Tensor):312 losses = loss_dict313 loss_dict = {"total_loss": loss_dict}314 else:315 losses = sum(loss_dict.values())316 if not self.zero_grad_before_forward:317 """318 If you need to accumulate gradients or do something similar, you can319 wrap the optimizer with your custom `zero_grad()` method.320 """321 self.optimizer.zero_grad()322 losses.backward()323 324 self.after_backward()325 326 if self.async_write_metrics:327 # write metrics asynchronically328 self.concurrent_executor.submit(329 self._write_metrics, loss_dict, data_time, iter=self.iter330 )331 else:332 self._write_metrics(loss_dict, data_time)333 334 """335 If you need gradient clipping/scaling or other processing, you can336 wrap the optimizer with your custom `step()` method. But it is337 suboptimal as explained in https://arxiv.org/abs/2006.15704 Sec 3.2.4338 """339 self.optimizer.step()340 341 @property342 def _data_loader_iter(self):343 # only create the data loader iterator when it is used344 if self._data_loader_iter_obj is None:345 self._data_loader_iter_obj = iter(self.data_loader)346 return self._data_loader_iter_obj347 348 def reset_data_loader(self, data_loader_builder):349 """350 Delete and replace the current data loader with a new one, which will be created351 by calling `data_loader_builder` (without argument).352 """353 del self.data_loader354 data_loader = data_loader_builder()355 self.data_loader = data_loader356 self._data_loader_iter_obj = None357 358 def _write_metrics(359 self,360 loss_dict: Mapping[str, torch.Tensor],361 data_time: float,362 prefix: str = "",363 iter: Optional[int] = None,364 ) -> None:365 logger = logging.getLogger(__name__)366 367 iter = self.iter if iter is None else iter368 if (iter + 1) % self.gather_metric_period == 0:369 try:370 SimpleTrainer.write_metrics(loss_dict, data_time, iter, prefix)371 except Exception:372 logger.exception("Exception in writing metrics: ")373 raise374 375 @staticmethod376 def write_metrics(377 loss_dict: Mapping[str, torch.Tensor],378 data_time: float,379 cur_iter: int,380 prefix: str = "",381 ) -> None:382 """383 Args:384 loss_dict (dict): dict of scalar losses385 data_time (float): time taken by the dataloader iteration386 prefix (str): prefix for logging keys387 """388 metrics_dict = {k: v.detach().cpu().item() for k, v in loss_dict.items()}389 metrics_dict["data_time"] = data_time390 391 # Gather metrics among all workers for logging392 # This assumes we do DDP-style training, which is currently the only393 # supported method in detectron2.394 all_metrics_dict = comm.gather(metrics_dict)395 396 if comm.is_main_process():397 storage = get_event_storage()398 399 # data_time among workers can have high variance. The actual latency400 # caused by data_time is the maximum among workers.401 data_time = np.max([x.pop("data_time") for x in all_metrics_dict])402 storage.put_scalar("data_time", data_time, cur_iter=cur_iter)403 404 # average the rest metrics405 metrics_dict = {406 k: np.mean([x[k] for x in all_metrics_dict]) for k in all_metrics_dict[0].keys()407 }408 total_losses_reduced = sum(metrics_dict.values())409 if not np.isfinite(total_losses_reduced):410 raise FloatingPointError(411 f"Loss became infinite or NaN at iteration={cur_iter}!\n"412 f"loss_dict = {metrics_dict}"413 )414 415 storage.put_scalar(416 "{}total_loss".format(prefix), total_losses_reduced, cur_iter=cur_iter417 )418 if len(metrics_dict) > 1:419 storage.put_scalars(cur_iter=cur_iter, **metrics_dict)420 421 def state_dict(self):422 ret = super().state_dict()423 ret["optimizer"] = self.optimizer.state_dict()424 return ret425 426 def load_state_dict(self, state_dict):427 super().load_state_dict(state_dict)428 self.optimizer.load_state_dict(state_dict["optimizer"])429 430 def after_train(self):431 super().after_train()432 self.concurrent_executor.shutdown(wait=True)433 434 435class AMPTrainer(SimpleTrainer):436 """437 Like :class:`SimpleTrainer`, but uses PyTorch's native automatic mixed precision438 in the training loop.439 """440 441 def __init__(442 self,443 model,444 data_loader,445 optimizer,446 gather_metric_period=1,447 zero_grad_before_forward=False,448 grad_scaler=None,449 precision: torch.dtype = torch.float16,450 log_grad_scaler: bool = False,451 async_write_metrics=False,452 ):453 """454 Args:455 model, data_loader, optimizer, gather_metric_period, zero_grad_before_forward,456 async_write_metrics: same as in :class:`SimpleTrainer`.457 grad_scaler: torch GradScaler to automatically scale gradients.458 precision: torch.dtype as the target precision to cast to in computations459 """460 unsupported = "AMPTrainer does not support single-process multi-device training!"461 if isinstance(model, DistributedDataParallel):462 assert not (model.device_ids and len(model.device_ids) > 1), unsupported463 assert not isinstance(model, DataParallel), unsupported464 465 super().__init__(466 model, data_loader, optimizer, gather_metric_period, zero_grad_before_forward467 )468 469 if grad_scaler is None:470 from torch.cuda.amp import GradScaler471 472 grad_scaler = GradScaler()473 self.grad_scaler = grad_scaler474 self.precision = precision475 self.log_grad_scaler = log_grad_scaler476 477 def run_step(self):478 """479 Implement the AMP training logic.480 """481 assert self.model.training, "[AMPTrainer] model was changed to eval mode!"482 assert torch.cuda.is_available(), "[AMPTrainer] CUDA is required for AMP training!"483 from torch.cuda.amp import autocast484 485 start = time.perf_counter()486 data = next(self._data_loader_iter)487 data_time = time.perf_counter() - start488 489 if self.zero_grad_before_forward:490 self.optimizer.zero_grad()491 with autocast(dtype=self.precision):492 loss_dict = self.model(data)493 if isinstance(loss_dict, torch.Tensor):494 losses = loss_dict495 loss_dict = {"total_loss": loss_dict}496 else:497 losses = sum(loss_dict.values())498 499 if not self.zero_grad_before_forward:500 self.optimizer.zero_grad()501 502 self.grad_scaler.scale(losses).backward()503 504 if self.log_grad_scaler:505 storage = get_event_storage()506 storage.put_scalar("[metric]grad_scaler", self.grad_scaler.get_scale())507 508 self.after_backward()509 510 if self.async_write_metrics:511 # write metrics asynchronically512 self.concurrent_executor.submit(513 self._write_metrics, loss_dict, data_time, iter=self.iter514 )515 else:516 self._write_metrics(loss_dict, data_time)517 518 self.grad_scaler.step(self.optimizer)519 self.grad_scaler.update()520 521 def state_dict(self):522 ret = super().state_dict()523 ret["grad_scaler"] = self.grad_scaler.state_dict()524 return ret525 526 def load_state_dict(self, state_dict):527 super().load_state_dict(state_dict)528 self.grad_scaler.load_state_dict(state_dict["grad_scaler"])529 