Метод predict и predict_on_batch

Kласс Model в Keras.js предоставляет два основных метода для получения прогнозов по данным: predict и predict_on_batch. Оба метода выполняют вычисление выходных значений нейронной сети на основе входных данных, но имеют различия в подходе к обработке батчей и синхронности выполнения.


Метод predict

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

model.predict(inputData, batchSize);

Параметры:

  • inputData — объект или массив объектов, содержащий входные данные. Данные должны быть представлены в формате Float32Array или аналогичных типов, совместимых с WebGL.
  • batchSize (необязательный) — число примеров, которые будут обрабатываться за один проход через модель. Если не указан, используется значение по умолчанию, установленное при инициализации модели.

Ключевые особенности метода predict:

  1. Автоматическое разделение на батчи. Если входные данные превышают размер одного батча, метод сам разделит их на несколько и объединит результаты.
  2. Асинхронность. Метод возвращает промис, что позволяет выполнять вычисления в фоне, не блокируя основной поток исполнения.
  3. Поддержка множественных входов и выходов. Входные данные могут быть объектом с несколькими ключами, соответствующими именам входов модели, а результатом будет объект или массив прогнозов.

Пример использования:

const input = new Float32Array([0.1, 0.2, 0.3, 0.4]);
model.predict(input, 2).then(predictions => {
    console.log(predictions);
});

В данном примере данные будут обработаны батчами по 2 элемента, а результатом станет массив прогнозов.


Метод predict_on_batch

Метод predict_on_batch выполняет прямое предсказание для одного батча данных, не разделяя входные данные на подбатчи. Это делает его более быстрым для сценариев, когда данные уже подготовлены и не требуют дополнительного разбиения.

Синтаксис:

const output = model.predict_on_batch(batchData);

Параметры:

  • batchData — массив или объект с входными данными, полностью соответствующий размеру батча, для которого требуется прогноз. Размер батча должен совпадать с первым измерением массива данных.

Ключевые особенности метода predict_on_batch:

  1. Синхронная обработка. Возвращает результат напрямую, без промисов, что позволяет использовать его внутри вычислительных циклов или для быстрого тестирования модели.
  2. Нет автоматического разбиения на батчи. Все данные должны помещаться в один батч, иначе возникнет ошибка.
  3. Оптимизация для inference. Используется в случаях, когда требуется быстро получить прогноз для заранее подготовленных батчей данных, например, при генерации изображений или обработке аудио.

Пример использования:

const batchInput = new Float32Array([0.5, 0.6, 0.7, 0.8]);
const batchOutput = model.predict_on_batch(batchInput);
console.log(batchOutput);

Здесь результат будет получен мгновенно для всего переданного батча без разбиения.


Сравнение predict и predict_on_batch

Характеристика predict predict_on_batch
Асинхронность Возвращает промис Синхронный результат
Разбиение на батчи Автоматическое Нет
Использование Удобно для больших наборов данных Оптимально для предобработанных батчей
Множественные входы/выходы Поддерживается Поддерживается
Скорость Медленнее из-за промисов и батчинга Быстрее при фиксированных батчах

Практические рекомендации

  • predict рекомендуется использовать для обучения и тестирования на больших наборах данных, когда точный размер батча заранее неизвестен.
  • predict_on_batch подходит для реального времени и интерактивных приложений, где важна скорость и данные уже организованы в батчи.
  • Для оптимизации использования GPU в браузере следует выбирать batchSize в predict кратным числу ядер устройства или размеру texture в WebGL, чтобы уменьшить накладные расходы.

Обработка множественных выходов

Оба метода поддерживают модели с несколькими выходами. В этом случае:

  • Результат predict будет объектом с ключами, соответствующими именам выходов модели.
  • predict_on_batch возвращает массив выходных значений в том же порядке, в котором они определены в модели.

Пример:

const multiOutputInput = new Float32Array([1.0, 2.0, 3.0, 4.0]);
const outputs = model.predict_on_batch(multiOutputInput);
console.log(outputs[0]); // прогноз для первого выхода
console.log(outputs[1]); // прогноз для второго выхода

Важные нюансы работы с Keras.js

  1. Входные данные должны быть предварительно нормализованы, если модель обучалась на нормализованных данных.
  2. Размерность входного массива должна соответствовать конфигурации модели (inputShape), иначе методы выдадут ошибку.
  3. Для больших моделей важно учитывать потребление памяти GPU, особенно при использовании predict с большим batchSize.

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