Подгонка модели — центральный этап обучения нейронной сети в
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.Пример создания 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).tf.callbacks.tensorBoard).Обратные вызовы критически важны при работе с большими данными или при длительном обучении, позволяя контролировать процесс и предотвращать переобучение.
fit.fitDataset.fitDataset всегда проверять корректность
преобразований через map и batch, чтобы
избежать ошибок при несовпадении форм тензоров.callbacks позволяет динамически управлять
обучением и сокращать время до достижения оптимальной модели.Методы fit и fitDataset обеспечивают
гибкость обучения нейронных сетей в TensorFlow.js, позволяя эффективно
обрабатывать как малые, так и масштабные наборы данных, интегрировать
сложные пайплайны обработки и управлять процессом обучения через
обратные вызовы.