Нейронные сети Генеративные модели: VAE, GAN, диффузия
0%

Генеративные модели: VAE, GAN, диффузия

Генеративные модели: VAE, GAN, диффузия

Все архитектуры, разобранные до сих пор — MLP, свёрточные, рекуррентные, трансформеры — мы обучали отвечать на вопрос про объект: какой это класс, какое число, какой следующий токен. Теперь задача переворачивается: научиться производить объекты, которых в обучающей выборке не было, но которые выглядят так, будто могли бы там быть.

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

1. Дискриминативное и генеративное: в чём реальная разница

Дискриминативная модель приближает условное распределение $p(y \mid x)$: дан объект — назови метку. Генеративная приближает само распределение данных $p(x)$ (или совместное $p(x, y)$, или условное $p(x \mid c)$ — «нарисуй кота в скафандре»).

Разница не косметическая. Чтобы отличить кота от собаки, достаточно выучить границу — модель может опираться на текстуру шерсти и вообще не «понимать», как устроена морда. Чтобы нарисовать кота, нужно выучить всё: анатомию, освещение, перспективу, согласованность левого и правого глаза. Поэтому у генеративных моделей появилось второе назначение — обучение представлений без разметки; подробнее об этом в статье про эмбеддинги.

Формально всё сводится к одной строчке. У нас есть выборка $x_1, \dots, x_N$ из неизвестного $p_{\text{data}}(x)$ и параметрическое семейство $p_\theta(x)$. Нужно приблизить одно к другому:

$$\theta^\ast = \arg\min_\theta ; D\big(p_{\text{data}} ,|, p_\theta\big)$$

Дальше вся история жанра — это ответы на два вопроса: как параметризовать $p_\theta$ и какое расхождение $D$ минимизировать. VAE, GAN и диффузия дают три разных ответа, и все их особенности (размытость, mode collapse, медленный сэмплинг) выводятся именно отсюда.

Проклятие нормировочной константы

Наивная идея: пусть сеть $f_\theta(x)$ выдаёт «оценку правдоподобия», и определим

$$p_\theta(x) = \frac{e^{f_\theta(x)}}{Z_\theta}, \qquad Z_\theta = \int e^{f_\theta(x)},dx$$

Это корректное распределение (энергетическая модель), но для картинки $256\times256\times3$ интеграл берётся по $\mathbb{R}^{196608}$. Посчитать $Z_\theta$ невозможно, а без него нельзя ни оценить правдоподобие, ни взять его градиент.

Весь дизайн генеративных моделей — это способы обойти $Z_\theta$. Варианты:

  • сделать $Z_\theta$ тривиальным — разложить $p(x)$ в произведение нормированных одномерных условных (авторегрессия) или строить $p_\theta$ обратимым преобразованием простого распределения (normalizing flows);
  • не считать $p_\theta$ точно, а оптимизировать нижнюю границу (VAE, диффузия);
  • вообще отказаться от плотности и уметь только сэмплировать (GAN).

Заметьте: авторегрессионные модели здесь тоже есть — большая языковая модель из статьи про LLM это полноценная генеративная модель над текстом с точным правдоподобием. В этой статье фокус на трёх семействах, доминирующих в непрерывных доменах: изображения, аудио, видео, молекулы.

2. Латентные переменные и вариационная нижняя граница

2.1 Интуиция

Картинка лица — это 150 тысяч чисел, но реальных степеней свободы куда меньше: поза, возраст, освещение, форма носа, эмоция. Гипотеза многообразия говорит, что данные лежат на маломерной поверхности внутри огромного пространства.

Отсюда конструкция: введём латентную переменную $z \in \mathbb{R}^d$ (скажем, $d = 128$) с простым априорным распределением $p(z) = \mathcal{N}(0, I)$ и обучим декодер $p_\theta(x \mid z)$, разворачивающий короткий код в объект:

$$p_\theta(x) = \int p_\theta(x \mid z), p(z), dz$$

Сэмплировать просто: взять $z \sim \mathcal{N}(0, I)$, прогнать через декодер. Проблема в обучении — интеграл по всем $z$ неберущийся. Метод Монте-Карло «насэмплируем $z$ и усредним» бесполезен: при $d = 128$ доля кодов, релевантных конкретной картинке $x$, астрономически мала, и оценка почти всегда равна нулю.

