Team Ai
Apppublic

Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
coco.py540 linesDownload Raw Back to datasets
1# Copyright (c) Facebook, Inc. and its affiliates.2import contextlib3import datetime4import io5import json6import logging7import numpy as np8import os9import shutil10import pycocotools.mask as mask_util11from fvcore.common.timer import Timer12from iopath.common.file_io import file_lock13from PIL import Image14 15from detectron2.structures import Boxes, BoxMode, PolygonMasks, RotatedBoxes16from detectron2.utils.file_io import PathManager17 18from .. import DatasetCatalog, MetadataCatalog19 20"""21This file contains functions to parse COCO-format annotations into dicts in "Detectron2 format".22"""23 24 25logger = logging.getLogger(__name__)26 27__all__ = ["load_coco_json", "load_sem_seg", "convert_to_coco_json", "register_coco_instances"]28 29 30def load_coco_json(json_file, image_root, dataset_name=None, extra_annotation_keys=None):31    """32    Load a json file with COCO's instances annotation format.33    Currently supports instance detection, instance segmentation,34    and person keypoints annotations.35 36    Args:37        json_file (str): full path to the json file in COCO instances annotation format.38        image_root (str or path-like): the directory where the images in this json file exists.39        dataset_name (str or None): the name of the dataset (e.g., coco_2017_train).40            When provided, this function will also do the following:41 42            * Put "thing_classes" into the metadata associated with this dataset.43            * Map the category ids into a contiguous range (needed by standard dataset format),44              and add "thing_dataset_id_to_contiguous_id" to the metadata associated45              with this dataset.46 47            This option should usually be provided, unless users need to load48            the original json content and apply more processing manually.49        extra_annotation_keys (list[str]): list of per-annotation keys that should also be50            loaded into the dataset dict (besides "iscrowd", "bbox", "keypoints",51            "category_id", "segmentation"). The values for these keys will be returned as-is.52            For example, the densepose annotations are loaded in this way.53 54    Returns:55        list[dict]: a list of dicts in Detectron2 standard dataset dicts format (See56        `Using Custom Datasets </tutorials/datasets.html>`_ ) when `dataset_name` is not None.57        If `dataset_name` is None, the returned `category_ids` may be58        incontiguous and may not conform to the Detectron2 standard format.59 60    Notes:61        1. This function does not read the image files.62           The results do not have the "image" field.63    """64    from pycocotools.coco import COCO65 66    timer = Timer()67    json_file = PathManager.get_local_path(json_file)68    with contextlib.redirect_stdout(io.StringIO()):69        coco_api = COCO(json_file)70    if timer.seconds() > 1:71        logger.info("Loading {} takes {:.2f} seconds.".format(json_file, timer.seconds()))72 73    id_map = None74    if dataset_name is not None:75        meta = MetadataCatalog.get(dataset_name)76        cat_ids = sorted(coco_api.getCatIds())77        cats = coco_api.loadCats(cat_ids)78        # The categories in a custom json file may not be sorted.79        thing_classes = [c["name"] for c in sorted(cats, key=lambda x: x["id"])]80        meta.thing_classes = thing_classes81 82        # In COCO, certain category ids are artificially removed,83        # and by convention they are always ignored.84        # We deal with COCO's id issue and translate85        # the category ids to contiguous ids in [0, 80).86 87        # It works by looking at the "categories" field in the json, therefore88        # if users' own json also have incontiguous ids, we'll89        # apply this mapping as well but print a warning.90        if not (min(cat_ids) == 1 and max(cat_ids) == len(cat_ids)):91            if "coco" not in dataset_name:92                logger.warning(93                    """94Category ids in annotations are not in [1, #categories]! We'll apply a mapping for you.95"""96                )97        id_map = {v: i for i, v in enumerate(cat_ids)}98        meta.thing_dataset_id_to_contiguous_id = id_map99 100    # sort indices for reproducible results101    img_ids = sorted(coco_api.imgs.keys())102    # imgs is a list of dicts, each looks something like:103    # {'license': 4,104    #  'url': 'http://farm6.staticflickr.com/5454/9413846304_881d5e5c3b_z.jpg',105    #  'file_name': 'COCO_val2014_000000001268.jpg',106    #  'height': 427,107    #  'width': 640,108    #  'date_captured': '2013-11-17 05:57:24',109    #  'id': 1268}110    imgs = coco_api.loadImgs(img_ids)111    # anns is a list[list[dict]], where each dict is an annotation112    # record for an object. The inner list enumerates the objects in an image113    # and the outer list enumerates over images. Example of anns[0]:114    # [{'segmentation': [[192.81,115    #     247.09,116    #     ...117    #     219.03,118    #     249.06]],119    #   'area': 1035.749,120    #   'iscrowd': 0,121    #   'image_id': 1268,122    #   'bbox': [192.81, 224.8, 74.73, 33.43],123    #   'category_id': 16,124    #   'id': 42986},125    #  ...]126    anns = [coco_api.imgToAnns[img_id] for img_id in img_ids]127    total_num_valid_anns = sum([len(x) for x in anns])128    total_num_anns = len(coco_api.anns)129    if total_num_valid_anns < total_num_anns:130        logger.warning(131            f"{json_file} contains {total_num_anns} annotations, but only "132            f"{total_num_valid_anns} of them match to images in the file."133        )134 135    if "minival" not in json_file:136        # The popular valminusminival & minival annotations for COCO2014 contain this bug.137        # However the ratio of buggy annotations there is tiny and does not affect accuracy.138        # Therefore we explicitly white-list them.139        ann_ids = [ann["id"] for anns_per_image in anns for ann in anns_per_image]140        assert len(set(ann_ids)) == len(ann_ids), "Annotation ids in '{}' are not unique!".format(141            json_file142        )143 144    imgs_anns = list(zip(imgs, anns))145    logger.info("Loaded {} images in COCO format from {}".format(len(imgs_anns), json_file))146 147    dataset_dicts = []148 149    ann_keys = ["iscrowd", "bbox", "keypoints", "category_id"] + (extra_annotation_keys or [])150 151    num_instances_without_valid_segmentation = 0152 153    for (img_dict, anno_dict_list) in imgs_anns:154        record = {}155        record["file_name"] = os.path.join(image_root, img_dict["file_name"])156        record["height"] = img_dict["height"]157        record["width"] = img_dict["width"]158        image_id = record["image_id"] = img_dict["id"]159 160        objs = []161        for anno in anno_dict_list:162            # Check that the image_id in this annotation is the same as163            # the image_id we're looking at.164            # This fails only when the data parsing logic or the annotation file is buggy.165 166            # The original COCO valminusminival2014 & minival2014 annotation files167            # actually contains bugs that, together with certain ways of using COCO API,168            # can trigger this assertion.169            assert anno["image_id"] == image_id170 171            assert anno.get("ignore", 0) == 0, '"ignore" in COCO json file is not supported.'172 173            obj = {key: anno[key] for key in ann_keys if key in anno}174            if "bbox" in obj and len(obj["bbox"]) == 0:175                raise ValueError(176                    f"One annotation of image {image_id} contains empty 'bbox' value! "177                    "This json does not have valid COCO format."178                )179 180            segm = anno.get("segmentation", None)181            if segm:  # either list[list[float]] or dict(RLE)182                if isinstance(segm, dict):183                    if isinstance(segm["counts"], list):184                        # convert to compressed RLE185                        segm = mask_util.frPyObjects(segm, *segm["size"])186                else:187                    # filter out invalid polygons (< 3 points)188                    segm = [poly for poly in segm if len(poly) % 2 == 0 and len(poly) >= 6]189                    if len(segm) == 0:190                        num_instances_without_valid_segmentation += 1191                        continue  # ignore this instance192                obj["segmentation"] = segm193 194            keypts = anno.get("keypoints", None)195            if keypts:  # list[int]196                for idx, v in enumerate(keypts):197                    if idx % 3 != 2:198                        # COCO's segmentation coordinates are floating points in [0, H or W],199                        # but keypoint coordinates are integers in [0, H-1 or W-1]200                        # Therefore we assume the coordinates are "pixel indices" and201                        # add 0.5 to convert to floating point coordinates.202                        keypts[idx] = v + 0.5203                obj["keypoints"] = keypts204 205            obj["bbox_mode"] = BoxMode.XYWH_ABS206            if id_map:207                annotation_category_id = obj["category_id"]208                try:209                    obj["category_id"] = id_map[annotation_category_id]210                except KeyError as e:211                    raise KeyError(212                        f"Encountered category_id={annotation_category_id} "213                        "but this id does not exist in 'categories' of the json file."214                    ) from e215            objs.append(obj)216        record["annotations"] = objs217        dataset_dicts.append(record)218 219    if num_instances_without_valid_segmentation > 0:220        logger.warning(221            "Filtered out {} instances without valid segmentation. ".format(222                num_instances_without_valid_segmentation223            )224            + "There might be issues in your dataset generation process.  Please "225            "check https://detectron2.readthedocs.io/en/latest/tutorials/datasets.html carefully"226        )227    return dataset_dicts228 229 230def load_sem_seg(gt_root, image_root, gt_ext="png", image_ext="jpg"):231    """232    Load semantic segmentation datasets. All files under "gt_root" with "gt_ext" extension are233    treated as ground truth annotations and all files under "image_root" with "image_ext" extension234    as input images. Ground truth and input images are matched using file paths relative to235    "gt_root" and "image_root" respectively without taking into account file extensions.236    This works for COCO as well as some other datasets.237 238    Args:239        gt_root (str): full path to ground truth semantic segmentation files. Semantic segmentation240            annotations are stored as images with integer values in pixels that represent241            corresponding semantic labels.242        image_root (str): the directory where the input images are.243        gt_ext (str): file extension for ground truth annotations.244        image_ext (str): file extension for input images.245 246    Returns:247        list[dict]:248            a list of dicts in detectron2 standard format without instance-level249            annotation.250 251    Notes:252        1. This function does not read the image and ground truth files.253           The results do not have the "image" and "sem_seg" fields.254    """255 256    # We match input images with ground truth based on their relative filepaths (without file257    # extensions) starting from 'image_root' and 'gt_root' respectively.258    def file2id(folder_path, file_path):259        # extract relative path starting from `folder_path`260        image_id = os.path.normpath(os.path.relpath(file_path, start=folder_path))261        # remove file extension262        image_id = os.path.splitext(image_id)[0]263        return image_id264 265    input_files = sorted(266        (os.path.join(image_root, f) for f in PathManager.ls(image_root) if f.endswith(image_ext)),267        key=lambda file_path: file2id(image_root, file_path),268    )269    gt_files = sorted(270        (os.path.join(gt_root, f) for f in PathManager.ls(gt_root) if f.endswith(gt_ext)),271        key=lambda file_path: file2id(gt_root, file_path),272    )273 274    assert len(gt_files) > 0, "No annotations found in {}.".format(gt_root)275 276    # Use the intersection, so that val2017_100 annotations can run smoothly with val2017 images277    if len(input_files) != len(gt_files):278        logger.warn(279            "Directory {} and {} has {} and {} files, respectively.".format(280                image_root, gt_root, len(input_files), len(gt_files)281            )282        )283        input_basenames = [os.path.basename(f)[: -len(image_ext)] for f in input_files]284        gt_basenames = [os.path.basename(f)[: -len(gt_ext)] for f in gt_files]285        intersect = list(set(input_basenames) & set(gt_basenames))286        # sort, otherwise each worker may obtain a list[dict] in different order287        intersect = sorted(intersect)288        logger.warn("Will use their intersection of {} files.".format(len(intersect)))289        input_files = [os.path.join(image_root, f + image_ext) for f in intersect]290        gt_files = [os.path.join(gt_root, f + gt_ext) for f in intersect]291 292    logger.info(293        "Loaded {} images with semantic segmentation from {}".format(len(input_files), image_root)294    )295 296    dataset_dicts = []297    for (img_path, gt_path) in zip(input_files, gt_files):298        record = {}299        record["file_name"] = img_path300        record["sem_seg_file_name"] = gt_path301        dataset_dicts.append(record)302 303    return dataset_dicts304 305 306def convert_to_coco_dict(dataset_name):307    """308    Convert an instance detection/segmentation or keypoint detection dataset309    in detectron2's standard format into COCO json format.310 311    Generic dataset description can be found here:312    https://detectron2.readthedocs.io/tutorials/datasets.html#register-a-dataset313 314    COCO data format description can be found here:315    http://cocodataset.org/#format-data316 317    Args:318        dataset_name (str):319            name of the source dataset320            Must be registered in DatastCatalog and in detectron2's standard format.321            Must have corresponding metadata "thing_classes"322    Returns:323        coco_dict: serializable dict in COCO json format324    """325 326    dataset_dicts = DatasetCatalog.get(dataset_name)327    metadata = MetadataCatalog.get(dataset_name)328 329    # unmap the category mapping ids for COCO330    if hasattr(metadata, "thing_dataset_id_to_contiguous_id"):331        reverse_id_mapping = {v: k for k, v in metadata.thing_dataset_id_to_contiguous_id.items()}332        reverse_id_mapper = lambda contiguous_id: reverse_id_mapping[contiguous_id]  # noqa333    else:334        reverse_id_mapper = lambda contiguous_id: contiguous_id  # noqa335 336    categories = [337        {"id": reverse_id_mapper(id), "name": name}338        for id, name in enumerate(metadata.thing_classes)339    ]340 341    logger.info("Converting dataset dicts into COCO format")342    coco_images = []343    coco_annotations = []344 345    for image_id, image_dict in enumerate(dataset_dicts):346        coco_image = {347            "id": image_dict.get("image_id", image_id),348            "width": int(image_dict["width"]),349            "height": int(image_dict["height"]),350            "file_name": str(image_dict["file_name"]),351        }352        coco_images.append(coco_image)353 354        anns_per_image = image_dict.get("annotations", [])355        for annotation in anns_per_image:356            # create a new dict with only COCO fields357            coco_annotation = {}358 359            # COCO requirement: XYWH box format for axis-align and XYWHA for rotated360            bbox = annotation["bbox"]361            if isinstance(bbox, np.ndarray):362                if bbox.ndim != 1:363                    raise ValueError(f"bbox has to be 1-dimensional. Got shape={bbox.shape}.")364                bbox = bbox.tolist()365            if len(bbox) not in [4, 5]:366                raise ValueError(f"bbox has to has length 4 or 5. Got {bbox}.")367            from_bbox_mode = annotation["bbox_mode"]368            to_bbox_mode = BoxMode.XYWH_ABS if len(bbox) == 4 else BoxMode.XYWHA_ABS369            bbox = BoxMode.convert(bbox, from_bbox_mode, to_bbox_mode)370 371            # COCO requirement: instance area372            if "segmentation" in annotation:373                # Computing areas for instances by counting the pixels374                segmentation = annotation["segmentation"]375                # TODO: check segmentation type: RLE, BinaryMask or Polygon376                if isinstance(segmentation, list):377                    polygons = PolygonMasks([segmentation])378                    area = polygons.area()[0].item()379                elif isinstance(segmentation, dict):  # RLE380                    area = mask_util.area(segmentation).item()381                else:382                    raise TypeError(f"Unknown segmentation type {type(segmentation)}!")383            else:384                # Computing areas using bounding boxes385                if to_bbox_mode == BoxMode.XYWH_ABS:386                    bbox_xy = BoxMode.convert(bbox, to_bbox_mode, BoxMode.XYXY_ABS)387                    area = Boxes([bbox_xy]).area()[0].item()388                else:389                    area = RotatedBoxes([bbox]).area()[0].item()390 391            if "keypoints" in annotation:392                keypoints = annotation["keypoints"]  # list[int]393                for idx, v in enumerate(keypoints):394                    if idx % 3 != 2:395                        # COCO's segmentation coordinates are floating points in [0, H or W],396                        # but keypoint coordinates are integers in [0, H-1 or W-1]397                        # For COCO format consistency we substract 0.5398                        # https://github.com/facebookresearch/detectron2/pull/175#issuecomment-551202163399                        keypoints[idx] = v - 0.5400                if "num_keypoints" in annotation:401                    num_keypoints = annotation["num_keypoints"]402                else:403                    num_keypoints = sum(kp > 0 for kp in keypoints[2::3])404 405            # COCO requirement:406            #   linking annotations to images407            #   "id" field must start with 1408            coco_annotation["id"] = len(coco_annotations) + 1409            coco_annotation["image_id"] = coco_image["id"]410            coco_annotation["bbox"] = [round(float(x), 3) for x in bbox]411            coco_annotation["area"] = float(area)412            coco_annotation["iscrowd"] = int(annotation.get("iscrowd", 0))413            coco_annotation["category_id"] = int(reverse_id_mapper(annotation["category_id"]))414 415            # Add optional fields416            if "keypoints" in annotation:417                coco_annotation["keypoints"] = keypoints418                coco_annotation["num_keypoints"] = num_keypoints419 420            if "segmentation" in annotation:421                seg = coco_annotation["segmentation"] = annotation["segmentation"]422                if isinstance(seg, dict):  # RLE423                    counts = seg["counts"]424                    if not isinstance(counts, str):425                        # make it json-serializable426                        seg["counts"] = counts.decode("ascii")427 428            coco_annotations.append(coco_annotation)429 430    logger.info(431        "Conversion finished, "432        f"#images: {len(coco_images)}, #annotations: {len(coco_annotations)}"433    )434 435    info = {436        "date_created": str(datetime.datetime.now()),437        "description": "Automatically generated COCO json file for Detectron2.",438    }439    coco_dict = {"info": info, "images": coco_images, "categories": categories, "licenses": None}440    if len(coco_annotations) > 0:441        coco_dict["annotations"] = coco_annotations442    return coco_dict443 444 445def convert_to_coco_json(dataset_name, output_file, allow_cached=True):446    """447    Converts dataset into COCO format and saves it to a json file.448    dataset_name must be registered in DatasetCatalog and in detectron2's standard format.449 450    Args:451        dataset_name:452            reference from the config file to the catalogs453            must be registered in DatasetCatalog and in detectron2's standard format454        output_file: path of json file that will be saved to455        allow_cached: if json file is already present then skip conversion456    """457 458    # TODO: The dataset or the conversion script *may* change,459    # a checksum would be useful for validating the cached data460 461    PathManager.mkdirs(os.path.dirname(output_file))462    with file_lock(output_file):463        if PathManager.exists(output_file) and allow_cached:464            logger.warning(465                f"Using previously cached COCO format annotations at '{output_file}'. "466                "You need to clear the cache file if your dataset has been modified."467            )468        else:469            logger.info(f"Converting annotations of dataset '{dataset_name}' to COCO format ...)")470            coco_dict = convert_to_coco_dict(dataset_name)471 472            logger.info(f"Caching COCO format annotations at '{output_file}' ...")473            tmp_file = output_file + ".tmp"474            with PathManager.open(tmp_file, "w") as f:475                json.dump(coco_dict, f)476            shutil.move(tmp_file, output_file)477 478 479def register_coco_instances(name, metadata, json_file, image_root):480    """481    Register a dataset in COCO's json annotation format for482    instance detection, instance segmentation and keypoint detection.483    (i.e., Type 1 and 2 in http://cocodataset.org/#format-data.484    `instances*.json` and `person_keypoints*.json` in the dataset).485 486    This is an example of how to register a new dataset.487    You can do something similar to this function, to register new datasets.488 489    Args:490        name (str): the name that identifies a dataset, e.g. "coco_2014_train".491        metadata (dict): extra metadata associated with this dataset.  You can492            leave it as an empty dict.493        json_file (str): path to the json instance annotation file.494        image_root (str or path-like): directory which contains all the images.495    """496    assert isinstance(name, str), name497    assert isinstance(json_file, (str, os.PathLike)), json_file498    assert isinstance(image_root, (str, os.PathLike)), image_root499    # 1. register a function which returns dicts500    DatasetCatalog.register(name, lambda: load_coco_json(json_file, image_root, name))501 502    # 2. Optionally, add metadata about this dataset,503    # since they might be useful in evaluation, visualization or logging504    MetadataCatalog.get(name).set(505        json_file=json_file, image_root=image_root, evaluator_type="coco", **metadata506    )507 508 509if __name__ == "__main__":510    """511    Test the COCO json dataset loader.512 513    Usage:514        python -m detectron2.data.datasets.coco \515            path/to/json path/to/image_root dataset_name516 517        "dataset_name" can be "coco_2014_minival_100", or other518        pre-registered ones519    """520    from detectron2.utils.logger import setup_logger521    from detectron2.utils.visualizer import Visualizer522    import detectron2.data.datasets  # noqa # add pre-defined metadata523    import sys524 525    logger = setup_logger(name=__name__)526    assert sys.argv[3] in DatasetCatalog.list()527    meta = MetadataCatalog.get(sys.argv[3])528 529    dicts = load_coco_json(sys.argv[1], sys.argv[2], sys.argv[3])530    logger.info("Done loading {} samples.".format(len(dicts)))531 532    dirname = "coco-data-vis"533    os.makedirs(dirname, exist_ok=True)534    for d in dicts:535        img = np.array(Image.open(d["file_name"]))536        visualizer = Visualizer(img, metadata=meta)537        vis = visualizer.draw_dataset_dict(d)538        fpath = os.path.join(dirname, os.path.basename(d["file_name"]))539        vis.save(fpath)540