Сегментація екземплярів за допомогою Mask R-CNN

Глибоке навчання для зображень із PyTorch

Michal Oleszak

Machine Learning Engineer

Faster R-CNN

Схема архітектури Faster R-CNN

Глибоке навчання для зображень із PyTorch

Mask R-CNN

Схема архітектури Mask R-CNN

Глибоке навчання для зображень із PyTorch

Попередньо натренована Mask R-CNN у PyTorch

from torchvision.models.detection import \
maskrcnn_resnet50_fpn


model = maskrcnn_resnet50_fpn(pretrained=True) model.eval()
image = Image.open("cat_and_laptop.jpg") transform = transforms.Compose([ transforms.ToTensor() ]) image_tensor = transform(image).unsqueeze(0)
with torch.no_grad(): prediction = model(image_tensor)
  • Імпортуйте модель Mask R-CNN
  • Завантажте попередньо натреновану модель
  • Завантажте тестове зображення та перетворіть на тензор

фото кота поруч із ноутбуком

  • Передайте тензор зображення в модель
Глибоке навчання для зображень із PyTorch

Виходи моделі

  • Мітки

    prediction[0]["labels"]
    
    tensor([
        17, 73, 76, 73, 67, 42, 63, 84,73, 65, 
        17, 73, 73, 73, 84, 72, 76, 76,17, 15
    ])
    
  • Назви класів

    print(class_names[17], class_names[73])
    
    cat laptop
    
  • Імовірності класів

    prediction[0]["scores"]
    
    tensor([
        0.9981, 0.9672, 0.9061, 0.6893, 0.3729, 
        ..., 
        0.0745, 0.0705, 0.0623, 0.0610, 0.0508
    ])
    
  • Маски

    prediction[0]["masks"]
    
    tensor([[[[0., 0., 0.,  ..., 0., 0., 0.],
              ...]]]])
    
Глибоке навчання для зображень із PyTorch

М'які маски

  • Унікальні значення маски

    prediction[0]["masks"].unique()
    
    tensor([0.0000e+00, 5.9713e-08, ..., 
            9.9989e-01, 9.9990e-01])
    
  • Маски Mask R-CNN:

    • Значення між 0 і 1
    • Показують упевненість моделі, що піксель належить об'єкту
    • Детальніші за бінарні маски
    • За потреби можна бінаризувати порогом
Глибоке навчання для зображень із PyTorch

Візуалізація м'яких масок

masks = prediction[0]["masks"]
labels = prediction[0]["labels"]


for i in range(2): plt.imshow(image)
plt.imshow( masks[i, 0], cmap="jet", alpha=0.5, )
plt.title( f"Object: {class_names[labels[i]]}" ) plt.show()
  • Витягніть маски й мітки з прогнозу
  • Пройдіть по двох найкращих об'єктах, виводячи оригінальне зображення
  • Для кожного об'єкта накладіть напівпрозору маску
  • Додайте заголовок і покажіть
Глибоке навчання для зображень із PyTorch

Візуалізація м'яких масок

Маски сегментації екземплярів, накладені на зображення

Глибоке навчання для зображень із PyTorch

Давайте потренуємось!

Глибоке навчання для зображень із PyTorch

Preparing Video For Download...