2.2 Вывод ELBO

Нужен способ сэмплировать те $z$, которые могли породить именно этот $x$, то есть из апостериорного $p_\theta(z \mid x)$. Оно тоже неберущееся, поэтому вводим вариационное приближение $q_\phi(z \mid x)$ — второй сети-энкодера, выдающей параметры гауссианы. Дальше — стандартный приём:

$$\log p_\theta(x) = \log \int p_\theta(x \mid z) p(z) dz = \log \mathbb{E}_ {q_\phi(z\mid x)}!\left[\frac{p_\theta(x \mid z) p(z)}{q_\phi(z \mid x)}\right]$$

По неравенству Йенсена ($\log$ вогнут) логарифм можно завести под матожидание:

$$\log p_\theta(x) ;\ge; \mathbb{E}_ {q_\phi(z\mid x)}\big[\log p_\theta(x \mid z)\big] ;-; D_{KL}\big(q_\phi(z\mid x),|,p(z)\big) ;=; \mathcal{L}_ {\text{ELBO}}(x)$$

Это ELBO (Evidence Lower BOund). Точная величина зазора известна:

$$\log p_\theta(x) - \mathcal{L}_ {\text{ELBO}}(x) = D_{KL}\big(q_\phi(z\mid x),|,p_\theta(z\mid x)\big) ;\ge; 0$$

Максимизируя ELBO по $\theta$ и $\phi$ одновременно, мы и подтягиваем правдоподобие, и делаем приближённый апостериор ближе к настоящему. Два слагаемых читаются инженерно:

  • реконструкция $\mathbb{E}[\log p_\theta(x\mid z)]$ — декодер должен восстановить $x$ из кода. Для гауссова декодера с фиксированной дисперсией это MSE, для бернуллиевского — бинарная кросс-энтропия;
  • регуляризатор $D_{KL}(q_\phi | p)$ — заставляет облака кодов разных объектов ложиться в общую $\mathcal{N}(0, I)$, без дыр между ними. Без него получился бы обычный автоэнкодер: реконструирует отлично, но сэмплировать из него нельзя — случайный $z$ попадёт в пустоту и даст мусор.

Для двух гауссиан KL берётся в замкнутой форме, что и делает VAE практичным:

$$D_{KL}\big(\mathcal{N}(\mu, \sigma^2 I),|,\mathcal{N}(0, I)\big) = \tfrac{1}{2}\sum_{j=1}^{d}\big(\mu_j^2 + \sigma_j^2 - \log \sigma_j^2 - 1\big)$$

2.3 Трюк репараметризации

Осталось препятствие: внутри графа стоит операция сэмплирования $z \sim q_\phi(z\mid x)$. Через случайный узел градиент не течёт — производной у «бросить кубик» нет.

Решение Кингмы и Веллинга (2013): вынести случайность из графа наружу.

$$z = \mu_\phi(x) + \sigma_\phi(x) \odot \varepsilon, \qquad \varepsilon \sim \mathcal{N}(0, I)$$

Распределение $z$ то же самое, но теперь $z$ — детерминированная дифференцируемая функция от выходов энкодера, а $\varepsilon$ — обычный вход, вроде данных.

Трюк репараметризации: наивное сэмплирование против детерминированного пути

Это один из тех приёмов, что стоит выучить отдельно от VAE: он работает везде, где нужно дифференцировать матожидание по параметрам распределения. Дискретный аналог — Gumbel-Softmax.

2.4 Реализация

import torch
import torch.nn as nn
import torch.nn.functional as F

class VAE(nn.Module):
    """VAE для MNIST: вход 784, латент 32."""

    def __init__(self, x_dim: int = 784, h: int = 512, z_dim: int = 32):
        super().__init__()
        self.enc = nn.Sequential(nn.Linear(x_dim, h), nn.ReLU(), nn.Linear(h, h), nn.ReLU())
        self.to_mu = nn.Linear(h, z_dim)
        self.to_logvar = nn.Linear(h, z_dim)   # предсказываем log σ², а не σ:
                                               # так выход не обязан быть положительным
        self.dec = nn.Sequential(
            nn.Linear(z_dim, h), nn.ReLU(),
            nn.Linear(h, h), nn.ReLU(),
            nn.Linear(h, x_dim),               # логиты Бернулли, сигмоида — внутри BCE
        )

    def encode(self, x):
        hh = self.enc(x)
        return self.to_mu(hh), self.to_logvar(hh)

    def reparameterize(self, mu, logvar):
        std = torch.exp(0.5 * logvar)
        eps = torch.randn_like(std)            # случайность — вход, а не узел графа
        return mu + std * eps

    def forward(self, x):
        mu, logvar = self.encode(x)
        z = self.reparameterize(mu, logvar)
        return self.dec(z), mu, logvar


