Метод fitDataset для работы с генераторами

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

Основные принципы работы

fitDataset принимает объект dataset, который реализует итератор данных. Каждый элемент этого итератора должен быть объектом с двумя ключами:

  • xs — входные данные модели;
  • ys — целевые значения (метки или ожидаемые результаты).

Объекты xs и ys могут быть представлены в виде Tensor, Tensor[] или словарей {'имя_входа': Tensor} для моделей с несколькими входами и выходами. Метод последовательно извлекает данные из генератора и выполняет оптимизацию весов модели на каждой порции.

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

model.fitDataset(dataset, {
  epochs: 10,
  stepsPerEpoch: 100,
  validationData: valDataset,
  validationSteps: 20,
  callbacks: [callback1, callback2]
});

Параметры:

  • dataset — объект типа tf.data.Dataset или совместимый генератор данных.
  • epochs — количество проходов по всему набору данных.
  • stepsPerEpoch — количество шагов (батчей), после которых считается завершённым один проход по данным. Если не указан, считается полное прохождение генератора до завершения итератора.
  • validationData — генератор или tf.data.Dataset для валидации модели.
  • validationSteps — количество шагов для валидации.
  • callbacks — массив функций обратного вызова для отслеживания прогресса обучения, сохранения модели и адаптации скорости обучения.

Формат генератора

Генератор должен возвращать объект в следующем формате:

{
  xs: tf.tensor([...]),  // или словарь входов
  ys: tf.tensor([...])   // или словарь выходов
}

Каждый вызов генератора возвращает одну порцию данных (батч). Метод fitDataset автоматически обрабатывает батчи и накапливает градиенты для обновления весов модели.

Использование батчей и потоков данных

Работа с потоками данных через fitDataset позволяет:

  • Обрабатывать большие наборы данных, которые не помещаются в оперативную память.
  • Применять динамическое предварительное преобразование входов (например, аугментацию изображений или нормализацию) на лету.
  • Объединять несколько источников данных в один генератор.

Пример создания генератора с батчами:

function* dataGenerator() {
  while (true) {
    const xs = tf.randomNormal([32, 28, 28, 1]);  // батч из 32 изображений
    const ys = tf.oneHot(tf.randomUniform([32], 0, 10, 'int32'), 10);  // соответствующие метки
    yield { xs, ys };
  }
}

const dataset = tf.data.generator(dataGenerator).batch(32);

Валидация модели на лету

Параметр validationData позволяет интегрировать проверку точности и потерь модели во время обучения. Он может быть представлен как:

  • генератор с объектами {xs, ys};
  • объект tf.data.Dataset.

validationSteps указывает, сколько батчей будет использоваться для расчёта метрик валидации на каждой эпохе.

Примеры обратных вызовов

Для контроля процесса обучения удобно использовать обратные вызовы (callbacks):

  • tf.callbacks.EarlyStopping — прекращение обучения при отсутствии улучшений метрики;
  • tf.callbacks.ModelCheckpoint — сохранение весов модели после каждой эпохи;
  • tf.callbacks.LearningRateScheduler — динамическая адаптация скорости обучения.

Пример:

const earlyStop = tf.callbacks.earlyStopping({ monitor: 'val_loss', patience: 5 });

await model.fitDataset(dataset, {
  epochs: 50,
  stepsPerEpoch: 100,
  validationData: valDataset,
  validationSteps: 20,
  callbacks: [earlyStop]
});

Отличия от метода fit

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

Оптимизация и производительность

Для ускорения обучения рекомендуется:

  • Кэшировать данные в tf.data.Dataset при возможности повторного использования.
  • Применять предварительную загрузку (prefetch) для скрытия задержек ввода/вывода.
  • Параллелить обработку с помощью dataset.map(fn, {numParallelCalls: ...}).

Пример:

const preparedDataset = dataset
  .map(preprocessFn, { numParallelCalls: 4 })
  .batch(32)
  .prefetch(2);

Такой подход позволяет сократить простои GPU или WebGL при обработке данных и повысить эффективность обучения.

Поддержка моделей с несколькими входами и выходами

fitDataset полностью поддерживает многовходовые и многоцелевые модели Keras. В этом случае объекты генератора должны возвращать словари:

yield {
  xs: { input1: tensor1, input2: tensor2 },
  ys: { output1: tensorOut1, output2: tensorOut2 }
};

Метод корректно вычисляет потери и метрики для всех выходов модели и суммирует их при обновлении весов.


Если требуется, могу подготовить дополнительный раздел с подробным примером обучения сложной модели CNN через fitDataset на реальных данных, с пояснением всех шагов. Это будет полезно для учебника. Хотите, чтобы я это сделал?