Сэмплинг и reparameterization trick

Основы сэмплинга

Сэмплинг (sampling) является ключевым компонентом при работе с вероятностными моделями и вариационными автокодировщиками (VAE). В контексте Keras.js, который позволяет запускать модели Keras непосредственно в браузере с использованием JavaScript, сэмплинг часто применяется для генерации новых данных на основе распределения латентного пространства.

Принцип работы сэмплинга:

  1. Модель обучается предсказывать параметры распределения, чаще всего среднего значения () и стандартного отклонения (), для каждого объекта в латентном пространстве.
  2. Чтобы сгенерировать новый объект, необходимо извлечь случайный вектор (z), который соответствует этому распределению: [ z = + , (0, 1)]

В Keras.js это реализуется с использованием обычных массивов JavaScript и методов генерации случайных чисел, таких как Math.random() или более специализированных библиотек для работы с нормальным распределением.

Реализация сэмплинга в JavaScript

Для практической реализации сэмплинга в Keras.js создается функция, которая получает на вход параметры () и () и возвращает случайный вектор (z) по формуле выше.

function sample(mu, sigma) {
    const epsilon = mu.map(() => gaussianRandom());
    return mu.map((m, i) => m + sigma[i] * epsilon[i]);
}

function gaussianRandom() {
    let u = 0, v = 0;
    while(u === 0) u = Math.random();
    while(v === 0) v = Math.random();
    return Math.sqrt(-2.0 * Math.log(u)) * Math.cos(2.0 * Math.PI * v);
}

Ключевой момент здесь — генерация нормально распределенной случайной величины epsilon через метод Бокса-Мюллера. Это позволяет корректно воспроизводить вероятностные свойства латентного пространства при генерации новых образцов.

Reparameterization trick

Reparameterization trick решает фундаментальную проблему: случайная выборка из распределения внутри нейронной сети не позволяет корректно вычислять градиенты для обратного распространения ошибки. Без этого трюка обучение вероятностных моделей, таких как VAE, невозможно через стандартный метод градиентного спуска.

Суть метода: Случайная величина (z (, ^2)) представляется как детерминированная функция параметров сети и вспомогательной случайной величины ((0, 1)): [ z = + ]

Это позволяет:

  • Сделать процесс генерации дифференцируемым.
  • Использовать стандартный backpropagation для обновления параметров () и ().

Применение в Keras.js

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

class SamplingLayer extends KerasJS.layers.Layer {
    constructor(config) {
        super(config);
    }

    call(inputs) {
        const [mu, logVar] = inputs;
        const sigma = logVar.map(v => Math.exp(0.5 * v));
        const z = sample(mu, sigma);
        return z;
    }
}

Здесь logVar — это логарифм дисперсии, что обеспечивает численную стабильность, а sample(mu, sigma) — функция, описанная ранее. Такой подход полностью совместим с Keras.js и позволяет выполнять генерацию с поддержкой градиентов для обучения моделей в браузере или Node.js.

Практические рекомендации

  • Использование логарифма дисперсии: хранение log(sigma^2) вместо sigma напрямую уменьшает риск переполнения и обеспечивает стабильность вычислений.
  • Пакетный сэмплинг: для эффективной генерации нескольких примеров одновременно следует использовать векторизованные операции, минимизируя количество циклов JavaScript.
  • Случайность и воспроизводимость: при необходимости воспроизводимости экспериментов можно использовать фиксированные seed для генератора случайных чисел.

Заключение концепции

Сэмплинг и reparameterization trick образуют основу работы с вероятностными латентными моделями. В Keras.js эти концепции реализуются через JavaScript-функции генерации нормальных случайных величин и слои, которые преобразуют параметры распределения в случайные образцы с возможностью дифференцирования. Такой подход позволяет переносить обучение и генерацию из Python/Keras непосредственно в браузер, сохраняя корректность градиентов и вероятность сходимости модели.