Fine-tuning — процесс дообучения уже предварительно обученной модели на новом наборе данных. В отличие от замораживания слоёв и обучения только верхних слоёв, fine-tuning всей сети позволяет корректировать веса всех слоёв, что особенно важно при работе с сильно отличающимися от исходных данных задачами.
Для начала необходимо загрузить предварительно обученную модель. В
TensorFlow.js это делается через tf.loadLayersModel или
загрузку моделей из TensorFlow Hub:
import * as tf from '@tensorflow/tfjs';
const modelUrl = 'https://example.com/model.json';
const pretrainedModel = await tf.loadLayersModel(modelUrl);
После загрузки важно определить, какие слои будут дообучаться. Для fine-tuning всей сети замораживание слоёв не применяется:
pretrainedModel.layers.forEach(layer => {
layer.trainable = true;
});
Данные должны быть приведены к формату, совместимому с моделью.
Например, для изображений используется tf.data API:
const imageSize = 224;
const preprocessImage = (image) => {
return tf.tidy(() => {
let tensor = tf.browser.fromPixels(image)
.resizeNearestNeighbor([imageSize, imageSize])
.toFloat();
return tensor.div(255.0).expandDims();
});
};
Для больших наборов данных рекомендуется использовать
tf.data.generator или tf.data.array для
эффективной загрузки и пакетирования данных:
const dataset = tf.data.generator(function* () {
for (let img of images) {
yield { xs: preprocessImage(img), ys: tf.tensor(labels.shift()) };
}
}).batch(32);
При fine-tuning всей сети важно выбрать оптимизатор с небольшим коэффициентом обучения, чтобы не разрушить уже обученные веса:
const optimizer = tf.train.adam(0.0001); // Малый learning rate
pretrainedModel.compile({
optimizer: optimizer,
loss: 'categoricalCrossentropy',
metrics: ['accuracy']
});
Ключевой момент: слишком высокий learning rate приведет к переобучению и разрушению полезных признаков, выученных предварительно.
Процесс обучения полностью аналогичен обычному обучению модели, только теперь обучаются все слои:
await pretrainedModel.fitDataset(dataset, {
epochs: 10,
callbacks: {
onEpochEnd: (epoch, logs) => {
console.log(`Эпоха ${epoch + 1}: потеря = ${logs.loss}, точность = ${logs.acc}`);
}
}
});
Можно использовать fit с массивами xs и
ys, если данные помещаются в память.
Fine-tuning всей сети несет риск переобучения, особенно при небольшом объеме данных. Для борьбы с этим применяются:
flip, rotation,
zoom) увеличивают разнообразие данных.const earlyStopping = tf.callbacks.earlyStopping({
monitor: 'val_loss',
patience: 3
});
После завершения fine-tuning модель сохраняется для последующего использования:
await pretrainedModel.save('localstorage://fine-tuned-model');
// или для скачивания
await pretrainedModel.save('downloads://fine-tuned-model');
tf.tidy для удаления временных тензоров, иначе возможны
утечки памяти.Fine-tuning всей сети позволяет адаптировать мощные предварительно обученные архитектуры к уникальным задачам, повышая точность и обеспечивая использование сложных признаков, выученных на больших датасетах.