Дообучение с разморозкой верхних слоёв

Keras.js представляет собой библиотеку для выполнения моделей, обученных с использованием Keras, непосредственно в браузере на языке JavaScript. Модели в Keras.js хранятся в формате JSON, который описывает структуру сети, а также набор бинарных файлов с весами, соответствующих слоям. Каждый слой модели может содержать веса (weights) и смещения (biases), а также дополнительные параметры для специфических слоев, таких как BatchNormalization (gamma, beta, moving_mean, moving_variance).

Инициализация модели и загрузка весов

Для работы с Keras.js необходимо сначала создать объект модели и загрузить веса:

import KerasJS from 'keras-js';

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

await model.ready();

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

  • filepath — путь к JSON-файлу модели.
  • Параметр gpu позволяет использовать WebGL для ускорения вычислений.
  • Метод ready() возвращает Promise, гарантирующий, что модель полностью загружена и готова к инференсу.

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

const layerWeights = await fetch('layer1_weights.bin').then(res => res.arrayBuffer());
model.layers[0].setWeights(layerWeights);

Механизм разморозки верхних слоев

Разморозка слоев — это процесс, при котором ранее “замороженные” слои модели становятся обучаемыми. В Keras.js это не встроенная функция, как в Python Keras, поэтому требуется ручное управление:

  1. Определение замороженных и обучаемых слоев. Каждому слою соответствует объект с параметром trainable. Для разморозки верхних слоев его необходимо выставить в true:

    for (let i = frozenLayerCount; i < model.layers.length; i++) {
      model.layers[i].trainable = true;
    }
  2. Обновление градиентов. Keras.js поддерживает вычисление градиентов для слоев, помеченных как trainable. При разморозке верхних слоев требуется пересчитать веса этих слоев через backpropagation, используя библиотеку для оптимизации (например, tfjs для работы с TensorFlow.js на стороне клиента):

    const optimizer = new tf.train.adam(0.0001);
    
    function trainStep(inputData, targetData) {
      tf.tidy(() => {
        const xs = tf.tensor(inputData);
        const ys = tf.tensor(targetData);
    
        optimizer.minimize(() => {
          const predictions = model.predict(xs);
          return tf.losses.meanSquaredError(ys, predictions);
        });
      });
    }

Применение дообучения

Дообучение с разморозкой верхних слоев особенно эффективно, если исходная модель была обучена на большом датасете, а новая задача имеет схожую структуру данных. Основные шаги:

  • Выбор верхних слоев для разморозки. Верхние слои обычно отвечают за извлечение признаков высокой абстракции. Их разморозка позволяет модели адаптироваться к новым признакам целевого датасета.
  • Настройка learning rate. Для уже обученных слоев рекомендуется использовать меньший шаг обучения, чтобы не разрушить предварительно выученные признаки.
  • Постепенное обучение. Сначала размораживаются только последние слои, затем можно постепенно подключать средние, сохраняя нижние слои замороженными.
const learningRates = [1e-5, 1e-4];
const layersToUnfreeze = model.layers.slice(-3);

layersToUnfreeze.forEach(layer => layer.trainable = true);

// Обучение с адаптивными скоростями
layersToUnfreeze.forEach((layer, i) => {
  optimizer.setLearningRate(learningRates[i]);
});

Ограничения и особенности

  • Keras.js ориентирован на инференс и дообучение на клиенте требует интеграции с TensorFlow.js для вычисления градиентов.
  • Заморозка и разморозка слоев управляется вручную через свойство trainable.
  • Для больших моделей возможны ограничения по памяти и производительности в браузере, особенно при использовании WebGL.

Практические советы по работе с размороженными слоями

  1. Разморозка только верхних слоев позволяет ускорить обучение и минимизировать переобучение.
  2. При дообучении стоит сохранять промежуточные веса, чтобы при необходимости откатить изменения.
  3. Использование tf.tidy() помогает управлять памятью, освобождая неиспользуемые тензоры после каждой итерации.
  4. Для больших данных рекомендуется выполнять пакетную обработку (batching) для снижения нагрузки на GPU.

Этот подход позволяет адаптировать предобученные модели Keras для специфичных задач прямо в браузере, комбинируя гибкость JavaScript и вычислительную мощность WebGL.