Диффузия и потоки

Оптимальный транспорт

Расстояние между распределениями, которое видит геометрию — и его связь с потоками

Шаг 100 из 117 · ~26 мин

Другая идея расстояния

KL из блока 4 сравнивает распределения поточечно: если носители не пересекаются, она бесконечна независимо от того, насколько далеко распределения друг от друга. Для генеративных моделей это плохо: две почти правильные модели, промахнувшиеся в разные стороны, неразличимы.

Оптимальный транспорт спрашивает иначе: сколько работы нужно, чтобы перевезти одно распределение в другое?

Wpp(μ,ν)=minγΠ(μ,ν)xypdγ(x,y)W_p^p(\mu,\nu) = \min_{\htmlData{k=plan}{\gamma \in \Pi(\mu,\nu)}} \int \htmlData{k=cost}{\|x - y\|^p}\, d\gamma(x,y)

γ\gamma — это совместное распределение, маргиналы которого равны μ\mu и ν\nu: сколько массы из точки xx уехало в точку yy. Задача Монжа требовала детерминированного отображения; Канторович разрешил дробить массу, и от этого задача стала линейной, а значит решаемой.

Три факта, которые стоит знать

В одномерном случае задача решается сортировкой. Оптимальный план — сопоставить kk-ю по порядку точку одной выборки с kk-й другой. Проверено перебором всех 720720 перестановок на двух выборках по шесть точек: сортировка даёт 21.089080276821.0890802768, полный перебор — то же число.

Между гауссианами есть замкнутая форма. В одномерном случае

W22=(μ1μ2)2+(σ1σ2)2W_2^2 = (\mu_1 - \mu_2)^2 + (\sigma_1 - \sigma_2)^2

распределенияW2W_2
N(0,1)\mathcal{N}(0,1) и N(3,1)\mathcal{N}(3,1)3.0003.000
N(0,1)\mathcal{N}(0,1) и N(0,4)\mathcal{N}(0,4)1.0001.000
N(1,0.25)\mathcal{N}(1, 0.25) и N(1,2.25)\mathcal{N}(-1, 2.25)2.2362.236

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

Метрика видит расстояние, а не пересечение. Для непересекающихся носителей KL равна бесконечности, а W2W_2 — расстоянию между ними. Именно поэтому Вассерштейн появился в GAN и почему он естественен для потоков: он измеряет перемещение, а потоки перемещением и занимаются.

Формулировка Бенаму–Бренье: транспорт как поток

Вот связь, ради которой урок стоит в этой ветке. Расстояние W2W_2 имеет динамическую формулировку:

W22(μ,ν)=minv01 ⁣ ⁣v(x,t)2pt(x)dxdtW_2^2(\mu,\nu) = \min_{v} \int_0^1\!\!\int \|v(x,t)\|^2\, p_t(x)\, dx\, dt

при условии, что поле vv переносит μ\mu в ν\nu. То есть W2W_2 — это минимальная кинетическая энергия потока, переводящего одно распределение в другое.

Отсюда сразу два следствия:

  1. Оптимальный поток движется по прямым с постоянной скоростью. Из всех способов проехать из точки в точку за единицу времени минимум v2dt\int\|v\|^2dt даёт равномерное прямолинейное движение — это неравенство Коши–Буняковского (блок 1, урок 050), а не свойство транспорта.
  2. Цель flow matching и цель OT — разные. Flow matching берёт случайные пары (x0,x1)(x_0, x_1); оптимальный транспорт выбирает пары так, чтобы суммарная стоимость была минимальна. Первое проще, второе даёт прямее.

Minibatch OT: как это используют

Практический приём из урока 090, теперь с объяснением. Вместо случайного сопоставления шума и данных внутри батча решают маленькую задачу транспорта: батч на батч, венгерский алгоритм или Синкхорн. Пары перестают пересекаться, усреднение искажает меньше, маргинальное поле выпрямляется.

способ парстоимостьпрямизна
случайнонулеваякак есть
minibatch OTO(B3)O(B^3) или Синкхорнзаметно лучше
rectified flowпереобучениелучше всего, но дороже

Оговорка, которую стоит держать: OT внутри батча — не глобальный OT. При размере батча BB решается задача на BB точках, а не на всём наборе данных, и с ростом BB приближение улучшается. Это компромисс, а не решение.

Посмотрите, что даёт прямизна для сэмплера: при четырёх шагах прямой путь проходится точно, а кривой — срезается.

шагов 4
длина пути
2.83
прямое расстояние
2.83
кривизна
×1
ошибка ломаной
1.6e-16
Кривизна ровно единица — это и есть тот путь, который выбрал бы оптимальный транспорт для данной пары концов: прямая с постоянной скоростью минимизирует кинетическую энергию.

Путь совпадает с прямой: минимум кинетической энергии по Бенаму–Бренье достигается ровно на равномерном прямолинейном движении.

Чего оптимальный транспорт не делает

Стоит закрыть популярное ожидание. Знание W2W_2 не даёт генеративной модели: расстояние — это число, а нужна выборка. И решение задачи транспорта между шумом и данными в полной постановке столь же дорого, как сама генерация, — при nn точках это O(n3)O(n^3) и требует всех данных сразу.

Практическая роль OT в этой ветке скромнее и точнее: он объясняет, почему прямые пути хороши (минимум энергии), и даёт дешёвое приближение внутри батча. Не более того — и этого достаточно.

Итог

  • WpW_p измеряет стоимость перевозки массы и остаётся конечной при непересекающихся носителях.
  • В одном измерении задача решается сортировкой; между гауссианами есть замкнутая форма.
  • Бенаму–Бренье переписывает W2W_2 как минимум кинетической энергии потока — отсюда прямые пути.
  • Flow matching со случайными парами не решает задачу OT; minibatch OT приближает её и выпрямляет поле.
  • Само по себе расстояние генеративной модели не даёт: полезны формулировка и приближение, а не число.

Источники

Проверки

0 из 2
  1. Транспорт, энергия и прямизна

    Отметьте все верные утверждения об оптимальном транспорте.

  2. Транспорт на прямой

    Реализуйте transport_1d(source, target) — два списка одинаковой длины nn, каждая точка несёт массу 1/n1/n. Стоимость назначения — квадрат расстояния. Верните [optimal_cost, w2, naive_cost, worst_cost]:

    • optimal_cost — минимальная суммарная стоимость по всем сопоставлениям «один к одному». Считать перебор не нужно: на прямой оптимум даёт сортировка обоих списков и сопоставление по порядку;
    • w2 = optimal_cost/n\sqrt{\texttt{optimal\_cost}/n} — расстояние Вассерштейна между двумя равномерными наборами точек;
    • naive_cost — стоимость сопоставления «по порядку поступления», то есть ii-й точки источника с ii-й точкой цели без всякой сортировки;
    • worst_cost — максимальная суммарная стоимость. Её тоже можно получить без перебора: она достигается на сопоставлении отсортированного источника с целью, отсортированной в обратном порядке.

    Проверить себя можно так: optimal_cost никогда не превосходит naive_cost, а worst_cost никогда не меньше обоих.

    функция transport_1d

    Загрузка редактора…

    Ctrl/⌘ + Enter