Team Ai
Apppublic

Arulkumar03/Fox_Sheep_Detector_Computer_Vision_model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
build.py557 linesDownload Raw Back to data
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