Вариационный автоэнкодер VAE: теория и реализация

Основы вариационного автоэнкодера

Вариационный автоэнкодер (Variational Autoencoder, VAE) — это тип генеративной модели, которая объединяет идеи автоэнкодеров и вероятностного моделирования. В отличие от классического автоэнкодера, который сжимает входные данные в детерминированное скрытое представление, VAE использует вероятностное кодирование. Это позволяет модели генерировать новые данные, схожие с обучающим набором, и работать с непрерывными латентными пространствами.

Ключевые компоненты VAE:

  • Энкодер (q_(z|x)) — преобразует входные данные (x) в параметры распределения латентного вектора (z), обычно среднее () и стандартное отклонение () для гауссовского распределения.

  • Сэмплинг — процесс генерации латентного вектора (z) из распределения, заданного энкодером. Используется техника «reparameterization trick» для возможности обратного распространения градиента.

  • Декодер (p_(x|z)) — восстанавливает исходные данные из латентного представления, формируя распределение вероятностей по выходным данным.

  • Функция потерь VAE сочетает два компонента:

    1. Реконструктивная потеря — измеряет, насколько хорошо декодер восстанавливает входные данные. Обычно используется среднеквадратичная ошибка (MSE) или бинарная кросс-энтропия.
    2. KL-дивергенция — штрафует отклонение распределения латентного вектора от стандартного нормального распределения (N(0, 1)), обеспечивая регуляризацию латентного пространства.

Архитектура VAE в Keras.js

Keras.js позволяет использовать модели, обученные в Python Keras, прямо в браузере на JavaScript. Для реализации VAE необходимы три основных блока: энкодер, сэмплер и декодер. Важно корректно обработать сэмплинг, так как он включает стохастический элемент.

Энкодер:

const encoderInput = tf.input({shape: [inputDim]});
const hiddenLayer = tf.layers.dense({units: 256, activation: 'relu'}).apply(encoderInput);

const zMean = tf.layers.dense({units: latentDim}).apply(hiddenLayer);
const zLogVar = tf.layers.dense({units: latentDim}).apply(hiddenLayer);

Сэмплинг с reparameterization trick:

function sampling([zMean, zLogVar]) {
    const epsilon = tf.randomNormal([latentDim]);
    return tf.add(zMean, tf.mul(tf.exp(tf.mul(0.5, zLogVar)), epsilon));
}
const z = sampling([zMean, zLogVar]);

Декодер:

const decoderInput = tf.input({shape: [latentDim]});
let x = tf.layers.dense({units: 256, activation: 'relu'}).apply(decoderInput);
const decoderOutput = tf.layers.dense({units: inputDim, activation: 'sigmoid'}).apply(x);

Обучение модели

Для обучения необходимо объединить энкодер и декодер в один VAE и задать комбинированную функцию потерь. В JavaScript Keras.js возможно реализовать кастомную функцию потерь через tf.losses и tf.tidy для управления памятью.

Функция потерь VAE:

function vaeLoss(xTrue, xPred) {
    const reconLoss = tf.metrics.meanSquaredError(xTrue, xPred);
    const klLoss = tf.mul(-0.5, tf.mean(tf.add(1, zLogVar, tf.neg(tf.square(zMean)), tf.neg(tf.exp(zLogVar)))));
    return tf.add(reconLoss, klLoss);
}

Модель компилируется стандартным образом:

const vae = tf.model({inputs: encoderInput, outputs: decoderOutput});
vae.compile({optimizer: tf.train.adam(0.001), loss: vaeLoss});

Генерация новых данных

После обучения VAE латентное пространство становится гладким и непрерывным. Для генерации новых объектов необходимо просто сэмплировать вектор (z) из стандартного нормального распределения и пропустить его через декодер:

const zSample = tf.randomNormal([1, latentDim]);
const generatedData = decoder.predict(zSample);

Важные особенности и рекомендации

  • Нормализация данных. Для стабильного обучения рекомендуется масштабировать входные данные в диапазон [0,1].
  • Размер латентного пространства. Слишком маленькое пространство ограничивает способность модели генерировать разнообразные объекты, слишком большое — усложняет обучение и ухудшает регуляризацию.
  • Баланс потерь. Иногда полезно масштабировать KL-дивергенцию, чтобы управлять компромиссом между качеством реконструкции и гладкостью латентного пространства.
  • Память и производительность. В браузерной реализации важно использовать tf.tidy для освобождения промежуточных тензоров и предотвращения утечек памяти.

Применение VAE в веб-приложениях

Keras.js позволяет интегрировать VAE прямо на фронтенд, что открывает возможности:

  • Генерация изображений, аудио или текстов без обращения к серверу.
  • Анимация латентного пространства с интерактивными слайдерами.
  • Реализация фильтров и стилизаций в реальном времени.

Использование вариационного автоэнкодера в браузере делает возможным эксперименты с генеративными моделями на клиентской стороне, минимизируя задержки и нагрузку на сервер.