Агрегация обновлений модели

Агрегация обновлений модели — ключевой компонент в распределённом машинном обучении, при котором несколько устройств или клиентов обучают локальные копии модели, а затем их результаты объединяются в центральной модели. В контексте 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 часто используется подход:

  1. Сохранение центральной копии модели в формате tf.Model.
  2. Рассылка текущих весов клиентам для локального обучения.
  3. Получение обновлений (дельт или градиентов).
  4. Агрегация с использованием среднего или взвешенного метода.
  5. Обновление центральной модели.

Пример обновления модели после агрегации:

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() для автоматического освобождения промежуточных тензоров.
  • Стараться агрегировать тензоры поэтапно, чтобы избежать переполнения памяти.
  • Применять сжатие данных (например, Float32ArrayFloat16) при передаче обновлений по сети.

Пример безопасной агрегации с 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;
  });
}

Расширенные стратегии

  • FedAvg с адаптивным шагом: позволяет учитывать разные скорости обучения на клиентах.
  • Агрегация с отсечением выбросов: исключение аномальных обновлений с использованием медианы или квантилей.
  • Компрессия градиентов: уменьшение размера передаваемых данных, сохраняя точность модели.

Эти методы особенно полезны для масштабных приложений с тысячами клиентов, где прямое усреднение может приводить к нестабильности.

Интеграция с веб-приложениями

TensorFlow.js позволяет интегрировать агрегацию обновлений в веб-клиенты:

  • Использование Web Workers для выполнения обучения без блокировки интерфейса.
  • Передача обновлений на сервер через fetch или WebSocket.
  • Локальное кэширование весов в IndexedDB для повторного использования и повышения отказоустойчивости.

Пример отправки обновлений на сервер:

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, минимизируют риски утечек памяти и сетевых задержек, а также повышают точность агрегированных моделей.