Подгонка модели: fit и fitDataset

Подгонка модели — центральный этап обучения нейронной сети в TensorFlow.js. Этот процесс заключается в оптимизации весов модели с целью минимизации функции потерь на тренировочных данных. TensorFlow.js предоставляет два основных метода для подгонки: fit и fitDataset. Каждый из них ориентирован на разные сценарии использования и форматы данных.


Метод fit

Метод fit используется для подгонки модели на данных, представленных в виде тензоров или массивов JavaScript. Он подходит для небольших и средних наборов данных, помещающихся в память браузера или Node.js.

Сигнатура метода:

model.fit(x, y, {
  epochs: 10,
  batchSize: 32,
  validationSplit: 0.2,
  callbacks: [...]
});

Аргументы:

  • x — входные данные. Может быть tf.Tensor или массив, где каждая запись соответствует одному примеру.
  • y — целевые значения (метки), аналогично формата x.
  • epochs — количество проходов по всему тренировочному набору.
  • batchSize — размер мини-батча. Оптимизация весов происходит после обработки каждого батча.
  • validationSplit — доля данных, выделяемая для валидации. TensorFlow.js случайным образом отделяет указанную часть.
  • callbacks — массив объектов обратного вызова (tf.callbacks) для мониторинга процесса обучения (например, tf.callbacks.earlyStopping или tf.callbacks.tensorBoard).

Особенности метода fit:

  • Все данные должны помещаться в память. Для очень больших наборов данных использование fit неэффективно.
  • Подходит для статических данных, загруженных заранее.
  • Автоматически поддерживается разбиение на батчи и валидация, что упрощает настройку обучения.

Пример использования fit:

const xs = tf.tensor2d([[0, 0], [0, 1], [1, 0], [1, 1]]);
const ys = tf.tensor2d([[0], [1], [1], [0]]);

await model.fit(xs, ys, {
  epochs: 100,
  batchSize: 4,
  validationSplit: 0.25,
  callbacks: tf.callbacks.earlyStopping({monitor: 'loss', patience: 10})
});

Метод fitDataset

Метод fitDataset предназначен для обучения на потоковых или очень больших данных, которые невозможно полностью загрузить в память. Данные подаются через объект tf.data.Dataset, поддерживающий ленивую загрузку и батчирование.

Сигнатура метода:

model.fitDataset(dataset, {
  epochs: 10,
  batchesPerEpoch: 100,
  callbacks: [...]
});

Аргументы:

  • dataset — объект tf.data.Dataset. Обычно создается из массивов, файлов или потоков данных.
  • epochs — количество проходов по всему датасету.
  • batchesPerEpoch — количество батчей, обрабатываемых за одну эпоху. Используется для ограничения времени обучения при бесконечных или очень больших потоках.
  • callbacks — аналогично методу fit, поддерживаются обратные вызовы.

Особенности метода fitDataset:

  • Поддерживает ленивую загрузку, что позволяет работать с набором данных любого размера.
  • Можно легко использовать аугментацию данных и сложные пайплайны предварительной обработки через функции dataset.map, dataset.shuffle, dataset.batch.
  • Полезен для интеграции с потоками данных из внешних источников (например, CSV-файлы, изображения, веб-API).

Пример создания Dataset и обучения:

const dataset = tf.data.array([
  {xs: [0, 0], ys: [0]},
  {xs: [0, 1], ys: [1]},
  {xs: [1, 0], ys: [1]},
  {xs: [1, 1], ys: [0]}
]).map(item => ({
  xs: tf.tensor2d([item.xs]),
  ys: tf.tensor2d([item.ys])
}))
.batch(2);

await model.fitDataset(dataset, {
  epochs: 50,
  batchesPerEpoch: 2
});

Отличия fit и fitDataset

Характеристика fit fitDataset
Размер данных Ограничен памятью Любой, потоковый
Формат входных данных tf.Tensor или массив tf.data.Dataset
Поддержка батчирования Встроена через batchSize Обрабатывается в Dataset
Валидация validationSplit или validationData Не встроена, нужна ручная реализация
Применимость Малые и средние наборы Большие наборы, стриминговые данные

Использование обратных вызовов

Методы fit и fitDataset поддерживают общие обратные вызовы, позволяющие:

  • Следить за метриками (onEpochEnd, onBatchEnd).
  • Раннее прекращение обучения (tf.callbacks.earlyStopping).
  • Логирование в TensorBoard (tf.callbacks.tensorBoard).

Обратные вызовы критически важны при работе с большими данными или при длительном обучении, позволяя контролировать процесс и предотвращать переобучение.


Практические рекомендации

  • Для небольших наборов данных проще и удобнее использовать fit.
  • Для больших наборов или потоковых данных необходимо использовать fitDataset.
  • При работе с fitDataset всегда проверять корректность преобразований через map и batch, чтобы избежать ошибок при несовпадении форм тензоров.
  • Использование callbacks позволяет динамически управлять обучением и сокращать время до достижения оптимальной модели.

Методы fit и fitDataset обеспечивают гибкость обучения нейронных сетей в TensorFlow.js, позволяя эффективно обрабатывать как малые, так и масштабные наборы данных, интегрировать сложные пайплайны обработки и управлять процессом обучения через обратные вызовы.