def vae_loss(logits, x, mu, logvar, beta: float = 1.0):
    """Отрицательный ELBO, усреднённый по батчу."""
    # реконструкция: сумма по пикселям, среднее по объектам
    rec = F.binary_cross_entropy_with_logits(logits, x, reduction="none").sum(dim=1)
    # KL в замкнутой форме
    kld = -0.5 * (1 + logvar - mu.pow(2) - logvar.exp()).sum(dim=1)
    return (rec + beta * kld).mean(), rec.mean().item(), kld.mean().item()


@torch.no_grad()
def sample(model: VAE, n: int, device) -> torch.Tensor:
    """Генерация: априор -> декодер. Энкодер на инференсе не нужен вообще."""
    z = torch.randn(n, model.to_mu.out_features, device=device)
    return torch.sigmoid(model.dec(z)).view(n, 1, 28, 28)

Сложность: forward-pass энкодера и декодера, $O(1)$ проходов сети на объект — это самый дешёвый сэмплинг из всех трёх семейств. Память — как у обычного автоэнкодера плюс 2 $d$ чисел на объект.

2.5 Почему VAE размывает картинки и что с этим делают

Классическая претензия: сэмплы VAE выглядят как фотография сквозь запотевшее стекло. Причины конкретны:

  1. Форма правдоподобия. Гауссов декодер + MSE означает, что при неопределённости («усы влево или вправо?») оптимальный по MSE ответ — усреднить оба варианта. Среднее двух валидных картинок не является валидной картинкой.
  2. Прямой KL в ELBO — режим mass-covering: модель штрафуется за нулевую вероятность там, где данные есть, но почти не штрафуется за размазывание массы туда, где данных нет. GAN, наоборот, mode-seeking — отсюда его резкость и его mode collapse.
  3. Posterior collapse. Если декодер достаточно мощный (например, авторегрессионный), ему выгоднее игнорировать $z$: KL-член обнуляется при $q_\phi(z\mid x) = p(z)$, реконструкция всё равно неплоха. Латент становится мёртвым, модель вырождается в безусловную.

Практические противоядия:

  • KL annealing / warm-up: линейно поднимать $\beta$ от 0 до 1 за первые эпохи, чтобы латент успел стать полезным до того, как включится штраф;
  • free bits: не штрафовать KL, пока он ниже порога $\lambda$ на измерение — kld_j = torch.clamp(kld_j, min=lambda_);
  • $\beta$-VAE ($\beta > 1$) — жертвуем реконструкцией ради факторизованного, «распутанного» латента (Higgins et al., 2017);
  • VQ-VAE (van den Oord et al., 2017) — заменить непрерывный латент дискретным словарём, а априор над кодами выучить отдельной авторегрессионной моделью. Это убирает и размытость, и posterior collapse.

Последний пункт важнее, чем кажется: сегодня VAE почти никогда не используют как генератор напрямую. Его используют как обучаемый компрессор — тот самый автоэнкодер, что превращает изображение $512\times512$ в тензор $64\times64\times4$, внутри которого работает Stable Diffusion. Об этом в разделе 5.

3. GAN: обучение через соперничество

3.1 Идея

Гудфеллоу и соавторы (2014) предложили обойти правдоподобие целиком. Пусть есть:

  • генератор $G_\theta: z \mapsto x$, превращающий шум в объект;
  • дискриминатор $D_\psi: x \mapsto [0,1]$, оценивающий вероятность, что объект настоящий.

Они играют в минимаксную игру:

$$\min_\theta \max_\psi ; V(D, G) = \mathbb{E}_ {x \sim p_{\text{data}}}[\log D(x)]

  • \mathbb{E}_ {z \sim p(z)}[\log (1 - D(G(z)))]$$

