TensorFlow.js предоставляет гибкие возможности для обучения моделей
нейронных сетей прямо в браузере или на сервере с помощью Node.js. Одним
из ключевых инструментов является метод fitDataset, который
позволяет обучать модели на данных, поступающих в виде потоков
(Dataset). Это особенно важно для работы с большими объёмами
данных, которые не помещаются целиком в оперативной памяти.
fitDatasetМетод fitDataset используется для обучения модели на
объектах tf.data.Dataset. В отличие от стандартного
fit, который принимает данные в виде массивов или тензоров,
fitDataset поддерживает ленивую загрузку данных,
батчирование и шифрование, что делает его подходящим для масштабируемых
решений.
Сигнатура метода:
model.fitDataset(dataset, {
epochs: number,
callbacks?: tf.Callback[],
validationData?: tf.data.Dataset,
verbose?: number
});
Основные параметры:
dataset — объект tf.data.Dataset, который
генерирует пары [features, labels].epochs — количество проходов по всему датасету.callbacks — массив обратных вызовов для мониторинга
процесса обучения.validationData — объект tf.data.Dataset
для валидации модели.verbose — уровень детализации вывода: 0 —
без вывода, 1 — прогресс-бар, 2 — строка
прогресса по эпохам.Объект Dataset можно создавать из массивов, тензоров,
генераторов или файлов. Пример создания из массивов:
const xs = tf.tensor2d([[1], [2], [3], [4]]);
const ys = tf.tensor2d([[1], [3], [5], [7]]);
const dataset = tf.data.array([{xs: xs, ys: ys}])
.batch(2)
.shuffle(4);
Здесь batch(2) объединяет данные в пакеты по 2 элемента,
а shuffle(4) перемешивает их с буфером размером 4. При
работе с большими данными важно использовать shuffle,
чтобы избежать корреляции между соседними примерами и ускорить
обучение.
Для динамической генерации данных удобно применять генераторы:
function* dataGenerator() {
for (let i = 0; i < 100; i++) {
const x = tf.tensor2d([i], [1, 1]);
const y = tf.tensor2d([2 * i], [1, 1]);
yield {xs: x, ys: y};
}
}
const dataset = tf.data.generator(dataGenerator).batch(5);
Генераторы позволяют создавать потоковые данные, что полезно при работе с потоками изображений, видео или большими CSV-файлами.
Создание простой модели для линейной регрессии:
const model = tf.sequential();
model.add(tf.layers.dense({units: 1, inputShape: [1]}));
model.compile({
optimizer: tf.train.sgd(0.01),
loss: 'meanSquaredError'
});
const dataset = tf.data.generator(dataGenerator).batch(5);
await model.fitDataset(dataset, {
epochs: 10,
verbose: 1
});
В данном примере:
sgd) с
небольшой скоростью обучения.fitDatasetВалидационный датасет передаётся через параметр
validationData. Пример:
const valDataset = tf.data.generator(validationGenerator).batch(5);
await model.fitDataset(dataset, {
epochs: 20,
validationData: valDataset,
verbose: 1
});
Валидация на лету позволяет контролировать переобучение и корректировать гиперпараметры в процессе обучения.
TensorFlow.js поддерживает обратные вызовы для мониторинга и управления процессом обучения:
const callbacks = [
tf.callbacks.earlyStopping({monitor: 'val_loss', patience: 3}),
tf.callbacks.modelCheckpoint({filepath: 'model.json', saveWeightsOnly: true})
];
await model.fitDataset(dataset, {
epochs: 50,
validationData: valDataset,
callbacks: callbacks
});
earlyStopping останавливает обучение при отсутствии
улучшений.modelCheckpoint сохраняет веса модели на диск, что
критично при долгих обучениях с потоковыми данными.fitDatasetfitDataset является ключевым инструментом для работы с
реальными, динамически изменяющимися данными и незаменим при обучении
моделей на больших объемах информации в TensorFlow.js.