Работа с предобученными моделями

Глубокое обучение для работы с изображениями на PyTorch

Michal Oleszak

Machine Learning Engineer

Использование предобученных моделей

  • Обучение модели с нуля:

    • Длительный процесс
    • Требует большого объёма данных
  • Предобученные модели — модели, уже обученные на определённой задаче

    • Применимы напрямую к новой задаче
    • Требуют адаптации к новой задаче (трансферное обучение)
  • Использование предобученных моделей:

    • Сохранение и загрузка моделей локально
    • Загрузка моделей torchvision
Глубокое обучение для работы с изображениями на PyTorch

Сохранение модели PyTorch

  • torch.save()
  • Расширение файла модели: .pt или .pth
  • Сохранение весов модели с помощью .state_dict()
    torch.save(model.state_dict(), "BinaryCNN.pth")
    
Глубокое обучение для работы с изображениями на PyTorch

Загрузка моделей PyTorch

  • Создайте экземпляр новой модели

    new_model = BinaryCNN()
    
  • Загрузите сохранённые параметры

    new_model.load_state_dict(torch.load('BinaryCNN.pth'))
    
Глубокое обучение для работы с изображениями на PyTorch

Загрузка моделей torchvision

from torchvision.models import (
    resnet18, ResNet18_Weights
)


weights = ResNet18_Weights.DEFAULT
model = resnet18(weights=weights)
transforms = weights.transforms()
  • Импорт архитектуры resnet и весов
  • Извлечение весов
  • Создание экземпляра модели с передачей весов
  • Сохранение необходимых преобразований данных
Глубокое обучение для работы с изображениями на PyTorch

Подготовка новых входных изображений

from PIL import Image

image = Image.open("cat013.jpg")

image_tensor = transform(image)
image_reshaped = image_tensors.unsqueeze(0)

 

изображение кошки

  • Загрузка изображения
  • Преобразование изображения
  • Изменение формы изображения
Глубокое обучение для работы с изображениями на PyTorch

Получение нового предсказания

model.eval()


with torch.no_grad():
pred = model(image_reshaped).squeeze(0)
pred_cls = pred.softmax(0)
cls_id = pred_cls.argmax().item()
cls_name = weights.meta["categories"][cls_id]
print(cls_name)
Egyptian cat
  • Режим оценки для инференса
  • Отключение градиентов
  • Передача изображения в модель и удаление размерности батча
  • Применение softmax
  • Выбор класса с наибольшей вероятностью и извлечение его индекса
  • Сопоставление индекса класса с меткой
  • Вывод метки класса
Глубокое обучение для работы с изображениями на PyTorch

Давайте потренируемся!

Глубокое обучение для работы с изображениями на PyTorch

Preparing Video For Download...