Ручной тренировочный цикл

Ручной тренировочный цикл предоставляет полный контроль над процессом обучения нейронной сети. В отличие от стандартного метода model.fit, где обучение происходит автоматически, ручной цикл позволяет задавать точные шаги оптимизации, вычислять потери, управлять градиентами и более гибко интегрировать кастомные функции потерь или метрики.


Основные компоненты

  1. Определение модели

Для ручного цикла можно использовать как функциональный API, так и tf.Sequential. Пример простой модели для регрессии:

const model = tf.sequential();
model.add(tf.layers.dense({units: 10, activation: 'relu', inputShape: [5]}));
model.add(tf.layers.dense({units: 1}));
  1. Выбор оптимизатора и функции потерь

Оптимизатор управляет обновлением весов. Наиболее часто используются:

const optimizer = tf.train.adam(0.01); // Adam с шагом обучения 0.01
const lossFn = (pred, label) => pred.sub(label).square().mean(); // MSE

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

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

const xs = tf.tensor2d([[1, 2, 3, 4, 5], [5, 4, 3, 2, 1]]);
const ys = tf.tensor2d([[10], [10]]);

Важно нормализовать данные и следить за формой тензоров, чтобы она соответствовала входным слоям модели.


Основная структура ручного цикла

Ручной цикл состоит из нескольких шагов:

  1. Проход вперед (forward pass) Вычисление предсказаний модели на текущем батче данных:

    const preds = model.predict(xs);
  2. Вычисление функции потерь Сравнение предсказаний с целевыми значениями:

    const loss = lossFn(preds, ys);
  3. Вычисление градиентов TensorFlow.js предоставляет tf.variableGrads или tf.grad для вычисления производных:

    const grads = tf.variableGrads(() => lossFn(model.predict(xs), ys));
  4. Обновление весов Оптимизатор применяет вычисленные градиенты к весам модели:

    optimizer.applyGradients(grads.grads);
  5. Освобождение памяти Для предотвращения утечек памяти тензоры, которые больше не нужны, необходимо очищать:

    preds.dispose();
    loss.dispose();
    for (const key in grads.grads) {
      grads.grads[key].dispose();
    }

Итеративное обучение

Ручной цикл обычно оборачивается в цикл по эпохам и батчам. Пример обучения в течение 100 эпох:

const epochs = 100;

for (let epoch = 0; epoch < epochs; epoch++) {
  tf.tidy(() => {
    const preds = model.predict(xs);
    const loss = lossFn(preds, ys);

    const grads = tf.variableGrads(() => lossFn(model.predict(xs), ys));
    optimizer.applyGradients(grads.grads);

    console.log(`Эпоха ${epoch + 1}: Потеря = ${loss.dataSync()[0]}`);

    // Очистка промежуточных тензоров
    for (const key in grads.grads) {
      grads.grads[key].dispose();
    }
  });
}

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


Работа с батчами

Для больших наборов данных обучение происходит по батчам:

const batchSize = 2;
const numBatches = Math.ceil(xs.shape[0] / batchSize);

for (let i = 0; i < numBatches; i++) {
  const start = i * batchSize;
  const end = start + batchSize;
  const batchXs = xs.slice([start, 0], [batchSize, xs.shape[1]]);
  const batchYs = ys.slice([start, 0], [batchSize, ys.shape[1]]);

  tf.tidy(() => {
    const grads = tf.variableGrads(() => lossFn(model.predict(batchXs), batchYs));
    optimizer.applyGradients(grads.grads);
    for (const key in grads.grads) {
      grads.grads[key].dispose();
    }
  });
}

Использование батчей уменьшает потребление памяти и стабилизирует градиенты.


Кастомные метрики и контроль градиентов

Ручной цикл позволяет:

  • Вычислять собственные метрики в каждой эпохе.
  • Ограничивать градиенты через tf.clipByValue или tf.clipByNorm:
for (const key in grads.grads) {
  grads.grads[key] = grads.grads[key].clipByValue(-1, 1);
}
  • Сохранять или визуализировать значения потерь на каждом шаге.

Преимущества ручного подхода

  • Полный контроль над процессом обучения.
  • Возможность применять сложные схемы оптимизации и кастомные функции потерь.
  • Легкая интеграция с другими вычислительными процессами и нестандартными архитектурами моделей.
  • Точная настройка обновления весов и управления памятью.

Ручной тренировочный цикл является основой для экспериментов с нейросетями и глубокого понимания механики обучения моделей в TensorFlow.js.