Обучение моделей на сервере

Keras.js — это библиотека, позволяющая запускать нейронные сети, созданные с использованием Keras, непосредственно в браузере с помощью JavaScript. Она не предоставляет возможностей для обучения моделей на клиентской стороне, поэтому обучение должно происходить на сервере с последующим экспортом модели в формат, совместимый с Keras.js.

Подготовка модели на сервере

  1. Создание модели в Keras Модель создается с использованием стандартных API Keras в Python:

    from keras.models import Sequential
    from keras.layers import Dense, Conv2D, Flatten
    
    model = Sequential([
        Conv2D(32, (3,3), activation='relu', input_shape=(28,28,1)),
        Flatten(),
        Dense(10, activation='softmax')
    ])
  2. Компиляция и обучение модели Модель компилируется с выбором оптимизатора, функции потерь и метрик:

    model.compile(optimizer='adam', loss='categorical_crossentropy', metrics=['accuracy'])
    model.fit(x_train, y_train, epochs=10, batch_size=32, validation_split=0.2)
  3. Экспорт модели в формат Keras.js Для работы в браузере необходимо сохранить веса и структуру модели в JSON и бинарный формат:

    model.save('model.h5')

    Затем с помощью инструмента kerasjs-converter выполняется преобразование:

    kerasjs-converter --input model.h5 --output model_js

    В результате получается папка с файлами model.json и весами в бинарном формате, готовыми к загрузке в Keras.js.

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

Keras.js работает с уже обученными моделями. Процесс включает несколько этапов:

Инициализация модели

import KerasJS from 'keras-js';

const model = new KerasJS.Model({
  filepath: 'model_js/model.json',
  gpu: true
});

Параметр gpu: true включает использование WebGL для ускорения вычислений. Если устройство не поддерживает WebGL, вычисления будут выполняться на CPU.

Подготовка данных для предсказания

Данные должны быть приведены к формату Float32Array и соответствовать входной форме модели. Например, для изображения 28x28 пикселей:

const input = {
  input_1: new Float32Array(28 * 28) // данные должны быть нормализованы
};

Нормализация включает приведение значений пикселей к диапазону [0,1] или [-1,1], в зависимости от того, как модель была обучена на сервере.

Выполнение предсказания

model.ready().then(() => {
  model.predict(input).then(outputData => {
    console.log(outputData);
  });
});

Метод ready() гарантирует, что модель загружена и готова к использованию, а predict() возвращает объект с именами выходных слоев и массивами предсказанных значений.

Управление производительностью

Использование GPU через WebGL позволяет значительно ускорить инференс, особенно для сложных сверточных и рекуррентных сетей. Для небольших моделей и простых операций использование CPU может быть достаточным и более совместимым с мобильными устройствами.

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

Совместимость и ограничения

  • Поддерживаются слои Dense, Conv2D, Flatten, Activation, MaxPooling2D, Dropout, BatchNormalization.
  • Рекуррентные слои (например, LSTM) и кастомные функции активации могут требовать дополнительных конвертаций или не поддерживаться полностью.
  • Обучение непосредственно в браузере не поддерживается: все операции тренировки должны выполняться на сервере или в Python-среде Keras.

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

Keras.js идеально подходит для:

  • Распознавания изображений на клиенте без серверной нагрузки.
  • Приложений, где важна оффлайн работа: модель загружается и выполняется локально.
  • Быстрой прототипизации и демонстрации обученных моделей без необходимости разворачивать полноценный сервер для инференса.

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