chendl/compositional_test
1
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 