Агрегация обновлений модели — ключевой компонент в распределённом машинном обучении, при котором несколько устройств или клиентов обучают локальные копии модели, а затем их результаты объединяются в центральной модели. В контексте TensorFlow.js это позволяет создавать приложения с обучением на стороне клиента, сохраняя при этом согласованность модели.
В TensorFlow.js параметры модели представлены в виде
тензоров (tf.Tensor). Каждый слой модели
содержит набор весов (weights) и
смещений (biases). Для агрегации
обновлений необходимо аккумулировать изменения этих тензоров.
Основные подходы к представлению обновлений:
Формат данных обычно представлен как объект с массивами тензоров, например:
{
'dense/kernel': tf.tensor([...]),
'dense/bias': tf.tensor([...])
}
Среднее арифметическое — наиболее распространённый метод агрегации:
function aggregateWeights(updates) {
const keys = Object.keys(updates[0]);
const aggregated = {};
keys.forEach(key => {
const tensors = updates.map(update => update[key]);
aggregated[key] = tf.stack(tensors).mean(0);
});
return aggregated;
}
Особенности метода:
Взвешенное среднее позволяет учитывать размер локальных выборок:
function weightedAggregate(updates, weights) {
const keys = Object.keys(updates[0]);
const aggregated = {};
keys.forEach(key => {
let weightedTensors = updates.map((update, i) => update[key].mul(weights[i]));
aggregated[key] = tf.addN(weightedTensors).div(tf.scalar(tf.sum(weights)));
});
return aggregated;
}
Использование весов важно для ситуаций, когда клиенты обучаются на выборках разного объёма.
Агрегация должна учитывать порядок и консистентность состояния модели. В TensorFlow.js часто используется подход:
tf.Model.Пример обновления модели после агрегации:
async function applyAggregatedWeights(model, aggregated) {
const weightNames = model.weights.map(w => w.name);
const newWeights = weightNames.map(name => aggregated[name]);
await model.setWeights(newWeights);
}
Важно использовать асинхронные операции, так как
работа с тензорами требует освобождения памяти с помощью
tf.dispose().
При агрегации большого количества обновлений необходимо:
tf.tidy() для автоматического освобождения
промежуточных тензоров.Float32Array → Float16) при передаче
обновлений по сети.Пример безопасной агрегации с tf.tidy:
function safeAggregate(updates) {
return tf.tidy(() => {
const keys = Object.keys(updates[0]);
const aggregated = {};
keys.forEach(key => {
const tensors = updates.map(update => update[key]);
aggregated[key] = tf.stack(tensors).mean(0);
});
return aggregated;
});
}
Эти методы особенно полезны для масштабных приложений с тысячами клиентов, где прямое усреднение может приводить к нестабильности.
TensorFlow.js позволяет интегрировать агрегацию обновлений в веб-клиенты:
Web Workers для выполнения обучения без
блокировки интерфейса.fetch или
WebSocket.Пример отправки обновлений на сервер:
async function sendUpdates(updates) {
const serialized = {};
Object.keys(updates).forEach(key => {
serialized[key] = Array.from(updates[key].dataSync());
});
await fetch('/aggregate', {
method: 'POST',
headers: { 'Content-Type': 'application/json' },
body: JSON.stringify(serialized)
});
}
await при работе с весами модели.Эти практики обеспечивают стабильное обучение распределённых моделей на стороне клиента с использованием TensorFlow.js, минимизируют риски утечек памяти и сетевых задержек, а также повышают точность агрегированных моделей.