छवि विभाजन

TensorFlow.org पर देखें Google Colab में चलाएं GitHub पर स्रोत देखें नोटबुक डाउनलोड करें

यह ट्यूटोरियल संशोधित यू-नेट का उपयोग करते हुए छवि विभाजन के कार्य पर केंद्रित है।

छवि विभाजन क्या है?

एक छवि वर्गीकरण कार्य में नेटवर्क प्रत्येक इनपुट छवि को एक लेबल (या वर्ग) प्रदान करता है। हालाँकि, मान लीजिए कि आप उस वस्तु का आकार जानना चाहते हैं, कौन सा पिक्सेल किस वस्तु का है, आदि। इस स्थिति में आप छवि के प्रत्येक पिक्सेल को एक वर्ग निर्दिष्ट करना चाहेंगे। इस कार्य को विभाजन के रूप में जाना जाता है। एक सेगमेंटेशन मॉडल छवि के बारे में अधिक विस्तृत जानकारी देता है। छवि विभाजन में चिकित्सा इमेजिंग, सेल्फ-ड्राइविंग कारों और उपग्रह इमेजिंग में कुछ नाम रखने के लिए कई अनुप्रयोग हैं।

यह ट्यूटोरियल ऑक्सफोर्ड-आईआईआईटी पेट डेटासेट ( पार्खी एट अल, 2012 ) का उपयोग करता है। डेटासेट में 37 पालतू नस्लों की छवियां होती हैं, प्रति नस्ल 200 छवियां (~ 100 प्रत्येक प्रशिक्षण और परीक्षण विभाजन में)। प्रत्येक छवि में संबंधित लेबल और पिक्सेल-वार मास्क शामिल होते हैं। मास्क प्रत्येक पिक्सेल के लिए क्लास-लेबल होते हैं। प्रत्येक पिक्सेल को तीन श्रेणियों में से एक दिया जाता है:

  • कक्षा 1: पालतू जानवर से संबंधित पिक्सेल।
  • कक्षा 2: पालतू जानवर की सीमा पर पिक्सेल।
  • कक्षा 3: उपरोक्त में से कोई नहीं/आसपास का पिक्सेल।
pip install git+https://github.com/tensorflow/examples.git
import tensorflow as tf

import tensorflow_datasets as tfds
from tensorflow_examples.models.pix2pix import pix2pix

from IPython.display import clear_output
import matplotlib.pyplot as plt

ऑक्सफोर्ड-आईआईआईटी पेट्स डेटासेट डाउनलोड करें

डेटासेट TensorFlow Datasets से उपलब्ध है । विभाजन मास्क संस्करण 3+ में शामिल हैं।

dataset, info = tfds.load('oxford_iiit_pet:3.*.*', with_info=True)

इसके अलावा, छवि रंग मान [0,1] श्रेणी में सामान्यीकृत होते हैं। अंत में, जैसा कि ऊपर उल्लेख किया गया है कि सेगमेंटेशन मास्क में पिक्सेल या तो {1, 2, 3} लेबल किए गए हैं। सुविधा के लिए, सेगमेंटेशन मास्क से 1 घटाएं, जिसके परिणामस्वरूप लेबल हैं: {0, 1, 2}।

def normalize(input_image, input_mask):
  input_image = tf.cast(input_image, tf.float32) / 255.0
  input_mask -= 1
  return input_image, input_mask
def load_image(datapoint):
  input_image = tf.image.resize(datapoint['image'], (128, 128))
  input_mask = tf.image.resize(datapoint['segmentation_mask'], (128, 128))

  input_image, input_mask = normalize(input_image, input_mask)

  return input_image, input_mask

डेटासेट में पहले से ही आवश्यक प्रशिक्षण और परीक्षण विभाजन शामिल हैं, इसलिए समान विभाजन का उपयोग करना जारी रखें।

TRAIN_LENGTH = info.splits['train'].num_examples
BATCH_SIZE = 64
BUFFER_SIZE = 1000
STEPS_PER_EPOCH = TRAIN_LENGTH // BATCH_SIZE
train_images = dataset['train'].map(load_image, num_parallel_calls=tf.data.AUTOTUNE)
test_images = dataset['test'].map(load_image, num_parallel_calls=tf.data.AUTOTUNE)

निम्न वर्ग एक छवि को बेतरतीब ढंग से फ़्लिप करके एक साधारण वृद्धि करता है। अधिक जानने के लिए छवि वृद्धि ट्यूटोरियल पर जाएं।

class Augment(tf.keras.layers.Layer):
  def __init__(self, seed=42):
    super().__init__()
    # both use the same seed, so they'll make the same random changes.
    self.augment_inputs = tf.keras.layers.RandomFlip(mode="horizontal", seed=seed)
    self.augment_labels = tf.keras.layers.RandomFlip(mode="horizontal", seed=seed)

  def call(self, inputs, labels):
    inputs = self.augment_inputs(inputs)
    labels = self.augment_labels(labels)
    return inputs, labels

इनपुट की बैचिंग के बाद ऑग्मेंटेशन लागू करते हुए इनपुट पाइपलाइन का निर्माण करें।

train_batches = (
    train_images
    .cache()
    .shuffle(BUFFER_SIZE)
    .batch(BATCH_SIZE)
    .repeat()
    .map(Augment())
    .prefetch(buffer_size=tf.data.AUTOTUNE))

test_batches = test_images.batch(BATCH_SIZE)

डेटासेट से एक छवि उदाहरण और उसके संबंधित मास्क की कल्पना करें।

def display(display_list):
  plt.figure(figsize=(15, 15))

  title = ['Input Image', 'True Mask', 'Predicted Mask']

  for i in range(len(display_list)):
    plt.subplot(1, len(display_list), i+1)
    plt.title(title[i])
    plt.imshow(tf.keras.utils.array_to_img(display_list[i]))
    plt.axis('off')
  plt.show()
for images, masks in train_batches.take(2):
  sample_image, sample_mask = images[0], masks[0]
  display([sample_image, sample_mask])
Corrupt JPEG data: 240 extraneous bytes before marker 0xd9
Corrupt JPEG data: premature end of data segment

पीएनजी