Keras.js — это библиотека, позволяющая запускать нейронные сети, созданные с использованием Keras, непосредственно в браузере с помощью JavaScript. Она не предоставляет возможностей для обучения моделей на клиентской стороне, поэтому обучение должно происходить на сервере с последующим экспортом модели в формат, совместимый с Keras.js.
Создание модели в 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')
])Компиляция и обучение модели Модель компилируется с выбором оптимизатора, функции потерь и метрик:
model.compile(optimizer='adam', loss='categorical_crossentropy', metrics=['accuracy'])
model.fit(x_train, y_train, epochs=10, batch_size=32, validation_split=0.2)Экспорт модели в формат 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) и кастомные функции
активации могут требовать дополнительных конвертаций или не
поддерживаться полностью.Keras.js идеально подходит для:
Для динамических приложений с изменением модели или дообучением рекомендуется комбинировать Keras.js с серверным API, которое занимается обучением и пересозданием модели в формате, совместимом с браузером.