Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model
0
1# Copyright (c) Facebook, Inc. and its affiliates.2import datetime3import logging4import time5from collections import OrderedDict, abc6from contextlib import ExitStack, contextmanager7from typing import List, Union8import torch9from torch import nn10 11from detectron2.utils.comm import get_world_size, is_main_process12from detectron2.utils.logger import log_every_n_seconds13 14 15class DatasetEvaluator:16 """17 Base class for a dataset evaluator.18 19 The function :func:`inference_on_dataset` runs the model over20 all samples in the dataset, and have a DatasetEvaluator to process the inputs/outputs.21 22 This class will accumulate information of the inputs/outputs (by :meth:`process`),23 and produce evaluation results in the end (by :meth:`evaluate`).24 """25 26 def reset(self):27 """28 Preparation for a new round of evaluation.29 Should be called before starting a round of evaluation.30 """31 pass32 33 def process(self, inputs, outputs):34 """35 Process the pair of inputs and outputs.36 If they contain batches, the pairs can be consumed one-by-one using `zip`:37 38 .. code-block:: python39 40 for input_, output in zip(inputs, outputs):41 # do evaluation on single input/output pair42 ...43 44 Args:45 inputs (list): the inputs that's used to call the model.46 outputs (list): the return value of `model(inputs)`47 """48 pass49 50 def evaluate(self):51 """52 Evaluate/summarize the performance, after processing all input/output pairs.53 54 Returns:55 dict:56 A new evaluator class can return a dict of arbitrary format57 as long as the user can process the results.58 In our train_net.py, we expect the following format:59 60 * key: the name of the task (e.g., bbox)61 * value: a dict of {metric name: score}, e.g.: {"AP50": 80}62 """63 pass64 65 66class DatasetEvaluators(DatasetEvaluator):67 """68 Wrapper class to combine multiple :class:`DatasetEvaluator` instances.69 70 This class dispatches every evaluation call to71 all of its :class:`DatasetEvaluator`.72 """73 74 def __init__(self, evaluators):75 """76 Args:77 evaluators (list): the evaluators to combine.78 """79 super().__init__()80 self._evaluators = evaluators81 82 def reset(self):83 for evaluator in self._evaluators:84 evaluator.reset()85 86 def process(self, inputs, outputs):87 for evaluator in self._evaluators:88 evaluator.process(inputs, outputs)89 90 def evaluate(self):91 results = OrderedDict()92 for evaluator in self._evaluators:93 result = evaluator.evaluate()94 if is_main_process() and result is not None:95 for k, v in result.items():96 assert (97 k not in results98 ), "Different evaluators produce results with the same key {}".format(k)99 results[k] = v100 return results101 102 103def inference_on_dataset(104 model, data_loader, evaluator: Union[DatasetEvaluator, List[DatasetEvaluator], None]105):106 """107 Run model on the data_loader and evaluate the metrics with evaluator.108 Also benchmark the inference speed of `model.__call__` accurately.109 The model will be used in eval mode.110 111 Args:112 model (callable): a callable which takes an object from113 `data_loader` and returns some outputs.114 115 If it's an nn.Module, it will be temporarily set to `eval` mode.116 If you wish to evaluate a model in `training` mode instead, you can117 wrap the given model and override its behavior of `.eval()` and `.train()`.118 data_loader: an iterable object with a length.119 The elements it generates will be the inputs to the model.120 evaluator: the evaluator(s) to run. Use `None` if you only want to benchmark,121 but don't want to do any evaluation.122 123 Returns:124 The return value of `evaluator.evaluate()`125 """126 num_devices = get_world_size()127 logger = logging.getLogger(__name__)128 logger.info("Start inference on {} batches".format(len(data_loader)))129 130 total = len(data_loader) # inference data loader must have a fixed length131 if evaluator is None:132 # create a no-op evaluator133 evaluator = DatasetEvaluators([])134 if isinstance(evaluator, abc.MutableSequence):135 evaluator = DatasetEvaluators(evaluator)136 evaluator.reset()137 138 num_warmup = min(5, total - 1)139 start_time = time.perf_counter()140 total_data_time = 0141 total_compute_time = 0142 total_eval_time = 0143 with ExitStack() as stack:144 if isinstance(model, nn.Module):145 stack.enter_context(inference_context(model))146 stack.enter_context(torch.no_grad())147 148 start_data_time = time.perf_counter()149 for idx, inputs in enumerate(data_loader):150 total_data_time += time.perf_counter() - start_data_time151 if idx == num_warmup:152 start_time = time.perf_counter()153 total_data_time = 0154 total_compute_time = 0155 total_eval_time = 0156 157 start_compute_time = time.perf_counter()158 outputs = model(inputs)159 if torch.cuda.is_available():160 torch.cuda.synchronize()161 total_compute_time += time.perf_counter() - start_compute_time162 163 start_eval_time = time.perf_counter()164 evaluator.process(inputs, outputs)165 total_eval_time += time.perf_counter() - start_eval_time166 167 iters_after_start = idx + 1 - num_warmup * int(idx >= num_warmup)168 data_seconds_per_iter = total_data_time / iters_after_start169 compute_seconds_per_iter = total_compute_time / iters_after_start170 eval_seconds_per_iter = total_eval_time / iters_after_start171 total_seconds_per_iter = (time.perf_counter() - start_time) / iters_after_start172 if idx >= num_warmup * 2 or compute_seconds_per_iter > 5:173 eta = datetime.timedelta(seconds=int(total_seconds_per_iter * (total - idx - 1)))174 log_every_n_seconds(175 logging.INFO,176 (177 f"Inference done {idx + 1}/{total}. "178 f"Dataloading: {data_seconds_per_iter:.4f} s/iter. "179 f"Inference: {compute_seconds_per_iter:.4f} s/iter. "180 f"Eval: {eval_seconds_per_iter:.4f} s/iter. "181 f"Total: {total_seconds_per_iter:.4f} s/iter. "182 f"ETA={eta}"183 ),184 n=5,185 )186 start_data_time = time.perf_counter()187 188 # Measure the time only for this worker (before the synchronization barrier)189 total_time = time.perf_counter() - start_time190 total_time_str = str(datetime.timedelta(seconds=total_time))191 # NOTE this format is parsed by grep192 logger.info(193 "Total inference time: {} ({:.6f} s / iter per device, on {} devices)".format(194 total_time_str, total_time / (total - num_warmup), num_devices195 )196 )197 total_compute_time_str = str(datetime.timedelta(seconds=int(total_compute_time)))198 logger.info(199 "Total inference pure compute time: {} ({:.6f} s / iter per device, on {} devices)".format(200 total_compute_time_str, total_compute_time / (total - num_warmup), num_devices201 )202 )203 204 results = evaluator.evaluate()205 # An evaluator may return None when not in main process.206 # Replace it by an empty dict instead to make it easier for downstream code to handle207 if results is None:208 results = {}209 return results210 211 212@contextmanager213def inference_context(model):214 """215 A context where the model is temporarily changed to eval mode,216 and restored to previous mode afterwards.217 218 Args:219 model: a torch Module220 """221 training_mode = model.training222 model.eval()223 yield224 model.train(training_mode)225 