Метрика качества здесь не задана руками (никакого MSE), а выучивается второй сетью — и подстраивается по мере роста генератора. Именно поэтому GAN даёт резкость: дискриминатор мгновенно ловит размытость как признак подделки.

3.2 Что на самом деле минимизируется

При фиксированном $G$ оптимальный дискриминатор выписывается явно:

$$D^\ast (x) = \frac{p_{\text{data}}(x)}{p_{\text{data}}(x) + p_g(x)}$$

Подставив его в $V$, получаем

$$V(D^\ast , G) = 2,D_{JS}\big(p_{\text{data}} ,|, p_g\big) - \log 4$$

То есть исходный GAN минимизирует дивергенцию Йенсена–Шеннона. Красиво — и здесь же спрятан главный дефект. Если носители $p_{\text{data}}$ и $p_g$ не пересекаются (а для маломерных многообразий в пространстве картинок это почти всегда так), $D_{JS} = \log 2$ константа, её градиент равен нулю. Идеальный дискриминатор не даёт генератору никакого сигнала.

Отсюда два практических следствия, которые видит каждый, кто обучал GAN:

  • non-saturating loss: вместо минимизации $\log(1 - D(G(z)))$ генератор максимизирует $\log D(G(z))$. Математически тот же оптимум, но градиент на старте (когда $D(G(z)) \approx 0$) на порядки больше;
  • дискриминатор нельзя переобучать. Слишком сильный $D$ убивает обучение.

Ключевая деталь, на которой спотыкаются все: detach() на шаге 1 и свежий сэмпл $z’$ на шаге 2. Без detach градиент от лосса дискриминатора потечёт в генератор и будет толкать его в противоположную сторону.

3.3 Реализация одного шага

import torch
import torch.nn as nn

bce = nn.BCEWithLogitsLoss()

def gan_step(G, D, opt_g, opt_d, real, z_dim, device):
    """Один шаг обучения DCGAN-подобной пары. real: (B, C, H, W) в [-1, 1]."""
    B = real.size(0)
    ones = torch.ones(B, 1, device=device)
    zeros = torch.zeros(B, 1, device=device)

    # --- 1. дискриминатор ---
    opt_d.zero_grad(set_to_none=True)
    z = torch.randn(B, z_dim, device=device)
    fake = G(z)
    # односторонний label smoothing: цель для реальных 0.9, не 1.0 —
    # мешает D становиться переуверенным и обнулять градиент генератора
    loss_d = bce(D(real), ones * 0.9) + bce(D(fake.detach()), zeros)
    loss_d.backward()
    opt_d.step()

    # --- 2. генератор ---
    opt_g.zero_grad(set_to_none=True)
    z = torch.randn(B, z_dim, device=device)          # новый шум, не переиспользуем
    # non-saturating: хотим, чтобы D назвал фейки настоящими
    loss_g = bce(D(G(z)), ones)
    loss_g.backward()
    opt_g.step()

    return loss_d.item(), loss_g.item()

Диагностика: если loss_d уходит в 0, а loss_g растёт — дискриминатор победил, градиента нет. Если оба лосса болтаются около $\ln 2 \approx 0.69$ и картинки улучшаются — всё идёт правильно. Значение лосса GAN не является метрикой качества: оно измеряет баланс игры, а не близость распределений.

3.4 Режимы отказа и стабилизация

Работающий на практике набор мер:

Приём Что чинит Источник
WGAN-GP — расстояние Вассерштейна + штраф на норму градиента нулевые градиенты при непересекающихся носителях Gulrajani et al., 2017
Spectral normalization — деление весов $D$ на старшее сингулярное число ограничивает липшицевость $D$, дёшево и надёжно Miyato et al., 2018
TTUR — разные lr для $G$ и $D$ (например 1e-4 / 4e-4) осцилляции Heusel et al., 2017
EMA весов генератора для инференса шум последних шагов стандарт в StyleGAN
ADA — адаптивные дифференцируемые аугментации переобучение $D$ на малых датасетах Karras et al., 2020
minibatch stddev — фича «разнообразие батча» на вход $D$ mode collapse ProGAN/StyleGAN

