Team Ai
Apppublic

Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
train_loop.py529 linesDownload Raw Back to engine
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