Arulkumar03/Fox_Sheep_Detector_Computer_Vision_model
0
1# Copyright (c) Facebook, Inc. and its affiliates.2import itertools3import logging4import numpy as np5import operator6import pickle7from typing import Any, Callable, Dict, List, Optional, Union8import torch9import torch.utils.data as torchdata10from tabulate import tabulate11from termcolor import colored12 13from detectron2.config import configurable14from detectron2.structures import BoxMode15from detectron2.utils.comm import get_world_size16from detectron2.utils.env import seed_all_rng17from detectron2.utils.file_io import PathManager18from detectron2.utils.logger import _log_api_usage, log_first_n19 20from .catalog import DatasetCatalog, MetadataCatalog21from .common import AspectRatioGroupedDataset, DatasetFromList, MapDataset, ToIterableDataset22from .dataset_mapper import DatasetMapper23from .detection_utils import check_metadata_consistency24from .samplers import (25 InferenceSampler,26 RandomSubsetTrainingSampler,27 RepeatFactorTrainingSampler,28 TrainingSampler,29)30 31"""32This file contains the default logic to build a dataloader for training or testing.33"""34 35__all__ = [36 "build_batch_data_loader",37 "build_detection_train_loader",38 "build_detection_test_loader",39 "get_detection_dataset_dicts",40 "load_proposals_into_dataset",41 "print_instances_class_histogram",42]43 44 45def filter_images_with_only_crowd_annotations(dataset_dicts):46 """47 Filter out images with none annotations or only crowd annotations48 (i.e., images without non-crowd annotations).49 A common training-time preprocessing on COCO dataset.50 51 Args:52 dataset_dicts (list[dict]): annotations in Detectron2 Dataset format.53 54 Returns:55 list[dict]: the same format, but filtered.56 """57 num_before = len(dataset_dicts)58 59 def valid(anns):60 for ann in anns:61 if ann.get("iscrowd", 0) == 0:62 return True63 return False64 65 dataset_dicts = [x for x in dataset_dicts if valid(x["annotations"])]66 num_after = len(dataset_dicts)67 logger = logging.getLogger(__name__)68 logger.info(69 "Removed {} images with no usable annotations. {} images left.".format(70 num_before - num_after, num_after71 )72 )73 return dataset_dicts74 75 76def filter_images_with_few_keypoints(dataset_dicts, min_keypoints_per_image):77 """78 Filter out images with too few number of keypoints.79 80 Args:81 dataset_dicts (list[dict]): annotations in Detectron2 Dataset format.82 83 Returns:84 list[dict]: the same format as dataset_dicts, but filtered.85 """86 num_before = len(dataset_dicts)87 88 def visible_keypoints_in_image(dic):89 # Each keypoints field has the format [x1, y1, v1, ...], where v is visibility90 annotations = dic["annotations"]91 return sum(92 (np.array(ann["keypoints"][2::3]) > 0).sum()93 for ann in annotations94 if "keypoints" in ann95 )96 97 dataset_dicts = [98 x for x in dataset_dicts if visible_keypoints_in_image(x) >= min_keypoints_per_image99 ]100 num_after = len(dataset_dicts)101 logger = logging.getLogger(__name__)102 logger.info(103 "Removed {} images with fewer than {} keypoints.".format(104 num_before - num_after, min_keypoints_per_image105 )106 )107 return dataset_dicts108 109 110def load_proposals_into_dataset(dataset_dicts, proposal_file):111 """112 Load precomputed object proposals into the dataset.113 114 The proposal file should be a pickled dict with the following keys:115 116 - "ids": list[int] or list[str], the image ids117 - "boxes": list[np.ndarray], each is an Nx4 array of boxes corresponding to the image id118 - "objectness_logits": list[np.ndarray], each is an N sized array of objectness scores119 corresponding to the boxes.120 - "bbox_mode": the BoxMode of the boxes array. Defaults to ``BoxMode.XYXY_ABS``.121 122 Args:123 dataset_dicts (list[dict]): annotations in Detectron2 Dataset format.124 proposal_file (str): file path of pre-computed proposals, in pkl format.125 126 Returns:127 list[dict]: the same format as dataset_dicts, but added proposal field.128 """129 logger = logging.getLogger(__name__)130 logger.info("Loading proposals from: {}".format(proposal_file))131 132 with PathManager.open(proposal_file, "rb") as f:133 proposals = pickle.load(f, encoding="latin1")134 135 # Rename the key names in D1 proposal files136 rename_keys = {"indexes": "ids", "scores": "objectness_logits"}137 for key in rename_keys:138 if key in proposals:139 proposals[rename_keys[key]] = proposals.pop(key)140 141 # Fetch the indexes of all proposals that are in the dataset142 # Convert image_id to str since they could be int.143 img_ids = set({str(record["image_id"]) for record in dataset_dicts})144 id_to_index = {str(id): i for i, id in enumerate(proposals["ids"]) if str(id) in img_ids}145 146 # Assuming default bbox_mode of precomputed proposals are 'XYXY_ABS'147 bbox_mode = BoxMode(proposals["bbox_mode"]) if "bbox_mode" in proposals else BoxMode.XYXY_ABS148 149 for record in dataset_dicts:150 # Get the index of the proposal151 i = id_to_index[str(record["image_id"])]152 153 boxes = proposals["boxes"][i]154 objectness_logits = proposals["objectness_logits"][i]155 # Sort the proposals in descending order of the scores156 inds = objectness_logits.argsort()[::-1]157 record["proposal_boxes"] = boxes[inds]158 record["proposal_objectness_logits"] = objectness_logits[inds]159 record["proposal_bbox_mode"] = bbox_mode160 161 return dataset_dicts162 163 164def print_instances_class_histogram(dataset_dicts, class_names):165 """166 Args:167 dataset_dicts (list[dict]): list of dataset dicts.168 class_names (list[str]): list of class names (zero-indexed).169 """170 num_classes = len(class_names)171 hist_bins = np.arange(num_classes + 1)172 histogram = np.zeros((num_classes,), dtype=np.int)173 for entry in dataset_dicts:174 annos = entry["annotations"]175 classes = np.asarray(176 [x["category_id"] for x in annos if not x.get("iscrowd", 0)], dtype=np.int177 )178 if len(classes):179 assert classes.min() >= 0, f"Got an invalid category_id={classes.min()}"180 assert (181 classes.max() < num_classes182 ), f"Got an invalid category_id={classes.max()} for a dataset of {num_classes} classes"183 histogram += np.histogram(classes, bins=hist_bins)[0]184 185 N_COLS = min(6, len(class_names) * 2)186 187 def short_name(x):188 # make long class names shorter. useful for lvis189 if len(x) > 13:190 return x[:11] + ".."191 return x192 193 data = list(194 itertools.chain(*[[short_name(class_names[i]), int(v)] for i, v in enumerate(histogram)])195 )196 total_num_instances = sum(data[1::2])197 data.extend([None] * (N_COLS - (len(data) % N_COLS)))198 if num_classes > 1:199 data.extend(["total", total_num_instances])200 data = itertools.zip_longest(*[data[i::N_COLS] for i in range(N_COLS)])201 table = tabulate(202 data,203 headers=["category", "#instances"] * (N_COLS // 2),204 tablefmt="pipe",205 numalign="left",206 stralign="center",207 )208 log_first_n(209 logging.INFO,210 "Distribution of instances among all {} categories:\n".format(num_classes)211 + colored(table, "cyan"),212 key="message",213 )214 215 216def get_detection_dataset_dicts(217 names,218 filter_empty=True,219 min_keypoints=0,220 proposal_files=None,221 check_consistency=True,222):223 """224 Load and prepare dataset dicts for instance detection/segmentation and semantic segmentation.225 226 Args:227 names (str or list[str]): a dataset name or a list of dataset names228 filter_empty (bool): whether to filter out images without instance annotations229 min_keypoints (int): filter out images with fewer keypoints than230 `min_keypoints`. Set to 0 to do nothing.231 proposal_files (list[str]): if given, a list of object proposal files232 that match each dataset in `names`.233 check_consistency (bool): whether to check if datasets have consistent metadata.234 235 Returns:236 list[dict]: a list of dicts following the standard dataset dict format.237 """238 if isinstance(names, str):239 names = [names]240 assert len(names), names241 dataset_dicts = [DatasetCatalog.get(dataset_name) for dataset_name in names]242 243 if isinstance(dataset_dicts[0], torchdata.Dataset):244 if len(dataset_dicts) > 1:245 # ConcatDataset does not work for iterable style dataset.246 # We could support concat for iterable as well, but it's often247 # not a good idea to concat iterables anyway.248 return torchdata.ConcatDataset(dataset_dicts)249 return dataset_dicts[0]250 251 for dataset_name, dicts in zip(names, dataset_dicts):252 assert len(dicts), "Dataset '{}' is empty!".format(dataset_name)253 254 if proposal_files is not None:255 assert len(names) == len(proposal_files)256 # load precomputed proposals from proposal files257 dataset_dicts = [258 load_proposals_into_dataset(dataset_i_dicts, proposal_file)259 for dataset_i_dicts, proposal_file in zip(dataset_dicts, proposal_files)260 ]261 262 dataset_dicts = list(itertools.chain.from_iterable(dataset_dicts))263 264 has_instances = "annotations" in dataset_dicts[0]265 if filter_empty and has_instances:266 dataset_dicts = filter_images_with_only_crowd_annotations(dataset_dicts)267 if min_keypoints > 0 and has_instances:268 dataset_dicts = filter_images_with_few_keypoints(dataset_dicts, min_keypoints)269 270 if check_consistency and has_instances:271 try:272 class_names = MetadataCatalog.get(names[0]).thing_classes273 check_metadata_consistency("thing_classes", names)274 print_instances_class_histogram(dataset_dicts, class_names)275 except AttributeError: # class names are not available for this dataset276 pass277 278 assert len(dataset_dicts), "No valid data found in {}.".format(",".join(names))279 return dataset_dicts280 281 282def build_batch_data_loader(283 dataset,284 sampler,285 total_batch_size,286 *,287 aspect_ratio_grouping=False,288 num_workers=0,289 collate_fn=None,290):291 """292 Build a batched dataloader. The main differences from `torch.utils.data.DataLoader` are:293 1. support aspect ratio grouping options294 2. use no "batch collation", because this is common for detection training295 296 Args:297 dataset (torch.utils.data.Dataset): a pytorch map-style or iterable dataset.298 sampler (torch.utils.data.sampler.Sampler or None): a sampler that produces indices.299 Must be provided iff. ``dataset`` is a map-style dataset.300 total_batch_size, aspect_ratio_grouping, num_workers, collate_fn: see301 :func:`build_detection_train_loader`.302 303 Returns:304 iterable[list]. Length of each list is the batch size of the current305 GPU. Each element in the list comes from the dataset.306 """307 world_size = get_world_size()308 assert (309 total_batch_size > 0 and total_batch_size % world_size == 0310 ), "Total batch size ({}) must be divisible by the number of gpus ({}).".format(311 total_batch_size, world_size312 )313 batch_size = total_batch_size // world_size314 315 if isinstance(dataset, torchdata.IterableDataset):316 assert sampler is None, "sampler must be None if dataset is IterableDataset"317 else:318 dataset = ToIterableDataset(dataset, sampler)319 320 if aspect_ratio_grouping:321 data_loader = torchdata.DataLoader(322 dataset,323 num_workers=num_workers,324 collate_fn=operator.itemgetter(0), # don't batch, but yield individual elements325 worker_init_fn=worker_init_reset_seed,326 ) # yield individual mapped dict327 data_loader = AspectRatioGroupedDataset(data_loader, batch_size)328 if collate_fn is None:329 return data_loader330 return MapDataset(data_loader, collate_fn)331 else:332 return torchdata.DataLoader(333 dataset,334 batch_size=batch_size,335 drop_last=True,336 num_workers=num_workers,337 collate_fn=trivial_batch_collator if collate_fn is None else collate_fn,338 worker_init_fn=worker_init_reset_seed,339 )340 341 342def _train_loader_from_config(cfg, mapper=None, *, dataset=None, sampler=None):343 if dataset is None:344 dataset = get_detection_dataset_dicts(345 cfg.DATASETS.TRAIN,346 filter_empty=cfg.DATALOADER.FILTER_EMPTY_ANNOTATIONS,347 min_keypoints=cfg.MODEL.ROI_KEYPOINT_HEAD.MIN_KEYPOINTS_PER_IMAGE348 if cfg.MODEL.KEYPOINT_ON349 else 0,350 proposal_files=cfg.DATASETS.PROPOSAL_FILES_TRAIN if cfg.MODEL.LOAD_PROPOSALS else None,351 )352 _log_api_usage("dataset." + cfg.DATASETS.TRAIN[0])353 354 if mapper is None:355 mapper = DatasetMapper(cfg, True)356 357 if sampler is None:358 sampler_name = cfg.DATALOADER.SAMPLER_TRAIN359 logger = logging.getLogger(__name__)360 if isinstance(dataset, torchdata.IterableDataset):361 logger.info("Not using any sampler since the dataset is IterableDataset.")362 sampler = None363 else:364 logger.info("Using training sampler {}".format(sampler_name))365 if sampler_name == "TrainingSampler":366 sampler = TrainingSampler(len(dataset))367 elif sampler_name == "RepeatFactorTrainingSampler":368 repeat_factors = RepeatFactorTrainingSampler.repeat_factors_from_category_frequency(369 dataset, cfg.DATALOADER.REPEAT_THRESHOLD370 )371 sampler = RepeatFactorTrainingSampler(repeat_factors)372 elif sampler_name == "RandomSubsetTrainingSampler":373 sampler = RandomSubsetTrainingSampler(374 len(dataset), cfg.DATALOADER.RANDOM_SUBSET_RATIO375 )376 else:377 raise ValueError("Unknown training sampler: {}".format(sampler_name))378 379 return {380 "dataset": dataset,381 "sampler": sampler,382 "mapper": mapper,383 "total_batch_size": cfg.SOLVER.IMS_PER_BATCH,384 "aspect_ratio_grouping": cfg.DATALOADER.ASPECT_RATIO_GROUPING,385 "num_workers": cfg.DATALOADER.NUM_WORKERS,386 }387 388 389@configurable(from_config=_train_loader_from_config)390def build_detection_train_loader(391 dataset,392 *,393 mapper,394 sampler=None,395 total_batch_size,396 aspect_ratio_grouping=True,397 num_workers=0,398 collate_fn=None,399):400 """401 Build a dataloader for object detection with some default features.402 403 Args:404 dataset (list or torch.utils.data.Dataset): a list of dataset dicts,405 or a pytorch dataset (either map-style or iterable). It can be obtained406 by using :func:`DatasetCatalog.get` or :func:`get_detection_dataset_dicts`.407 mapper (callable): a callable which takes a sample (dict) from dataset and408 returns the format to be consumed by the model.409 When using cfg, the default choice is ``DatasetMapper(cfg, is_train=True)``.410 sampler (torch.utils.data.sampler.Sampler or None): a sampler that produces411 indices to be applied on ``dataset``.412 If ``dataset`` is map-style, the default sampler is a :class:`TrainingSampler`,413 which coordinates an infinite random shuffle sequence across all workers.414 Sampler must be None if ``dataset`` is iterable.415 total_batch_size (int): total batch size across all workers.416 aspect_ratio_grouping (bool): whether to group images with similar417 aspect ratio for efficiency. When enabled, it requires each418 element in dataset be a dict with keys "width" and "height".419 num_workers (int): number of parallel data loading workers420 collate_fn: a function that determines how to do batching, same as the argument of421 `torch.utils.data.DataLoader`. Defaults to do no collation and return a list of422 data. No collation is OK for small batch size and simple data structures.423 If your batch size is large and each sample contains too many small tensors,424 it's more efficient to collate them in data loader.425 426 Returns:427 torch.utils.data.DataLoader:428 a dataloader. Each output from it is a ``list[mapped_element]`` of length429 ``total_batch_size / num_workers``, where ``mapped_element`` is produced430 by the ``mapper``.431 """432 if isinstance(dataset, list):433 dataset = DatasetFromList(dataset, copy=False)434 if mapper is not None:435 dataset = MapDataset(dataset, mapper)436 437 if isinstance(dataset, torchdata.IterableDataset):438 assert sampler is None, "sampler must be None if dataset is IterableDataset"439 else:440 if sampler is None:441 sampler = TrainingSampler(len(dataset))442 assert isinstance(sampler, torchdata.Sampler), f"Expect a Sampler but got {type(sampler)}"443 return build_batch_data_loader(444 dataset,445 sampler,446 total_batch_size,447 aspect_ratio_grouping=aspect_ratio_grouping,448 num_workers=num_workers,449 collate_fn=collate_fn,450 )451 452 453def _test_loader_from_config(cfg, dataset_name, mapper=None):454 """455 Uses the given `dataset_name` argument (instead of the names in cfg), because the456 standard practice is to evaluate each test set individually (not combining them).457 """458 if isinstance(dataset_name, str):459 dataset_name = [dataset_name]460 461 dataset = get_detection_dataset_dicts(462 dataset_name,463 filter_empty=False,464 proposal_files=[465 cfg.DATASETS.PROPOSAL_FILES_TEST[list(cfg.DATASETS.TEST).index(x)] for x in dataset_name466 ]467 if cfg.MODEL.LOAD_PROPOSALS468 else None,469 )470 if mapper is None:471 mapper = DatasetMapper(cfg, False)472 return {473 "dataset": dataset,474 "mapper": mapper,475 "num_workers": cfg.DATALOADER.NUM_WORKERS,476 "sampler": InferenceSampler(len(dataset))477 if not isinstance(dataset, torchdata.IterableDataset)478 else None,479 }480 481 482@configurable(from_config=_test_loader_from_config)483def build_detection_test_loader(484 dataset: Union[List[Any], torchdata.Dataset],485 *,486 mapper: Callable[[Dict[str, Any]], Any],487 sampler: Optional[torchdata.Sampler] = None,488 batch_size: int = 1,489 num_workers: int = 0,490 collate_fn: Optional[Callable[[List[Any]], Any]] = None,491) -> torchdata.DataLoader:492 """493 Similar to `build_detection_train_loader`, with default batch size = 1,494 and sampler = :class:`InferenceSampler`. This sampler coordinates all workers495 to produce the exact set of all samples.496 497 Args:498 dataset: a list of dataset dicts,499 or a pytorch dataset (either map-style or iterable). They can be obtained500 by using :func:`DatasetCatalog.get` or :func:`get_detection_dataset_dicts`.501 mapper: a callable which takes a sample (dict) from dataset502 and returns the format to be consumed by the model.503 When using cfg, the default choice is ``DatasetMapper(cfg, is_train=False)``.504 sampler: a sampler that produces505 indices to be applied on ``dataset``. Default to :class:`InferenceSampler`,506 which splits the dataset across all workers. Sampler must be None507 if `dataset` is iterable.508 batch_size: the batch size of the data loader to be created.509 Default to 1 image per worker since this is the standard when reporting510 inference time in papers.511 num_workers: number of parallel data loading workers512 collate_fn: same as the argument of `torch.utils.data.DataLoader`.513 Defaults to do no collation and return a list of data.514 515 Returns:516 DataLoader: a torch DataLoader, that loads the given detection517 dataset, with test-time transformation and batching.518 519 Examples:520 ::521 data_loader = build_detection_test_loader(522 DatasetRegistry.get("my_test"),523 mapper=DatasetMapper(...))524 525 # or, instantiate with a CfgNode:526 data_loader = build_detection_test_loader(cfg, "my_test")527 """528 if isinstance(dataset, list):529 dataset = DatasetFromList(dataset, copy=False)530 if mapper is not None:531 dataset = MapDataset(dataset, mapper)532 if isinstance(dataset, torchdata.IterableDataset):533 assert sampler is None, "sampler must be None if dataset is IterableDataset"534 else:535 if sampler is None:536 sampler = InferenceSampler(len(dataset))537 return torchdata.DataLoader(538 dataset,539 batch_size=batch_size,540 sampler=sampler,541 drop_last=False,542 num_workers=num_workers,543 collate_fn=trivial_batch_collator if collate_fn is None else collate_fn,544 )545 546 547def trivial_batch_collator(batch):548 """549 A batch collator that does nothing.550 """551 return batch552 553 554def worker_init_reset_seed(worker_id):555 initial_seed = torch.initial_seed() % 2**31556 seed_all_rng(initial_seed + worker_id)557 