การจำแนกรูปภาพ

ดูบน TensorFlow.org ทำงานใน Google Colab ดูแหล่งที่มาบน GitHub ดาวน์โหลดโน๊ตบุ๊ค

บทช่วยสอนนี้แสดงวิธีจำแนกรูปภาพดอกไม้ มันสร้างตัวแยกประเภทรูปภาพโดยใช้โมเดล tf.keras.Sequential และโหลดข้อมูลโดยใช้ tf.keras.utils.image_dataset_from_directory คุณจะได้รับประสบการณ์เชิงปฏิบัติด้วยแนวคิดต่อไปนี้:

  • โหลดชุดข้อมูลออกจากดิสก์อย่างมีประสิทธิภาพ
  • ระบุการใช้เทคนิคมากเกินไปเพื่อบรรเทาปัญหา รวมถึงการเสริมข้อมูลและการออกกลางคัน

บทช่วยสอนนี้เป็นไปตามเวิร์กโฟลว์การเรียนรู้ของเครื่องขั้นพื้นฐาน:

  1. ตรวจสอบและทำความเข้าใจข้อมูล
  2. สร้างไปป์ไลน์อินพุต
  3. สร้างแบบจำลอง
  4. ฝึกโมเดล
  5. ทดสอบโมเดล
  6. ปรับปรุงแบบจำลองและทำซ้ำกระบวนการ

นำเข้า TensorFlow และไลบรารีอื่นๆ

import matplotlib.pyplot as plt
import numpy as np
import os
import PIL
import tensorflow as tf

from tensorflow import keras
from tensorflow.keras import layers
from tensorflow.keras.models import Sequential

ดาวน์โหลดและสำรวจชุดข้อมูล

บทช่วยสอนนี้ใช้ชุดข้อมูลภาพถ่ายดอกไม้ประมาณ 3,700 ภาพ ชุดข้อมูลประกอบด้วยไดเร็กทอรีย่อยห้าไดเร็กทอรี หนึ่งไดเร็กทอรีต่อคลาส:

flower_photo/
  daisy/
  dandelion/
  roses/
  sunflowers/
  tulips/
import pathlib
dataset_url = "https://storage.googleapis.com/download.tensorflow.org/example_images/flower_photos.tgz"
data_dir = tf.keras.utils.get_file('flower_photos', origin=dataset_url, untar=True)
data_dir = pathlib.Path(data_dir)

หลังจากดาวน์โหลด คุณควรมีสำเนาของชุดข้อมูลที่พร้อมใช้งาน มีทั้งหมด 3,670 ภาพ:

image_count = len(list(data_dir.glob('*/*.jpg')))
print(image_count)
3670

นี่คือดอกกุหลาบบางส่วน:

roses = list(data_dir.glob('roses/*'))
PIL.Image.open(str(roses[0]))

png

PIL.Image.open(str(roses[1]))

png

และดอกทิวลิปบางส่วน:

tulips = list(data_dir.glob('tulips/*'))
PIL.Image.open(str(tulips[0]))

png

PIL.Image.open(str(tulips[1]))

png

โหลดข้อมูลโดยใช้ยูทิลิตี้ Keras

มาโหลดอิมเมจเหล่านี้จากดิสก์โดยใช้ยูทิลิตี้ tf.keras.utils.image_dataset_from_directory ที่เป็นประโยชน์ ซึ่งจะนำคุณจากไดเร็กทอรีของรูปภาพบนดิสก์ไปยัง tf.data.Dataset ในโค้ดเพียงไม่กี่บรรทัด หากต้องการ คุณยังสามารถเขียนโค้ดการโหลดข้อมูลของคุณเองตั้งแต่ต้นโดยไปที่บทแนะนำการ โหลดและประมวลผลภาพล่วงหน้า

สร้างชุดข้อมูล

กำหนดพารามิเตอร์บางอย่างสำหรับตัวโหลด:

batch_size = 32
img_height = 180
img_width = 180

แนวทางปฏิบัติที่ดีในการใช้การแยกการตรวจสอบความถูกต้องเมื่อพัฒนาแบบจำลองของคุณ ลองใช้รูปภาพ 80% สำหรับการฝึกอบรมและ 20% สำหรับการตรวจสอบ

train_ds = tf.keras.utils.image_dataset_from_directory(
  data_dir,
  validation_split=0.2,
  subset="training",
  seed=123,
  image_size=(img_height, img_width),
  batch_size=batch_size)
Found 3670 files belonging to 5 classes.
Using 2936 files for training.
val_ds = tf.keras.utils.image_dataset_from_directory(
  data_dir,
  validation_split=0.2,
  subset="validation",
  seed=123,
  image_size=(img_height, img_width),
  batch_size=batch_size)
Found 3670 files belonging to 5 classes.
Using 734 files for validation.

คุณสามารถค้นหาชื่อคลาสได้ในแอตทริบิวต์ class_names บนชุดข้อมูลเหล่านี้ ซึ่งสอดคล้องกับชื่อไดเร็กทอรีตามลำดับตัวอักษร

class_names = train_ds.class_names
print(class_names)
['daisy', 'dandelion', 'roses', 'sunflowers', 'tulips']

เห็นภาพข้อมูล

นี่คือภาพเก้าภาพแรกจากชุดข้อมูลการฝึกอบรม:

import matplotlib.pyplot as plt

plt.figure(figsize=(10, 10))
for images, labels in train_ds.take(1):
  for i in range(9):
    ax = plt.subplot(3, 3, i + 1)
    plt.imshow(images[i].numpy().astype("uint8"))
    plt.title(class_names[labels[i]])
    plt.axis("off")

png

คุณจะฝึกโมเดลโดยใช้ชุดข้อมูลเหล่านี้โดยส่งต่อไปยัง Model.fit ในอีกสักครู่ หากต้องการ คุณยังสามารถวนซ้ำชุดข้อมูลด้วยตนเองและเรียกค้นชุดรูปภาพ: