Дообучение верхних слоёв

Дообучение (fine-tuning) верхних слоёв нейронной сети позволяет адаптировать предварительно обученную модель к новой задаче или конкретному набору данных, сохранив при этом уже изученные низкоуровневые признаки. В TensorFlow.js это реализуется через работу с предобученными моделями, их модификацию и последующую тренировку ограниченного числа слоёв.


Предобученные модели в TensorFlow.js

TensorFlow.js предоставляет ряд моделей, обученных на больших датасетах, таких как ImageNet. Среди популярных:

  • MobileNet — компактная свёрточная сеть для задач классификации изображений.
  • Inception — глубокая сеть для точной классификации, но с большим числом параметров.
  • Coco-SSD — для задач детекции объектов в реальном времени.

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


Загрузка и подготовка модели

Пример загрузки MobileNet и отделения верхнего слоя:

import * as tf from '@tensorflow/tfjs';
import * as mobilenet from '@tensorflow-models/mobilenet';

// Загрузка модели без верхнего классификационного слоя
const loadBaseModel = async () => {
    const model = await mobilenet.load({version: 2, alpha: 1.0});
    const layer = model.model.getLayer('conv_pw_13_relu');
    return tf.model({inputs: model.model.inputs, outputs: layer.output});
};

Ключевой момент: использование слоя до последнего классификационного уровня позволяет сохранить извлечение признаков и добавлять новые слои для целевой задачи.


Заморозка нижних слоёв

Для предотвращения разрушения уже изученных признаков нижние слои фиксируются:

baseModel.layers.forEach(layer => {
    layer.trainable = false;
});

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


Добавление новых слоёв

После выделения базовой модели добавляются новые слои для конкретной задачи:

const model = tf.sequential();
model.add(baseModel);
model.add(tf.layers.flatten());
model.add(tf.layers.dense({units: 128, activation: 'relu'}));
model.add(tf.layers.dropout({rate: 0.5}));
model.add(tf.layers.dense({units: NUM_CLASSES, activation: 'softmax'}));

Пояснение ключевых компонентов:

  • flatten() преобразует выходные признаки свёрточной сети в одномерный вектор.
  • dense() формирует полносвязные слои для классификации.
  • dropout() снижает переобучение за счёт случайного “выключения” нейронов во время тренировки.

Компиляция модели

Выбор оптимизатора и функции потерь критичен для успешного дообучения:

model.compile({
    optimizer: tf.train.adam(0.0001),
    loss: 'categoricalCrossentropy',
    metrics: ['accuracy']
});

Особенности:

  • Малый learning rate (0.0001) предотвращает разрушение ранее изученных признаков.
  • categoricalCrossentropy подходит для многоклассовой классификации.

Подготовка данных

Для дообучения необходимо привести изображения к единому формату и нормализовать пиксели:

const preprocessImage = (img) => {
    return tf.tidy(() => {
        let tensor = tf.browser.fromPixels(img)
            .resizeNearestNeighbor([224, 224])
            .toFloat()
            .div(tf.scalar(127))
            .sub(tf.scalar(1));
        return tensor.expandDims();
    });
};

Заметка: tf.tidy автоматически освобождает промежуточные тензоры, что критично для работы в браузере.


Тренировка верхних слоёв

Процесс обучения ограничивается новым полносвязным слоем:

await model.fitDataset(trainDataset, {
    epochs: 10,
    validationData: valDataset,
    callbacks: tf.callbacks.earlyStopping({monitor: 'val_loss', patience: 3})
});

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

  • Использование fitDataset позволяет обрабатывать большие наборы данных без загрузки всех данных в память.
  • EarlyStopping предотвращает переобучение, прерывая тренировку при отсутствии улучшений.

Проверка модели

Для оценки точности на тестовых данных используется метод evaluate:

const evalResult = await model.evaluateDataset(testDataset);
console.log(`Test Loss: ${evalResult[0].dataSync()}, Test Accuracy: ${evalResult[1].dataSync()}`);

Сохранение и использование модели

После дообучения модель сохраняется для дальнейшего использования в браузере или Node.js:

await model.save('localstorage://my-finetuned-model');

Загрузка модели для инференса:

const loadedModel = await tf.loadLayersModel('localstorage://my-finetuned-model');

Вывод: дообучение верхних слоёв позволяет адаптировать мощные предобученные модели под специфические задачи, экономя вычислительные ресурсы и обеспечивая высокую точность при ограниченных данных.