Team Ai
Apppublic

Arulkumar03/Fox_Sheep_Detector_Computer_Vision_model

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