Вершина линии — StyleGAN2/3 (Karras et al., 2019): отображение $z \to w$ отдельной MLP, модуляция весов свёрток стилем на каждом уровне разрешения, отсюда — управляемая иерархия признаков (поза на низких разрешениях, текстура кожи на высоких). До сих пор непревзойдён по скорости: одна прямая подача сети на изображение против десятков шагов у диффузии.

4. Диффузия: разрушить и научиться восстанавливать

4.1 Интуиция

Обе предыдущие модели пытались одним махом прыгнуть из шума в изображение. Диффузия разбивает этот прыжок на сотни крошечных шагов, каждый из которых — почти тривиальная задача «убери чуть-чуть шума». Обучающий сигнал берётся бесплатно: чтобы получить пару (зашумлённое, чистое), достаточно взять реальную картинку и подмешать шум.

Это меняет всё: нет второй сети, нет игры, нет коллапса мод. Есть обычная регрессия с MSE — самая устойчивая вещь в глубоком обучении.

4.2 Прямой процесс

Задаём расписание шума $\beta_1, \dots, \beta_T$ (например, линейно от $10^{-4}$ до $0.02$ при $T = 1000$) и марковскую цепь:

$$q(x_t \mid x_{t-1}) = \mathcal{N}\big(x_t; \sqrt{1 - \beta_t}, x_{t-1},; \beta_t I\big)$$

Ключевое свойство — композиция гауссиан остаётся гауссианой. Обозначив $\alpha_t = 1 - \beta_t$ и $\bar\alpha_t = \prod_{s \le t} \alpha_s$, получаем замкнутую форму для любого шага сразу:

$$q(x_t \mid x_0) = \mathcal{N}\big(x_t; \sqrt{\bar\alpha_t}, x_0,; (1 - \bar\alpha_t) I\big) \quad\Longleftrightarrow\quad x_t = \sqrt{\bar\alpha_t}, x_0 + \sqrt{1 - \bar\alpha_t}, \varepsilon$$

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

Диффузия: прямой процесс зашумления и обратный процесс расшумления

4.3 Обратный процесс и функция потерь

Обратные переходы $q(x_{t-1} \mid x_t)$ неизвестны, но при малом $\beta_t$ они тоже приблизительно гауссовы — этим и пользуемся, параметризуя их сетью:

$$p_\theta(x_{t-1} \mid x_t) = \mathcal{N}\big(x_{t-1}; \mu_\theta(x_t, t), \Sigma_t\big)$$

Вывод ELBO для этой цепи даёт сумму KL-членов между $q(x_{t-1} \mid x_t, x_0)$ (это уже вычислимо в замкнутой форме) и $p_\theta$. Хо и соавторы (DDPM, 2020) показали, что после перепараметризации $\mu_\theta$ через предсказание шума и отбрасывания весовых коэффициентов остаётся поразительно простой лосс:

$$\mathcal{L}_ {\text{simple}} = \mathbb{E}_ {x_0,, \varepsilon,, t} \Big[\big| \varepsilon - \varepsilon_\theta\big(\sqrt{\bar\alpha_t} x_0 + \sqrt{1-\bar\alpha_t}\varepsilon,; t\big) \big|^2\Big]$$

Сеть просто угадывает, какой шум был подмешан. MSE. Всё.

Есть и вторая, эквивалентная оптика: предсказание шума с точностью до множителя совпадает с оценкой $\nabla_x \log p_t(x)$ — score-функции зашумлённого распределения (Song & Ermon, 2019). Score не зависит от нормировочной константы $Z_\theta$ (градиент логарифма её убивает) — вот как диффузия обходит проклятие из раздела 1. Обобщение обоих взглядов через СДУ — Song et al., 2021.

Асимметрия на схеме — главный trade-off диффузии: обучение эффективно и распараллеливается по $t$, а генерация требует $N$ последовательных проходов сети. Один сэмпл стоит $O(N)$ форвардов; у DDPM $N = 1000$, у DDIM хватает 20–50, у дистиллированных моделей — 1–4.

4.4 Минимальная рабочая реализация

import torch
import torch.nn.functional as F

