Детекция объектов

Детекция объектов — задача компьютерного зрения, направленная на определение наличия объектов на изображении, их классификацию и локализацию с помощью ограничивающих рамок (bounding boxes). TensorFlow.js позволяет реализовать детекцию объектов прямо в браузере или на сервере с использованием JavaScript, что открывает возможности для интерактивных веб-приложений и визуализаций.


Архитектура моделей детекции

Модели детекции объектов обычно строятся на основе сверточных нейронных сетей (CNN) и включают два основных подхода:

  1. Одностадийные модели (One-Stage Detection) Примеры: YOLO, SSD. Отличие заключается в том, что модель сразу предсказывает координаты ограничивающих рамок и классы объектов без промежуточного этапа предложения регионов. Этот подход обеспечивает высокую скорость, что важно для веб-приложений и мобильных устройств.

  2. Двухстадийные модели (Two-Stage Detection) Примеры: Faster R-CNN. Сначала генерируются предложения регионов, затем каждый регион классифицируется и уточняется. Двухстадийные модели обычно точнее, но медленнее.

TensorFlow.js поддерживает оба подхода через портированные модели из TensorFlow и готовые модели, доступные в @tensorflow-models.


Загрузка и использование предобученной модели

TensorFlow.js предоставляет несколько готовых моделей детекции объектов, среди которых coco-ssd и mobilenet. Основные шаги:

import * as tf from '@tensorflow/tfjs';
import * as cocossd from '@tensorflow-models/coco-ssd';

async function detectObjects(imageElement) {
    const model = await cocossd.load();
    const predictions = await model.detect(imageElement);

    predictions.forEach(prediction => {
        console.log(`Объект: ${prediction.class}, вероятность: ${prediction.score}`);
        console.log(`Координаты: x=${prediction.bbox[0]}, y=${prediction.bbox[1]}, width=${prediction.bbox[2]}, height=${prediction.bbox[3]}`);
    });
}

Ключевые моменты:

  • Метод cocossd.load() загружает предобученную модель COCO-SSD.
  • Метод detect() принимает HTML-элемент <img>, <video> или <canvas> и возвращает массив объектов с полями class, score и bbox.
  • bbox — массив [x, y, width, height], задающий ограничивающую рамку объекта.

Подготовка данных для обучения и инференса

Для обучения собственной модели детекции объектов требуется:

  1. Разметка данных Формат COCO или Pascal VOC содержит координаты bounding box и классы объектов для каждого изображения.
  2. Преобразование изображений в тензоры Используется tf.browser.fromPixels(image) для преобразования изображения в тензор с формой [height, width, 3].
  3. Нормализация Значения пикселей масштабируются в диапазон [0, 1] или [-1, 1] в зависимости от модели.

Пример преобразования:

const imgTensor = tf.browser.fromPixels(imageElement).toFloat().div(255.0);
const batched = imgTensor.expandDims(0); // добавление батч-размера

Создание и обучение собственной модели

  1. Базовая архитектура

Можно использовать сверточные слои с несколькими выходами: один для координат ограничивающей рамки, другой для вероятностей классов.

const model = tf.sequential();
model.add(tf.layers.conv2d({inputShape: [224, 224, 3], filters: 16, kernelSize: 3, activation: 'relu'}));
model.add(tf.layers.maxPooling2d({poolSize: 2}));
model.add(tf.layers.conv2d({filters: 32, kernelSize: 3, activation: 'relu'}));
model.add(tf.layers.flatten());
model.add(tf.layers.dense({units: 4, activation: 'linear'})); // координаты bbox
model.add(tf.layers.dense({units: numClasses, activation: 'softmax'})); // классы
  1. Функция потерь

Для детекции обычно комбинируют два типа потерь:

  • Localization loss — ошибка в предсказанных координатах (например, MSE или Smooth L1).
  • Classification loss — ошибка классификации (обычно categoricalCrossentropy).
const loss = (yTrue, yPred) => {
    const locLoss = tf.losses.meanSquaredError(yTrue.bbox, yPred.bbox);
    const clsLoss = tf.losses.softmaxCrossEntropy(yTrue.classes, yPred.classes);
    return locLoss.add(clsLoss);
};
  1. Тренировка
model.compile({optimizer: 'adam', loss: loss});
await model.fit(trainDataset, {
    epochs: 50,
    validationData: valDataset,
});

Оптимизация модели для работы в браузере

Для детекции объектов в реальном времени важно учитывать производительность:

  • Использовать модели с малым числом параметров (MobileNet, TinyYOLO).
  • Применять tf.tidy() для очистки промежуточных тензоров и уменьшения утечек памяти.
  • Сжимать изображения до небольшого разрешения при инференсе, чтобы ускорить вычисления.
  • Использовать WebGL-ускорение (tf.setBackend('webgl')).

Визуализация результатов

Для отображения bounding box на изображении удобно использовать <canvas>:

const ctx = canvas.getContext('2d');
ctx.drawImage(imageElement, 0, 0);

predictions.forEach(prediction => {
    ctx.strokeStyle = 'red';
    ctx.lineWidth = 2;
    ctx.strokeRect(prediction.bbox[0], prediction.bbox[1], prediction.bbox[2], prediction.bbox[3]);
    ctx.fillStyle = 'red';
    ctx.fillText(`${prediction.class} (${(prediction.score*100).toFixed(1)}%)`, prediction.bbox[0], prediction.bbox[1] - 5);
});

Потенциал интеграции с веб-приложениями

TensorFlow.js позволяет использовать детекцию объектов для:

  • Интерфейсов с дополненной реальностью (AR) в браузере.
  • Веб-приложений с контролем качества или мониторингом безопасности.
  • Интерактивных обучающих платформ и игр, распознающих объекты в реальном времени.

Эта библиотека обеспечивает сочетание высокой производительности и гибкости, позволяя работать как с готовыми моделями, так и создавать собственные решения.