Предварительно обученные модели и трансферное обучение
- 10 мин
Обучение CNN может занять значительное количество времени, и для этого требуется большой объем данных. Большая часть времени тратится на эксперимент, чтобы найти лучшие низкоуровневые фильтры, необходимые сети для извлечения шаблонов из изображений. Возникает естественный вопрос. Можно ли использовать нейронную сеть, обученную на одном наборе данных, и адаптировать ее к классификации различных изображений без полного процесса обучения?
Этот подход называется передачей обучения, так как мы переносим некоторые знания из одной модели нейронной сети в другую. В трансферном обучении мы обычно начинаем с предварительно обученной модели, которая была натренирована на крупном наборе данных изображений, таком как ImageNet. Эти модели уже выполняют хорошую работу, извлекая различные функции из универсальных образов, и во многих случаях просто создание классификатора на основе этих извлеченных признаков может дать хороший результат.
import tensorflow as tf
import keras
import matplotlib.pyplot as plt
import numpy as np
import os
import glob
from PIL import Image
Набор данных кошек и собак
В этом уроке мы решим реальную проблему классификации изображений кошек и собак. По этой причине мы будем использовать набор данных Kaggle Cats vs. Dogs, который также можно скачать из Microsoft.
Давайте скачайте этот набор данных и извлеким его в data каталог:
import urllib.request
import zipfile
dataset_url = 'https://download.microsoft.com/download/3/E/1/3E1C3F21-ECDB-4869-8368-6DEBA77B919F/kagglecatsanddogs_5340.zip'
data_dir = 'data'
os.makedirs(data_dir, exist_ok=True)
zip_path = os.path.join(data_dir, 'kagglecatsanddogs_5340.zip')
if not os.path.exists(zip_path):
urllib.request.urlretrieve(dataset_url, zip_path)
if not os.path.exists(os.path.join(data_dir, 'PetImages')):
with zipfile.ZipFile(zip_path, 'r') as zip_ref:
zip_ref.extractall(data_dir)
Набор данных может содержать несколько поврежденных файлов изображений. Давайте определим вспомогательную функцию, чтобы проверить и удалить их перед загрузкой:
def check_image(fn):
try:
im = Image.open(fn)
im.verify()
return True
except (IOError, SyntaxError):
return False
def check_image_dir(dir_path):
for fn in glob.glob(dir_path):
if not check_image(fn):
print(f"Corrupt image: {fn}")
os.remove(fn)
# Remove any corrupt images from the dataset
check_image_dir('data/PetImages/Cat/*.jpg')
check_image_dir('data/PetImages/Dog/*.jpg')
Загрузка набора данных
В предыдущих примерах мы загружали наборы данных, встроенные в Keras. Теперь мы будем использовать собственный набор данных, который необходимо загрузить из каталога изображений. Keras включает вспомогательную функцию image_dataset_from_directory, которая может создавать tf.data.Dataset из каталога изображений, организованных в подкаталоги по классу.
data_dir = 'data/PetImages'
batch_size = 32
ds_train = keras.utils.image_dataset_from_directory(
data_dir,
validation_split=0.2,
subset='training',
seed=13,
image_size=(224, 224),
batch_size=batch_size
)
ds_test = keras.utils.image_dataset_from_directory(
data_dir,
validation_split=0.2,
subset='validation',
seed=13,
image_size=(224, 224),
batch_size=batch_size
)
Замечание
Мы используем одно и то же seed значение при создании разделений обучения и проверки, чтобы обеспечить отсутствие перекрытия между двумя подмножествами.
Мы можем проверить имена классов, которые были автоматически выведены из структуры каталогов:
# Expected output: ['Cat', 'Dog']
ds_train.class_names
Давайте определим вспомогательный элемент для визуализации примеров из нашего набора данных (это новая версия display_dataset , адаптированная для пакетных данных):
def display_dataset(images, labels, classes=None, cols=8):
n = len(images)
rows = (n + cols - 1) // cols
fig, axes = plt.subplots(rows, cols, figsize=(cols * 1.5, rows * 1.5))
axes = axes.flatten() if n > 1 else [axes]
for i, ax in enumerate(axes):
if i < n:
ax.imshow(images[i])
label = int(labels[i][0]) if labels[i].ndim > 0 else int(labels[i])
title = classes[label] if classes else str(label)
ax.set_title(title, fontsize=8)
ax.axis('off')
plt.tight_layout()
plt.show()
Набор данных выдает пакеты изображений и меток. Каждый пакет содержит 32 изображения размером 224×224 с 3 цветовыми каналами и соответствующими метками:
for x, y in ds_train:
print(f"Training batch shape: features={x.shape}, labels={y.shape}")
x_sample, y_sample = x, y
break
# Expected output: Training batch shape: features=(32, 224, 224, 3), labels=(32,)
display_dataset(x_sample.numpy().astype(np.uint8), np.expand_dims(y_sample, 1), classes=ds_train.class_names)
Замечание
Значения пикселей изображения находятся в диапазоне от 0 до 255. Для некоторых моделей требуется масштабировать входные данные до 0–1 или предварительно обработать с помощью функции, конкретной модели. VGG-16 имеет собственную preprocess_input функцию, которую мы используем позже.
Предварительно обученные модели
Существует множество предварительно обученных нейронных сетей для классификации изображений, которые были обучены в наборе данных ImageNet, который содержит более 14 миллионов изображений в 1000 категориях. Одной из самых известных архитектур является VGG-16, которая обеспечивает хорошую точность при простом понимании. Давайте загрузим модель VGG-16 с предварительно обученными весами:
vgg = keras.applications.VGG16()
Давайте попробуем использовать эту предварительно обученную сеть для классификации одного из наших образов. Сеть VGG-16 была обучена на ImageNet, которая включает категории для различных собак и кошачьих пород:
inp = keras.applications.vgg16.preprocess_input(x_sample[:1])
res = vgg(inp)
# tf.argmax returns the index of the highest-probability class
print(f"Most probable class = {tf.argmax(res, 1)}")
# decode_predictions maps class indices to human-readable labels
keras.applications.vgg16.decode_predictions(res.numpy())
Функция preprocess_input масштабирует значения пикселей соответствующим образом для модели VGG-16. Функция decode_predictions возвращает наиболее вероятные классы ImageNet top-5, а также их оценки достоверности.
Давайте рассмотрим архитектуру VGG-16:
# Shows all layers including convolutional blocks and final Dense classifier
vgg.summary()
Вычисления GPU
Глубокие нейронные сети требуют достаточно основной вычислительной мощности для обучения. Использование GPU может значительно ускорить процесс обучения. Давайте проверим, доступен ли GPU:
# Lists available GPU devices; an empty list means CPU-only
tf.config.list_physical_devices('GPU')
Извлечение функций VGG
Если мы хотим использовать VGG-16 для извлечения признаков из изображений, нам нужна модель без окончательных уровней классификации. Для этого можно указать include_top=False:
vgg = keras.applications.VGG16(include_top=False)
inp = keras.applications.vgg16.preprocess_input(x_sample[:1])
res = vgg(inp)
# The output is a 7x7 grid of 512 feature maps
print(f"Shape after applying VGG-16: {res[0].shape}")
plt.figure(figsize=(15, 3))
plt.imshow(res[0].numpy().reshape(-1, 512))
Результирующий вектор признаков имеет фигуру 7×7×512 = 25088 значений. Это представляет высокоуровневые признаки, которые VGG-16 научилась извлекать из изображения. Мы можем вручную предварительно вычислить эти признаки для всего набора данных, а затем обучить классификатор на основе этого.
Предупреждение
Мы используем .take(25) и .take(10) ниже, чтобы ограничить размер набора данных для ускорения обучения в этом примере. Каждый пакет содержит 32 изображения, поэтому мы используем только 800 обучающих образов и 320 тестовых образов. Точность около 90%, сообщаемая здесь, отражает это небольшое подмножество и может не обобщать полный набор данных. Для использования в рабочей среде обучите полный набор данных.
def preprocess(x, y):
return keras.applications.vgg16.preprocess_input(x), y
ds_features_train = ds_train.take(25).map(preprocess).map(lambda x, y: (vgg(x), y)).cache()
ds_features_test = ds_test.take(10).map(preprocess).map(lambda x, y: (vgg(x), y)).cache()
for x, y in ds_features_train:
# Expected output: (32, 7, 7, 512) (32,)
print(x.shape, y.shape)
break
Замечание
Мы вызываем .cache() после извлечения функций, чтобы модель VGG-16 выполнялось только один раз на каждый пакет вместо каждой эпохи.
Теперь мы можем создать простой классификатор на извлеченных функциях. Так как функции VGG уже высокоинформативны, даже один плотный слой может достичь хороших результатов:
model = keras.Sequential([
keras.layers.Input(shape=(7, 7, 512)),
keras.layers.Flatten(),
keras.layers.Dense(1, activation='sigmoid')
])
model.compile(optimizer='adam', loss='binary_crossentropy', metrics=['accuracy'])
hist = model.fit(ds_features_train, validation_data=ds_features_test)
# Expected: validation accuracy around 90%
С точностью около 90%, это демонстрирует возможности предварительно обученных функций! Однако предварительный расчёт признаков вручную является утомительным.
Передача обучения с помощью одной сети VGG
Мы можем избежать предварительной компиляции функций, объединив средство извлечения компонентов VGG-16 и классификатор в одну сеть. Ключ заключается в том, чтобы заморозить предварительно обученные слои, чтобы их веса не обновлялись во время обучения.
Мы перемещаем preprocess_input шаг в конвейер данных, а не встраиваем его в модель как Lambda слой. Это позволяет сериализовать модель, чтобы сохранить и загрузить ее позже:
def preprocess(x, y):
return keras.applications.vgg16.preprocess_input(x), y
ds_train_preprocessed = ds_train.map(preprocess)
ds_test_preprocessed = ds_test.map(preprocess)
Замечание
Так как предварительная обработка теперь является частью конвейера данных, а не модели, необходимо также применять preprocess_input к входным данным во время инференции.
Теперь мы создадим модель с замороженной базой VGG-16:
vgg_base = keras.applications.VGG16(include_top=False, input_shape=(224, 224, 3))
vgg_base.trainable = False
model = keras.Sequential([
keras.layers.Input(shape=(224, 224, 3)),
vgg_base,
keras.layers.Flatten(),
keras.layers.Dense(1, activation='sigmoid')
])
# Notice: ~15 million params are non-trainable (VGG-16), only ~25k are trainable
model.summary()
Заморозив слои VGG-16, нам нужно только обучить окончательный плотный слой, который имеет примерно 25 000 параметров вместо полного 15 миллионов. Это ускоряет обучение:
Предупреждение
Как и в предыдущем разделе, мы используем .take(50) и .take(10) ограничиваем набор данных для ускорения обучения. Это означает, что мы обучаемся на примерно 1600 изображениях и проверяем на 320. Результаты точности могут отличаться при обучении по полному набору данных.
model.compile(optimizer='adam', loss='binary_crossentropy', metrics=['accuracy'])
hist = model.fit(ds_train_preprocessed.take(50), validation_data=ds_test_preprocessed.take(10))
# Expected: validation accuracy around 90% or higher
Сохранение и загрузка модели
После того как у нас есть обученная модель, мы можем сохранить ее на диск и перезагрузить ее позже без повторного обучения:
model.save('data/cats_dogs.keras')
Замечание
Расширение .keras использует собственный формат Keras 3. Если вы используете старую версию TensorFlow/Keras, используйте .h5 (формат HDF5) или формат каталога SavedModel.
Чтобы загрузить сохраненную модель, выполните следующие действия.
model = keras.models.load_model('data/cats_dogs.keras')
Другие модели компьютерного зрения
VGG-16 является одной из самых простых глубоких архитектур сверточных нейронных сетей (CNN) для понимания, благодаря своей однородной структуре стека 3×3 сверток. Keras предоставляет гораздо больше предварительно обученных сетей. Наиболее часто используемыми среди них являются архитектуры ResNet, разработанные корпорацией Майкрософт и Inception от Google.
Улучшение результатов с помощью расширения данных
При работе с ограниченным объемом данных для обучения расширение данных может значительно улучшить генерализацию. Применяя случайные преобразования (такие как горизонтальное отражение, повороты и масштабирование) к обучающим изображениям, мы искусственно увеличиваем разнообразие набора данных. Keras предоставляет такие уровни расширения, как keras.layers.RandomFlip, keras.layers.RandomRotationи keras.layers.RandomZoom их можно добавить непосредственно в модель или конвейер данных.
Вывод
С помощью переноса обучения нам удалось быстро создать классификатор для нашей задачи классификации пользовательских объектов и добиться высокой точности. Этот пример не был полностью объективным, так как исходная сеть VGG-16 была предварительно обучена на ImageNet, который уже включает категории для различных пород кошек и собак, и поэтому мы всего лишь использовали большинство шаблонов, которые уже есть в сети. Вы можете ожидать более низкую точность для других объектов, относящихся к домену, таких как сведения о производственной линии в заводе или различных листьях дерева. Вы можете увидеть, что более сложные задачи требуют более высокой вычислительной мощности и часто используют ускорение GPU для обучения.
Проверьте свои знания
Обратная связь
Были ли сведения на этой странице полезными?
Нет
Нужна помощь с этой темой?
Хотите попробовать использовать Ask Learn для уточнения или руководства по этой теме?