class Diffusion:
    """DDPM: расписание, обучающий шаг и два сэмплера."""

    def __init__(self, T: int = 1000, device="cuda"):
        self.T = T
        self.betas = torch.linspace(1e-4, 0.02, T, device=device)
        self.alphas = 1.0 - self.betas
        self.abar = torch.cumprod(self.alphas, dim=0)          # ᾱ_t

    def loss(self, model, x0, cond=None):
        """L_simple. x0 нормализован в [-1, 1]."""
        B = x0.size(0)
        t = torch.randint(0, self.T, (B,), device=x0.device)   # свой t на каждый объект
        eps = torch.randn_like(x0)
        a = self.abar[t].view(-1, 1, 1, 1)
        x_t = a.sqrt() * x0 + (1 - a).sqrt() * eps             # прямой процесс за один шаг
        eps_hat = model(x_t, t, cond)
        return F.mse_loss(eps_hat, eps)

    @torch.no_grad()
    def sample_ddpm(self, model, shape, cond=None, guidance: float = 0.0):
        x = torch.randn(shape, device=self.betas.device)
        for i in reversed(range(self.T)):
            t = torch.full((shape[0],), i, device=x.device, dtype=torch.long)
            eps = model(x, t, cond)
            if guidance > 0:                                   # classifier-free guidance
                eps_uncond = model(x, t, None)
                eps = eps_uncond + guidance * (eps - eps_uncond)
            a, ab = self.alphas[i], self.abar[i]
            # среднее обратного перехода
            mean = (x - (1 - a) / (1 - ab).sqrt() * eps) / a.sqrt()
            if i > 0:
                x = mean + self.betas[i].sqrt() * torch.randn_like(x)
            else:
                x = mean                                        # на последнем шаге без шума
        return x

    @torch.no_grad()
    def sample_ddim(self, model, shape, steps: int = 50, eta: float = 0.0, cond=None):
        """DDIM: подсеть шагов, eta=0 -> детерминированное отображение шум->картинка."""
        seq = torch.linspace(self.T - 1, 0, steps, device=self.betas.device).long()
        x = torch.randn(shape, device=self.betas.device)
        for i, ti in enumerate(seq):
            t = torch.full((shape[0],), ti, device=x.device, dtype=torch.long)
            eps = model(x, t, cond)
            ab_t = self.abar[ti]
            ab_prev = self.abar[seq[i + 1]] if i + 1 < steps else torch.tensor(1.0)
            # шаг 1: восстановить оценку чистой картинки
            x0_hat = (x - (1 - ab_t).sqrt() * eps) / ab_t.sqrt()
            x0_hat = x0_hat.clamp(-1, 1)                        # важно: без клипа копится дрейф
            sigma = eta * ((1 - ab_prev) / (1 - ab_t) * (1 - ab_t / ab_prev)).sqrt()
            # шаг 2: снова зашумить, но уже до уровня t-1
            x = ab_prev.sqrt() * x0_hat + (1 - ab_prev - sigma ** 2).clamp(min=0).sqrt() * eps
            if eta > 0 and i + 1 < steps:
                x = x + sigma * torch.randn_like(x)
        return x

Про архитектуру $\varepsilon_\theta$: это U-Net со skip-соединениями (см. статью о CNN), в который на каждом уровне подмешивается синусоидальный эмбеддинг шага $t$ и, через cross-attention, эмбеддинг условия $c$ — текстового промпта. Современные модели (DiT, SD3, Flux) заменяют U-Net трансформером над патчами латента, ровно как в статье про трансформеры.

4.5 Три идеи, сделавшие диффузию практичной

1. Classifier-free guidance (Ho & Salimans, 2022). При обучении в 10% случаев условие $c$ заменяют на пустое — одна сеть учится быть и условной, и безусловной. На инференсе предсказания экстраполируют:

$$\tilde\varepsilon = \varepsilon_\theta(x_t, \varnothing) + w \cdot \big(\varepsilon_\theta(x_t, c) - \varepsilon_\theta(x_t, \varnothing)\big)$$

При $w = 1$ обычная условная генерация; $w \approx 7$ — типичный выбор для text-to-image. Это прямая ручка trade-off «соответствие промпту против разнообразия»: чем выше $w$, тем точнее по тексту, тем однообразнее и пересыщеннее картинки. Цена — два прохода сети на шаг.

