TensorFlow.js — это мощная библиотека для машинного обучения в среде JavaScript, которая позволяет создавать, обучать и использовать модели непосредственно в браузере или в Node.js. Одним из ключевых преимуществ TensorFlow.js является возможность обучения моделей на клиентских устройствах без передачи данных на сервер, что обеспечивает конфиденциальность, снижение задержек и независимость от серверной инфраструктуры.
Тензоры — это центральная структура данных
TensorFlow.js. Тензор представляет собой многомерный массив чисел с
определённой формой и типом данных (tf.Tensor). Все
операции с данными выполняются через тензоры, что позволяет эффективно
использовать GPU и CPU для вычислений.
Модели в TensorFlow.js могут быть созданы двумя основными способами:
Оптимизаторы управляют процессом обучения,
корректируя веса модели на основе вычисленных градиентов. Важнейшие
оптимизаторы в TensorFlow.js: sgd, adam,
rmsprop.
Функции потерь (loss functions) измеряют, насколько
предсказания модели отличаются от реальных значений. Для задач регрессии
обычно используется meanSquaredError, для классификации —
categoricalCrossentropy или
binaryCrossentropy.
Обучение на устройстве требует, чтобы данные были локальными. Для этого TensorFlow.js предлагает несколько способов работы с данными:
Важно: при обучении на клиенте необходимо контролировать размер
батчей, так как устройства могут иметь ограниченные ресурсы.
Использование метода .batch(batchSize) позволяет
регулировать нагрузку на память и процессор.
Пример создания простой модели для классификации изображений:
const model = tf.sequential();
model.add(tf.layers.flatten({ inputShape: [28, 28, 1] }));
model.add(tf.layers.dense({ units: 128, activation: 'relu' }));
model.add(tf.layers.dense({ units: 10, activation: 'softmax' }));
model.compile({
optimizer: 'adam',
loss: 'categoricalCrossentropy',
metrics: ['accuracy']
});
Ключевые моменты компиляции:
optimizer управляет скоростью обучения и направлением
корректировки весов.loss определяет функцию ошибки.metrics позволяет отслеживать показатели качества
модели во время обучения.Метод model.fit() выполняет обучение модели на
предоставленных данных. Для обучения в браузере необходимо учитывать
ограничения по ресурсам:
await model.fit(trainXs, trainYs, {
epochs: 10,
batchSize: 32,
validationSplit: 0.2,
callbacks: tf.callbacks.earlyStopping({ monitor: 'val_loss', patience: 3 })
});
Особенности обучения на клиенте:
batchSize влияет на использование памяти. Меньшие батчи
подходят для слабых устройств.validationSplit позволяет автоматически выделить часть
данных для проверки модели.callbacks дают возможность контролировать процесс
обучения: остановка при достижении определённого качества
(earlyStopping), визуализация прогресса
(tfvis) и сохранение промежуточных весов.Для работы с изображениями или большими массивами данных лучше использовать асинхронные генераторы, чтобы не блокировать основной поток браузера.
TensorFlow.js позволяет загружать модели, обученные на сервере, и дообучать их на устройстве:
const pretrainedModel = await tf.loadLayersModel('path/to/model.json');
// Замораживание исходных слоёв
for (const layer of pretrainedModel.layers) {
layer.trainable = false;
}
// Добавление нового выходного слоя
const output = tf.layers.dense({ units: 5, activation: 'softmax' }).apply(pretrainedModel.output);
const newModel = tf.model({ inputs: pretrainedModel.inputs, outputs: output });
newModel.compile({ optimizer: 'adam', loss: 'categoricalCrossentropy', metrics: ['accuracy'] });
Такая техника называется transfer learning и позволяет обучать модель на специфических данных пользователя без необходимости пересоздавать всю архитектуру.
tf.getBackend().tf.nextFrame() помогает разбивать вычисления на чанки,
чтобы интерфейс оставался отзывчивым.Модель можно сохранять локально в IndexedDB или скачивать на устройство:
// Сохранение в IndexedDB
await model.save('indexeddb://my-local-model');
// Загрузка
const loadedModel = await tf.loadLayersModel('indexeddb://my-local-model');
Это позволяет хранить пользовательские модели без передачи данных на сервер и использовать их при следующем запуске приложения.
Обученные на устройстве модели можно применять сразу:
const prediction = model.predict(testXs);
prediction.print();
Особенности работы:
TensorFlow.js предоставляет полный стек инструментов для локального обучения и инференса моделей, что делает возможным построение современных приложений машинного обучения без серверной инфраструктуры и без утечки данных. Использование tензоров, генераторов данных, transfer learning и локального хранения моделей создаёт безопасную и эффективную среду для ML прямо на устройстве пользователя.