Метод 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]
});
Параметры:
tf.data.Dataset
или совместимый генератор данных.tf.data.Dataset для валидации модели.Генератор должен возвращать объект в следующем формате:
{
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):
Пример:
const earlyStop = tf.callbacks.earlyStopping({ monitor: 'val_loss', patience: 5 });
await model.fitDataset(dataset, {
epochs: 50,
stepsPerEpoch: 100,
validationData: valDataset,
validationSteps: 20,
callbacks: [earlyStop]
});
fitfit требует полного набора данных в памяти, тогда как
fitDataset работает с потоками данных.fitDataset позволяет использовать генераторы для
непрерывного потока данных.Для ускорения обучения рекомендуется:
tf.data.Dataset
при возможности повторного использования.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 на реальных данных, с пояснением всех
шагов. Это будет полезно для учебника. Хотите, чтобы я это сделал?