2. Latent diffusion (Rombach et al., 2022). Диффузия в пространстве пикселей $512^2$ разорительна. Решение: обучить VAE, сжимающий изображение в 8 раз по каждой оси ($512\times512\times3 \to 64\times64\times4$, в 48 раз меньше элементов), и гонять диффузию в латенте. Так VAE вернулся — не как генератор, а как перцептивный компрессор. Это архитектура Stable Diffusion.

3. Дистилляция шагов. Consistency models, progressive distillation, adversarial diffusion distillation сжимают 50 шагов в 1–4, обучая студента повторять результат многошагового учителя. Здесь диффузия смыкается с GAN: у лучших дистилляторов в лоссе стоит дискриминатор.

Отдельно стоит упомянуть flow matching / rectified flow — современную переформулировку, где вместо марковской цепи учат векторное поле, переносящее шум в данные вдоль (почти) прямых траекторий. Математика проще, траектории прямее, шагов нужно меньше; на этом построены Stable Diffusion 3 и Flux (Lipman et al., 2022).

5. Сравнение: что выбрать под задачу

Свойство VAE GAN Диффузия
Правдоподобие нижняя граница (ELBO) недоступно нижняя граница, обычно лучшая
Стоимость сэмпла 1 форвард 1 форвард $N$ форвардов ($N$ = 1…1000)
Стабильность обучения высокая низкая, требует няньки высокая (обычный MSE)
Покрытие мод (recall) высокое низкое, склонность к коллапсу высокое
Резкость (precision) низкая очень высокая очень высокая
Латент для редактирования есть, гладкий есть ($W$-пространство StyleGAN) нет прямого; инверсия дорогая
Масштабирование на текст плохо средне отлично (cross-attention + CFG)

Метрики: чем вообще измерять «красиво»

  • FID (Heusel et al., 2017) — расстояние Фреше между гауссианами, подогнанными к активациям InceptionV3 на реальных и сгенерированных выборках. Де-факто стандарт, но: сильно зависит от числа сэмплов (нужно $\ge$ 10k, иначе смещён), от библиотеки ресайза, и вообще не видит запоминания обучающей выборки — модель, копирующая датасет, получит идеальный FID.
  • Precision / Recall для распределений (Kynkäänniemi et al., 2019) — разделяют «качество» и «разнообразие». FID их смешивает, поэтому mode collapse по одному FID можно не заметить.
  • CLIP score — соответствие картинки промпту для text-to-image.
  • Человеческая оценка — по-прежнему единственный арбитр для генерации по тексту.

6. Типичные ошибки

  1. Не согласованы диапазоны данных и выхода сети. Диффузия и GAN работают с tanh-выходом в $[-1, 1]$, а датасет часто отдаёт $[0, 1]$. Модель обучится, но качество упадёт вдвое и вы будете искать причину в архитектуре.
  2. Суммирование против усреднения в ELBO. Если реконструкцию усреднить по пикселям, а KL — нет, соотношение членов уедет в $D$ раз, и латент схлопнется. Держите оба члена в одинаковых единицах (сумма по измерениям, среднее по батчу).
  3. Забытый detach() в GAN — генератор получает градиент от лосса дискриминатора. Обучение «идёт», лоссы выглядят прилично, картинок нет.
  4. Оценка GAN по значению лосса. Лосс измеряет баланс игры. Смотрите FID и фиксированную сетку сэмплов из одного и того же $z$ по эпохам.
  5. Слишком большой guidance. $w = 20$ даёт пересыщенные, «пластиковые» изображения и обрушивает разнообразие. Если хочется силы — берите dynamic thresholding вместо накрутки $w$.
  6. Один $t$ на весь батч в диффузии. Оценка градиента становится высокодисперсной; сэмплируйте $t$ независимо для каждого объекта (см. код выше).
  7. Отсутствие EMA весов. И в GAN, и в диффузии инференс делают с экспоненциальным средним весов (decay 0.999–0.9999). Разница в FID — десятки процентов. Это самая дешёвая оптимизация из существующих.
  8. Сравнение FID, посчитанных разными реализациями. Числа несопоставимы между статьями, если различаются препроцессинг и число сэмплов. Сравнивайте только внутри своего пайплайна.
  9. Игнорирование утечки данных в метрику. Проверяйте ближайших соседей сгенерированных сэмплов в обучающей выборке — диффузионные модели умеют дословно запоминать многократно повторённые картинки (Carlini et al., 2023), что является юридическим и приватностным риском.

