CoolFace
Apppublic

chendl/compositional_test

sourceHugging Faceupdated 3y agoView on Hugging Face
1likes
run_semantic_segmentation.py521 linesDownload Raw Back to semantic-segmentation
1#!/usr/bin/env python2# coding=utf-83# Copyright 2022 The HuggingFace Inc. team. All rights reserved.4#5# Licensed under the Apache License, Version 2.0 (the "License");6# you may not use this file except in compliance with the License.7# You may obtain a copy of the License at8#9#     http://www.apache.org/licenses/LICENSE-2.010#11# Unless required by applicable law or agreed to in writing, software12# distributed under the License is distributed on an "AS IS" BASIS,13# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.14# See the License for the specific language governing permissions and15 16import json17import logging18import os19import random20import sys21from dataclasses import dataclass, field22from typing import Optional23 24import evaluate25import numpy as np26import torch27from datasets import load_dataset28from huggingface_hub import hf_hub_download29from PIL import Image30from torch import nn31from torchvision import transforms32from torchvision.transforms import functional33 34import transformers35from transformers import (36    AutoConfig,37    AutoImageProcessor,38    AutoModelForSemanticSegmentation,39    HfArgumentParser,40    Trainer,41    TrainingArguments,42    default_data_collator,43)44from transformers.trainer_utils import get_last_checkpoint45from transformers.utils import check_min_version, send_example_telemetry46from transformers.utils.versions import require_version47 48 49""" Finetuning any ๐Ÿค— Transformers model supported by AutoModelForSemanticSegmentation for semantic segmentation leveraging the Trainer API."""50 51logger = logging.getLogger(__name__)52 53# Will error if the minimal version of Transformers is not installed. Remove at your own risks.54check_min_version("4.28.0")55 56require_version("datasets>=2.0.0", "To fix: pip install -r examples/pytorch/semantic-segmentation/requirements.txt")57 58 59def pad_if_smaller(img, size, fill=0):60    size = (size, size) if isinstance(size, int) else size61    original_width, original_height = img.size62    pad_height = size[1] - original_height if original_height < size[1] else 063    pad_width = size[0] - original_width if original_width < size[0] else 064    img = functional.pad(img, (0, 0, pad_width, pad_height), fill=fill)65    return img66 67 68class Compose:69    def __init__(self, transforms):70        self.transforms = transforms71 72    def __call__(self, image, target):73        for t in self.transforms:74            image, target = t(image, target)75        return image, target76 77 78class Identity:79    def __init__(self):80        pass81 82    def __call__(self, image, target):83        return image, target84 85 86class Resize:87    def __init__(self, size):88        self.size = size89 90    def __call__(self, image, target):91        image = functional.resize(image, self.size)92        target = functional.resize(target, self.size, interpolation=transforms.InterpolationMode.NEAREST)93        return image, target94 95 96class RandomResize:97    def __init__(self, min_size, max_size=None):98        self.min_size = min_size99        if max_size is None:100            max_size = min_size101        self.max_size = max_size102 103    def __call__(self, image, target):104        size = random.randint(self.min_size, self.max_size)105        image = functional.resize(image, size)106        target = functional.resize(target, size, interpolation=transforms.InterpolationMode.NEAREST)107        return image, target108 109 110class RandomCrop:111    def __init__(self, size):112        self.size = size if isinstance(size, tuple) else (size, size)113 114    def __call__(self, image, target):115        image = pad_if_smaller(image, self.size)116        target = pad_if_smaller(target, self.size, fill=255)117        crop_params = transforms.RandomCrop.get_params(image, self.size)118        image = functional.crop(image, *crop_params)119        target = functional.crop(target, *crop_params)120        return image, target121 122 123class RandomHorizontalFlip:124    def __init__(self, flip_prob):125        self.flip_prob = flip_prob126 127    def __call__(self, image, target):128        if random.random() < self.flip_prob:129            image = functional.hflip(image)130            target = functional.hflip(target)131        return image, target132 133 134class PILToTensor:135    def __call__(self, image, target):136        image = functional.pil_to_tensor(image)137        target = torch.as_tensor(np.array(target), dtype=torch.int64)138        return image, target139 140 141class ConvertImageDtype:142    def __init__(self, dtype):143        self.dtype = dtype144 145    def __call__(self, image, target):146        image = functional.convert_image_dtype(image, self.dtype)147        return image, target148 149 150class Normalize:151    def __init__(self, mean, std):152        self.mean = mean153        self.std = std154 155    def __call__(self, image, target):156        image = functional.normalize(image, mean=self.mean, std=self.std)157        return image, target158 159 160class ReduceLabels:161    def __call__(self, image, target):162        if not isinstance(target, np.ndarray):163            target = np.array(target).astype(np.uint8)164        # avoid using underflow conversion165        target[target == 0] = 255166        target = target - 1167        target[target == 254] = 255168 169        target = Image.fromarray(target)170        return image, target171 172 173@dataclass174class DataTrainingArguments:175    """176    Arguments pertaining to what data we are going to input our model for training and eval.177    Using `HfArgumentParser` we can turn this class into argparse arguments to be able to specify178    them on the command line.179    """180 181    dataset_name: Optional[str] = field(182        default="segments/sidewalk-semantic",183        metadata={184            "help": "Name of a dataset from the hub (could be your own, possibly private dataset hosted on the hub)."185        },186    )187    dataset_config_name: Optional[str] = field(188        default=None, metadata={"help": "The configuration name of the dataset to use (via the datasets library)."}189    )190    train_val_split: Optional[float] = field(191        default=0.15, metadata={"help": "Percent to split off of train for validation."}192    )193    max_train_samples: Optional[int] = field(194        default=None,195        metadata={196            "help": (197                "For debugging purposes or quicker training, truncate the number of training examples to this "198                "value if set."199            )200        },201    )202    max_eval_samples: Optional[int] = field(203        default=None,204        metadata={205            "help": (206                "For debugging purposes or quicker training, truncate the number of evaluation examples to this "207                "value if set."208            )209        },210    )211    reduce_labels: Optional[bool] = field(212        default=False,213        metadata={"help": "Whether or not to reduce all labels by 1 and replace background by 255."},214    )215 216    def __post_init__(self):217        if self.dataset_name is None and (self.train_dir is None and self.validation_dir is None):218            raise ValueError(219                "You must specify either a dataset name from the hub or a train and/or validation directory."220            )221 222 223@dataclass224class ModelArguments:225    """226    Arguments pertaining to which model/config/tokenizer we are going to fine-tune from.227    """228 229    model_name_or_path: str = field(230        default="nvidia/mit-b0",231        metadata={"help": "Path to pretrained model or model identifier from huggingface.co/models"},232    )233    config_name: Optional[str] = field(234        default=None, metadata={"help": "Pretrained config name or path if not the same as model_name"}235    )236    cache_dir: Optional[str] = field(237        default=None, metadata={"help": "Where do you want to store the pretrained models downloaded from s3"}238    )239    model_revision: str = field(240        default="main",241        metadata={"help": "The specific model version to use (can be a branch name, tag name or commit id)."},242    )243    image_processor_name: str = field(default=None, metadata={"help": "Name or path of preprocessor config."})244    use_auth_token: bool = field(245        default=False,246        metadata={247            "help": (248                "Will use the token generated when running `huggingface-cli login` (necessary to use this script "249                "with private models)."250            )251        },252    )253 254 255def main():256    # See all possible arguments in src/transformers/training_args.py257    # or by passing the --help flag to this script.258    # We now keep distinct sets of args, for a cleaner separation of concerns.259 260    parser = HfArgumentParser((ModelArguments, DataTrainingArguments, TrainingArguments))261    if len(sys.argv) == 2 and sys.argv[1].endswith(".json"):262        # If we pass only one argument to the script and it's the path to a json file,263        # let's parse it to get our arguments.264        model_args, data_args, training_args = parser.parse_json_file(json_file=os.path.abspath(sys.argv[1]))265    else:266        model_args, data_args, training_args = parser.parse_args_into_dataclasses()267 268    # Sending telemetry. Tracking the example usage helps us better allocate resources to maintain them. The269    # information sent is the one passed as arguments along with your Python/PyTorch versions.270    send_example_telemetry("run_semantic_segmentation", model_args, data_args)271 272    # Setup logging273    logging.basicConfig(274        format="%(asctime)s - %(levelname)s - %(name)s - %(message)s",275        datefmt="%m/%d/%Y %H:%M:%S",276        handlers=[logging.StreamHandler(sys.stdout)],277    )278 279    if training_args.should_log:280        # The default of training_args.log_level is passive, so we set log level at info here to have that default.281        transformers.utils.logging.set_verbosity_info()282 283    log_level = training_args.get_process_log_level()284    logger.setLevel(log_level)285    transformers.utils.logging.set_verbosity(log_level)286    transformers.utils.logging.enable_default_handler()287    transformers.utils.logging.enable_explicit_format()288 289    # Log on each process the small summary:290    logger.warning(291        f"Process rank: {training_args.local_rank}, device: {training_args.device}, n_gpu: {training_args.n_gpu}"292        + f"distributed training: {bool(training_args.local_rank != -1)}, 16-bits training: {training_args.fp16}"293    )294    logger.info(f"Training/evaluation parameters {training_args}")295 296    # Detecting last checkpoint.297    last_checkpoint = None298    if os.path.isdir(training_args.output_dir) and training_args.do_train and not training_args.overwrite_output_dir:299        last_checkpoint = get_last_checkpoint(training_args.output_dir)300        if last_checkpoint is None and len(os.listdir(training_args.output_dir)) > 0:301            raise ValueError(302                f"Output directory ({training_args.output_dir}) already exists and is not empty. "303                "Use --overwrite_output_dir to overcome."304            )305        elif last_checkpoint is not None and training_args.resume_from_checkpoint is None:306            logger.info(307                f"Checkpoint detected, resuming training at {last_checkpoint}. To avoid this behavior, change "308                "the `--output_dir` or add `--overwrite_output_dir` to train from scratch."309            )310 311    # Load dataset312    # In distributed training, the load_dataset function guarantees that only one local process can concurrently313    # download the dataset.314    # TODO support datasets from local folders315    dataset = load_dataset(data_args.dataset_name, cache_dir=model_args.cache_dir)316 317    # Rename column names to standardized names (only "image" and "label" need to be present)318    if "pixel_values" in dataset["train"].column_names:319        dataset = dataset.rename_columns({"pixel_values": "image"})320    if "annotation" in dataset["train"].column_names:321        dataset = dataset.rename_columns({"annotation": "label"})322 323    # If we don't have a validation split, split off a percentage of train as validation.324    data_args.train_val_split = None if "validation" in dataset.keys() else data_args.train_val_split325    if isinstance(data_args.train_val_split, float) and data_args.train_val_split > 0.0:326        split = dataset["train"].train_test_split(data_args.train_val_split)327        dataset["train"] = split["train"]328        dataset["validation"] = split["test"]329 330    # Prepare label mappings.331    # We'll include these in the model's config to get human readable labels in the Inference API.332    if data_args.dataset_name == "scene_parse_150":333        repo_id = "huggingface/label-files"334        filename = "ade20k-id2label.json"335    else:336        repo_id = data_args.dataset_name337        filename = "id2label.json"338    id2label = json.load(open(hf_hub_download(repo_id, filename, repo_type="dataset"), "r"))339    id2label = {int(k): v for k, v in id2label.items()}340    label2id = {v: str(k) for k, v in id2label.items()}341 342    # Load the mean IoU metric from the datasets package343    metric = evaluate.load("mean_iou")344 345    # Define our compute_metrics function. It takes an `EvalPrediction` object (a namedtuple with a346    # predictions and label_ids field) and has to return a dictionary string to float.347    @torch.no_grad()348    def compute_metrics(eval_pred):349        logits, labels = eval_pred350        logits_tensor = torch.from_numpy(logits)351        # scale the logits to the size of the label352        logits_tensor = nn.functional.interpolate(353            logits_tensor,354            size=labels.shape[-2:],355            mode="bilinear",356            align_corners=False,357        ).argmax(dim=1)358 359        pred_labels = logits_tensor.detach().cpu().numpy()360        metrics = metric.compute(361            predictions=pred_labels,362            references=labels,363            num_labels=len(id2label),364            ignore_index=0,365            reduce_labels=image_processor.do_reduce_labels,366        )367        # add per category metrics as individual key-value pairs368        per_category_accuracy = metrics.pop("per_category_accuracy").tolist()369        per_category_iou = metrics.pop("per_category_iou").tolist()370 371        metrics.update({f"accuracy_{id2label[i]}": v for i, v in enumerate(per_category_accuracy)})372        metrics.update({f"iou_{id2label[i]}": v for i, v in enumerate(per_category_iou)})373 374        return metrics375 376    config = AutoConfig.from_pretrained(377        model_args.config_name or model_args.model_name_or_path,378        label2id=label2id,379        id2label=id2label,380        cache_dir=model_args.cache_dir,381        revision=model_args.model_revision,382        use_auth_token=True if model_args.use_auth_token else None,383    )384    model = AutoModelForSemanticSegmentation.from_pretrained(385        model_args.model_name_or_path,386        from_tf=bool(".ckpt" in model_args.model_name_or_path),387        config=config,388        cache_dir=model_args.cache_dir,389        revision=model_args.model_revision,390        use_auth_token=True if model_args.use_auth_token else None,391    )392    image_processor = AutoImageProcessor.from_pretrained(393        model_args.image_processor_name or model_args.model_name_or_path,394        cache_dir=model_args.cache_dir,395        revision=model_args.model_revision,396        use_auth_token=True if model_args.use_auth_token else None,397    )398 399    # Define torchvision transforms to be applied to each image + target.400    # Not that straightforward in torchvision: https://github.com/pytorch/vision/issues/9401    # Currently based on official torchvision references: https://github.com/pytorch/vision/blob/main/references/segmentation/transforms.py402    if "shortest_edge" in image_processor.size:403        # We instead set the target size as (shortest_edge, shortest_edge) to here to ensure all images are batchable.404        size = (image_processor.size["shortest_edge"], image_processor.size["shortest_edge"])405    else:406        size = (image_processor.size["height"], image_processor.size["width"])407    train_transforms = Compose(408        [409            ReduceLabels() if data_args.reduce_labels else Identity(),410            RandomCrop(size=size),411            RandomHorizontalFlip(flip_prob=0.5),412            PILToTensor(),413            ConvertImageDtype(torch.float),414            Normalize(mean=image_processor.image_mean, std=image_processor.image_std),415        ]416    )417    # Define torchvision transform to be applied to each image.418    # jitter = ColorJitter(brightness=0.25, contrast=0.25, saturation=0.25, hue=0.1)419    val_transforms = Compose(420        [421            ReduceLabels() if data_args.reduce_labels else Identity(),422            Resize(size=size),423            PILToTensor(),424            ConvertImageDtype(torch.float),425            Normalize(mean=image_processor.image_mean, std=image_processor.image_std),426        ]427    )428 429    def preprocess_train(example_batch):430        pixel_values = []431        labels = []432        for image, target in zip(example_batch["image"], example_batch["label"]):433            image, target = train_transforms(image.convert("RGB"), target)434            pixel_values.append(image)435            labels.append(target)436 437        encoding = {}438        encoding["pixel_values"] = torch.stack(pixel_values)439        encoding["labels"] = torch.stack(labels)440 441        return encoding442 443    def preprocess_val(example_batch):444        pixel_values = []445        labels = []446        for image, target in zip(example_batch["image"], example_batch["label"]):447            image, target = val_transforms(image.convert("RGB"), target)448            pixel_values.append(image)449            labels.append(target)450 451        encoding = {}452        encoding["pixel_values"] = torch.stack(pixel_values)453        encoding["labels"] = torch.stack(labels)454 455        return encoding456 457    if training_args.do_train:458        if "train" not in dataset:459            raise ValueError("--do_train requires a train dataset")460        if data_args.max_train_samples is not None:461            dataset["train"] = (462                dataset["train"].shuffle(seed=training_args.seed).select(range(data_args.max_train_samples))463            )464        # Set the training transforms465        dataset["train"].set_transform(preprocess_train)466 467    if training_args.do_eval:468        if "validation" not in dataset:469            raise ValueError("--do_eval requires a validation dataset")470        if data_args.max_eval_samples is not None:471            dataset["validation"] = (472                dataset["validation"].shuffle(seed=training_args.seed).select(range(data_args.max_eval_samples))473            )474        # Set the validation transforms475        dataset["validation"].set_transform(preprocess_val)476 477    # Initalize our trainer478    trainer = Trainer(479        model=model,480        args=training_args,481        train_dataset=dataset["train"] if training_args.do_train else None,482        eval_dataset=dataset["validation"] if training_args.do_eval else None,483        compute_metrics=compute_metrics,484        tokenizer=image_processor,485        data_collator=default_data_collator,486    )487 488    # Training489    if training_args.do_train:490        checkpoint = None491        if training_args.resume_from_checkpoint is not None:492            checkpoint = training_args.resume_from_checkpoint493        elif last_checkpoint is not None:494            checkpoint = last_checkpoint495        train_result = trainer.train(resume_from_checkpoint=checkpoint)496        trainer.save_model()497        trainer.log_metrics("train", train_result.metrics)498        trainer.save_metrics("train", train_result.metrics)499        trainer.save_state()500 501    # Evaluation502    if training_args.do_eval:503        metrics = trainer.evaluate()504        trainer.log_metrics("eval", metrics)505        trainer.save_metrics("eval", metrics)506 507    # Write model card and (optionally) push to hub508    kwargs = {509        "finetuned_from": model_args.model_name_or_path,510        "dataset": data_args.dataset_name,511        "tags": ["image-segmentation", "vision"],512    }513    if training_args.push_to_hub:514        trainer.push_to_hub(**kwargs)515    else:516        trainer.create_model_card(**kwargs)517 518 519if __name__ == "__main__":520    main()521