7. Как это выглядит в проде

Обучение. Диффузионные модели обучают неделями на сотнях GPU: mixed precision (bf16), gradient checkpointing, ZeRO/FSDP-шардинг оптимизатора. Датасет — предварительно закодированные VAE-латенты и текстовые эмбеддинги, сложенные в шардированные .tar/webdataset; кодировать на лету значит греть GPU впустую. Bucketing по соотношению сторон обязателен, иначе модель выучивает единственный квадратный формат.

Инференс. Основная работа — сократить число шагов, потому что стоимость линейна по ним и удваивается из-за CFG (два прохода на шаг). Стандартный стек:

  • солвер вместо наивного цикла: DPM-Solver++ или UniPC, 20–30 шагов вместо 1000;
  • батчинг условного и безусловного прохода в один форвард — CFG почти бесплатен по латентности;
  • дистиллированный вариант модели (LCM, Turbo) для интерактивных сценариев с 1–4 шагами;
  • квантизация UNet/DiT и компиляция графа — см. статью про инференс и деплой;
  • планировщик очереди: генерация — это долгая задача (секунды), её нельзя обслуживать тем же синхронным HTTP-обработчиком, что и классификацию.

Управляемость. В прод почти никогда не идёт «голая» модель. Поверх ставят LoRA-адаптеры под стиль (та же техника, что для LLM), ControlNet для условия на позу или карту глубины, IP-Adapter для условия на референсное изображение, inpainting-режим для локальных правок.

Безопасность. Обязательные элементы: фильтр промптов, NSFW-классификатор поверх выхода, невидимая водяная метка в сгенерированном изображении, логирование seed + промпт + версия модели для воспроизводимости инцидентов. Подробнее об атаках и защите — в статье про интерпретируемость и безопасность.

Где эти модели применяются помимо картинок: генерация молекул и белков (диффузия по координатам атомов), синтез табличных данных для приватного обмена, аугментация редких классов в промышленном контроле качества, детекция аномалий (объект с низким правдоподобием под $p_\theta$ подозрителен), сжатие с обучаемым кодеком, планирование в RL (Diffuser генерирует траектории), text-to-speech (диффузионные вокодеры).

8. Мини-итог

  • Любая генеративная модель — это ответ на вопрос «как обойти неберущуюся нормировочную константу $Z_\theta$». Три семейства дают три разных ответа, и все их сильные и слабые стороны выводятся именно отсюда.
  • VAE оптимизирует нижнюю границу правдоподобия; даёт гладкий латент и мгновенный сэмплинг, но размывает детали из-за mass-covering KL и усредняющего правдоподобия. Сегодня живёт в основном как компрессор внутри latent diffusion.
  • GAN отказывается от плотности и учит метрику второй сетью; даёт резкость и генерацию за один форвард, платит нестабильностью и коллапсом мод. Актуален там, где критична латентность и нужен редактируемый $W$-латент.
  • Диффузия разбивает генерацию на сотни лёгких шагов расшумления и сводит обучение к MSE. Отсюда устойчивость, покрытие мод и управляемость через CFG — ценой дорогого последовательного сэмплинга, который сейчас активно сжимают дистилляцией и flow matching.
  • Метрики врут поодиночке: FID смешивает качество и разнообразие, лосс GAN не измеряет ничего полезного. Смотрите precision/recall, сетку сэмплов и людей.

Источники

Что дальше

Мы прошли данные с решёткой (CNN), с порядком (RNN, трансформеры) и научились их порождать. Остался важный класс, где структура данных — произвольный граф связей: социальные сети, молекулы, транзакции, дорожные графы. О том, как обобщить свёртку на такую структуру, — в статье Графовые нейронные сети.

Нашли неточность? Выделите фрагмент текста — рядом появится жучок.

Нужен разбор именно вашей ситуации?

Статья описывает общий случай. Если у вас частный — можно разобрать его отдельно, платно. А если не хватает целого материала, предложите тему: её оплачивают вскладчину, и она выходит открытой для всех.

Доска запросов