Flow Matching: обучение и дистилляция. Flow Matching.. Flow Matching. дистилляция.. Flow Matching. дистилляция. диффузионные модели.
Рисунок 1. Примеры возможностей генеративных моделей: из шума появляются изображения, видео, звук и другие сложные объекты

Рисунок 1. Примеры возможностей генеративных моделей: из шума появляются изображения, видео, звук и другие сложные объекты

Flow Matching — современный подход генеративного моделирования. Основная идея — научить модель постепенно превращать простое распределение, например, случайный шум, в сложные данные — изображения, видео, звук, 3D-сцены или действия робота.

В рамках статьи разберём устройство Flow Matching’а: как задаётся путь от шума к данным, чему учится модель и как после обучения из случайного шума появляется новый объект. Обсудим главный недостаток подхода — медленная генерация (чтобы получить один объект, нужно выполнить много последовательных шагов).

Один из способов его решения — дистилляция. Мы заменим многошаговый процесс генератором, который за одно вычисление нейронной сети произведёт нам нужные картинки.

Однако такую модель нельзя просто обучить регрессией по принципу: «Вот — вход, вот — правильный ответ». Мы разберём, почему это так и как из задачи дистилляции естественным образом возникает минимаксная постановка: две модели начинают играть друг с другом, а результатом игры становится качественная генеративная модель!

Сферы применения

Flow Matching сегодня — один из ключевых подходов в современных генеративных моделях, «равный» диффузионным моделям (в каком-то смысле считается, что это одно и то же).

Известные примеры применения — Stable Diffusion 3.5, FLUX.2 и Kandinsky от Сбера. Эти модели способны превращать текстовое описание в новое изображение — от реалистичной фотографии до сложной фантастической сцены.

Рисунок 2. Пример генерации изображений в Stable Diffusion 3.5

Рисунок 2. Пример генерации изображений в Stable Diffusion 3.5

Но изображениями применение Flow Matching’а не ограничивается. С его помощью можно генерировать и более сложные объекты — например, видео вместе со звуком. MiniMax H3 получает текстовое описание сцены и постепенно превращает случайный шум не просто в одну картинку, а в целую последовательность кадров и соответствующий ей звук.

Но генеративные модели идут ещё дальше: сегодня они умеют создавать уже не отдельные изображения или видео, а целые интерактивные миры, с которыми можно взаимодействовать в реальном времени. Например, семейство Genie от Google DeepMind позволяет сгенерировать виртуальное окружение из изображения или текстового описания, а затем исследовать его, управляя персонажем. Получается, можно придумать игру и почти сразу оказаться внутри неё!

Определение

Flow Matching — один из современных подходов генерации изображений, видео, звука. Его основная идея очень похожа на интуитивное «расшумление»: мы начинаем со случайного шума и постепенно превращаем его в осмысленный объект — например, в чёткое изображение.

Рисунок 6. Интуиция «расшумления»: шум постепенно превращается в осмысленное изображение

Рисунок 6. Интуиция «расшумления»: шум постепенно превращается в осмысленное изображение

Формально у нас есть настоящие данные x_text{data} sim p_{text{data}}. Это могут быть, например, реальные изображения из интернета. Мы хотим научиться порождать новые объекты, которые похожи на данные из этого распределения, но не являются простым копированием обучающих примеров. То есть мы хотим выучить распределение картинок по набору примеров из него.

С другой стороны, нам нужно распределение, из которого легко начинать генерацию. Обычно для этого берут стандартное нормальное распределение: x_text{noise} sim mathcal{N}(0, I). Из него легко семплировать случайные точки, поэтому оно обычно выступает начальным шумом.

Как происходит переход от шума к данным с помощью непрерывного времени t in [0,1]? В момент t=0 мы находимся в шуме: x_0=x_text{noise} sim mathcal{N}(0, I), а в момент t=1 хотим оказаться в данных: x_1=x_text{data} sim p_{text{data}}.

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

Рисунок 7. Схема движения семплов из шумового распределения p_0 к семплам из распределения данных p_1; x_t — промежуточные состояния

Рисунок 7. Схема движения семплов из шумового распределения p_0 к семплам из распределения данных p_1; x_t — промежуточные состояния

На рисунке выше мы работаем в одномерном пространстве (по вертикали — координата, по горизонтали — время), либо в двумерном пространстве (шум — слева, картинки — справа, и мы движемся в этом двумерном пространстве во времени).

Но для двумерного пространства траектории могут быть и другие, например, петли:

Рисунок 8. Схема движения семплов из шумового распределения p_0 к семплам из распределения данных p_1 с образованием петель

Рисунок 8. Схема движения семплов из шумового распределения p_0 к семплам из распределения данных p_1 с образованием петель

Самый простой способ соединить шум x_0 и объект x_1 во времени — провести между ними прямую линию:

x_t=(1-t)x_0 + t x_1, qquad tin[0,1].

При t=0 эта формула даёт x_t=x_0, то есть шум. При t=1 она определяет x_t=x_1, то есть объект из данных. А при промежуточных значениях времени t точка находится где-то между шумом и изображением.

Скорость движения вдоль такой прямой траектории постоянна и равна:

frac{d x_t}{dt}=x_1 - x_0.

То есть здесь направление движения (вектор скорости) — вектор от начального шума к конечному объекту.

Рисунок 9. Простой случай: каждая точка идёт к выбранному объекту по прямой траектории, а скорость вдоль неё равна x_1-x_0

Рисунок 9. Простой случай: каждая точка идёт к выбранному объекту по прямой траектории, а скорость вдоль неё равна x_1-x_0

Траектории точек (или частиц) могут быть совершенно разными. Вместо предварительного знания конечной точки x_1 или точки на траектории (что кажется очень сложным, потому что нужно сопоставлять точки из шума с реальными картинками) попробуем описать само движение. Для этого удобно использовать векторное поле.

Векторное поле u(x,t) говорит нам, куда и с какой скоростью нужно двигаться из точки x в момент времени t. Например, если сейчас мы находимся в точке x_t, то поле u(x_t,t) показывает направление следующего шага.

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

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

На этой и следующей визуализациях векторы короче, чем они есть, потому что удобно брать далёкие x_0 и x_1. В таком случае x_1 - x_0 будет большим и сложным для отображения. Тогда мы будем «домножать» векторы на некоторое d tau.

Если сделать небольшой шаг по времени dt — новое положение точки можно записать так:

x_{t+dt}=x_t + dt cdot u(x_t, t).

Следовательно, мы берём текущую точку x_t, смотрим на скорость, которую предсказывает поле, и немного передвигаемся в этом направлении.

Рисунок 11. Семплирование x_1 от x_0 в направлении по u(x_t, t)

Рисунок 11. Семплирование x_1 от x_0 в направлении по u(x_t, t)

Итак, задачу генерации можно сформулировать так: необходимо обучить нейросеть u_{theta}(x,t), которая приближает правильное векторное поле. После этого мы можем взять случайный шум x_0 sim mathcal{N}(0,I) и постепенно двигать его по этому векторному полю от t=0 до t=1. Если поле обучено хорошо, то в конце этого движения мы получим реалистичное изображение.

В математике для формального обозначения движения по вектору скорости используют обыкновенные дифференциальные уравнения (ordinary differential equations – ODE). Мы говорим, что хотим найти x_1, зная x_0 sim mathcal{N}(0,I) и закон frac{d x_t}{d t}=u(x_t, t). Имеем начальную точку и закон скорости, значит, можем получить конечную точку, просто проинтегрировав x_s=x_0 + int_{0}^s u(x_t, t) dt и x_1=x_0 + int_{0}^1 u(x_t, t) dt, где x_0 известен. Это мы и аппроксимируем через маленькие шаги dt.

Если мы для генерации траектории каждый раз будем вычислять u(x_t, t) — мы не получим траектории, которые в определённый момент времени пересекутся и отойдут друг от друга. Дело в том, что при пересечении мы получим одинаковый x_t, и далее обе траектории пойдут по u(x_t, t), то есть уже не разъединятся.

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

Обучение

Вернёмся к простой конструкции с прямыми траекториями. Возьмём пару объектов x_0, x_1, где x_0 sim mathcal{N}(0,I) — случайный шум, а x_1 sim p_{text{data}} — реальное изображение. Между ними можно провести прямую траекторию:

x_t=(1-t)x_0 + t x_1.

Мы можем соединить все x_0 sim p(x_0) со всеми x_1 sim p(x_1) прямыми траекториями (каждый с каждым). В таком случае для текущего x_t бесконечное число пар x_0, x_1 задаёт прямую траекторию, проходящую через x_t. То есть промежуточная точка появляется из разных пар: шумы x_0 могли двигаться к различным изображениям x_1 и в какой-то момент оказаться примерно в одном и том же месте, имея скорости x_1 - x_0.

Итоговую скорость u(x_t, t) можно получить, если усреднить все скорости, которые проходят через точку x_t в момент времени t. То есть мы можем выразить глобальную скорость u(x_t, t) через среднее локальных скоростей x_1 - x_0 по всем парам x_0, x_1.

Итак, Flow Matching показывает, что нужно предсказывать среднюю скорость всех частиц, которые проходят через эту точку в данный момент времени, чтобы получить скорость для искомых траекторий. Формально это записывается через условное матожидание:

u^*(x,t)=mathbb{E}_{x_0,x_1} left[ x_1 - x_0 mid x_t=x right].

Условное математическое ожидание

Математическое ожидание — обобщение понятия среднего элементов. Но элементов может быть бесконечно много. Как и в нашем случае :)

Условное математическое ожидание — среднее элементов, полученных при каких-то условиях. К примеру, если мы хотим найти среднюю продолжительность сна перед сессией, то сессия — это условие.

Формула выше означает следующее: мы смотрим на все пары x_0,x_1, для которых промежуточная точка x_t оказалась равна x, берём их скорости x_1-x_0 и усредняем. Получившееся среднее направление и есть правильное векторное поле в точке x и времени t. Его мы и будем использовать для семплирования.

Здесь время t фиксировано: мы усредняем по парам (x_0,x_1), которые в этот момент проходят через точку x, а не по разным моментам времени.

Рисунок 12. Прохождение траекторий от разных пар (x_0,x_1) в одной точке; поле u^*(x,t) берёт среднее направление их скоростей, а мы в итоге получаем искомую траекторию

Рисунок 12. Прохождение траекторий от разных пар (x_0,x_1) в одной точке; поле u^*(x,t) берёт среднее направление их скоростей, а мы в итоге получаем искомую траекторию

Прямые траектории при соединении x_0 с x_1 и конечные траектории при семплировании — не одно и то же. Просто нам достаточно учить не все скорости, проходящие через текущую точку, а их среднее.

Почему можно брать именно среднее? Интуитивно, нас интересует не судьба одной конкретной частицы, а то, как меняется всё распределение точек. Если через одну область пространства проходит много частиц с разными скоростями, то общее движение плотности определяется их средним потоком. Поэтому для генерации достаточно выучить не каждую отдельную траекторию, а усреднённое векторное поле u^*(x,t).

На практике мы обучаем нейросеть u_{theta}(x,t) приближать поле. Для этого семплируем пару x_0,x_1, выбираем случайное время t, строим промежуточную точку x_t=(1-t)x_0 + t x_1 и просим модель по x_t,t предсказать скорость x_1 - x_0.

Получается обычная задача регрессии:

mathcal{L}(theta)=mathbb{E}_{x_0,x_1,t} left[ left| u_theta(x_t,t) - (x_1-x_0) right|^2 right].

Минимум этой MSE-задачи при каждом фиксированном времени t как раз и равен условному среднему:

u_{theta}(x,t) approx u^*(x,t)=mathbb{E}_{x_0,x_1} left[ x_1-x_0 mid x_t=x right].

Условное среднее

Покажем, что минимум MSE-задачи равен условному среднему.

Распишем наш лосс как математическое ожидание от условного математического ожидания (вынесем x_t):

begin{aligned}mathcal{L}(theta)&=mathbb{E}_{x_0,x_1,t}left[left|u_theta(x_t,t)-(x_1-x_0)right|^2right] \&=mathbb{E}_{t,x_t}left[mathbb{E}_{x_0,x_1 mid x_t,t}left[left|u_theta(x_t,t)-(x_1-x_0)right|^2right]right].end{aligned}

Отдельно обозначим mu(x_t, t) как условное среднее mu(x_t,t)=mathbb{E}_{x_0,x_1 mid x_t,t}left[x_1-x_0right]. Мы хотим показать, что u_theta(x_t, t) как раз выучит mu(x_t, t).

Распишем отдельно то, что стоит под математическим ожиданием:

begin{aligned} & mathbb{E}_{x_0,x_1 mid x_t,t} left[ left| u_theta(x_t,t)-(x_1-x_0) right|^2 right] \ &=mathbb{E}_{x_0,x_1 mid x_t,t} left[ left| u_theta(x_t,t)-mu(x_t,t) + mu(x_t,t)-(x_1-x_0) right|^2 right] \ &=mathbb{E}_{x_0,x_1 mid x_t,t} left[ left| u_theta(x_t,t)-mu(x_t,t) right|^2 right] \ &quad + mathbb{E}_{x_0,x_1 mid x_t,t} left[ left| mu(x_t,t)-(x_1-x_0) right|^2 right] \ &quad + 2 mathbb{E}_{x_0,x_1 mid x_t,t} left[ leftlangle u_theta(x_t,t)-mu(x_t,t), mu(x_t,t)-(x_1-x_0) rightrangle right] \ &=left| u_theta(x_t,t)-mu(x_t,t) right|^2 + mathbb{E}_{x_0,x_1 mid x_t,t} left[ left| mu(x_t,t)-(x_1-x_0) right|^2 right] \ &quad + 2 leftlangle u_theta(x_t,t)-mu(x_t,t), mathbb{E}_{x_0,x_1 mid x_t,t} left[ mu(x_t,t)-(x_1-x_0) right] rightrangle \ &=left| u_theta(x_t,t)-mu(x_t,t) right|^2 + mathbb{E}_{x_0,x_1 mid x_t,t} left[ left| mu(x_t,t)-(x_1-x_0) right|^2 right] \ &quad + 2 leftlangle u_theta(x_t,t)-mu(x_t,t), mu(x_t,t) - mathbb{E}_{x_0,x_1 mid x_t,t} left[x_1-x_0 right] rightrangle \ &=left| u_theta(x_t,t)-mu(x_t,t) right|^2 + mathbb{E}_{x_0,x_1 mid x_t,t} left[ left| mu(x_t,t)-(x_1-x_0) right|^2 right] \ &quad + 2 leftlangle u_theta(x_t,t)-mu(x_t,t), mu(x_t,t) - mu(x_t, t) rightrangle \ &=left| u_theta(x_t,t)-mu(x_t,t) right|^2 + mathbb{E}_{x_0,x_1 mid x_t,t} left[ left| mu(x_t,t)-(x_1-x_0) right|^2 right] \ &=left| u_theta(x_t,t) - mu(x_t,t) right|^2 + operatorname{tr} operatorname{Cov}_{x_0,x_1 mid x_t,t}(x_1-x_0). end{aligned}

В итоге получим, что оптимум у u_theta(x_t, t) равен mu(x_t, t), так как второе слагаемое не зависит от u_theta. Что мы и хотели доказать.

Именно поэтому Flow Matching можно обучать просто: мы сами строим промежуточные точки x_t и знаем целевую скорость x_1-x_0 для каждой обучающей пары.

Рисунок 13. Переход от траекторий к полю: в промежуточной точке много семпловых скоростей, усредняем их и получаем локальную стрелку векторного поля

Рисунок 13. Переход от траекторий к полю: в промежуточной точке много семпловых скоростей, усредняем их и получаем локальную стрелку векторного поля

В следующих трёх скрытых блоках можно посмотреть примеры траекторий Flow Matching’а на задаче перевода одного нормального распределения в два, а также подробнее почитать про интуицию метода.

Перевод одной гауссианы в две

Рассмотрим пример с отображением одного нормального распределения в смесь двух нормальных распределений. Для начала визуализируем семплы из x_0 sim p_0 и x_1 sim p_1, а также промежуточные семплы x_t sim p_t, которые были получены как x_t=(1 - t)x_0 + t x_1

Рисунок 14. Игрушечный пример с настоящими гауссианами: семплы из одного начального гауссиана переходят к смеси из двух гауссианов

Рисунок 14. Игрушечный пример с настоящими гауссианами: семплы из одного начального гауссиана переходят к смеси из двух гауссианов

Для каждой промежуточной точки мы можем найти x_0 и x_1, которые её породили, и соответствующую скорость x_1 - x_0 (на рисунке для удобства восприятия мы ужали векторы по длине).

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

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

Итоговая скорость в момент времени t для точки x — среднее среди скоростей x_1 - x_0, для которых x_t=x (для удобства восприятия векторы мы ужали по длине).

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

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

Представим, что в момент времени t мы смотрим на некоторую точку x_t. Она могла появиться из разных пар (x_0, x_1). Для каждой такой пары есть своя скорость v=x_1 - x_0. Можно сказать, что она получена применением некоторой силы, толчка. Поэтому в одной и той же точке x_t может быть не одна «правильная сила», а целое распределение возможных сил.

Теперь вместо того, чтобы сразу заменять все эти силы на среднюю, будем прикладывать их по одной. На очень маленьком промежутке времени dt случайно выберем одну скорость из условного распределения:

p(x_1 - x_0 mid x_t,t)

и сдвинем точку:

x_{t+dt}=x_t + dt cdot (x_1 - x_0).

Почему такая процедура сохраняет правильное распределение? Потому что в исходном процессе происходит ровно то же самое, если смотреть только на точку x_t. Среди всех частиц, которые оказались в x_t в момент времени t, скорости распределены по тому же условному закону:

p(x_1 - x_0 mid x_t,t).

Значит, если мы выбираем скорость из этого распределения, то за маленький шаг dt получаем такой же локальный перенос частиц, как и в исходном процессе с прямыми траекториями (где мы соединяли x_0 и x_1 и шли линейно от x_0 к x_1).

Теперь разобьём маленький промежуток времени dt на множество микропромежутков времени d tau. На каждом микрошаге мы можем выбрать одну из возможных скоростей из p(x_1 - x_0 mid x_t,t). Но отдельный сдвиг очень мал, потому что он умножается на d tau. Поэтому за каждый конкретный микрошаг мы почти не сдвигаемся, p(x_1 - x_0 mid x_t,t) не меняется, а за множество таких шагов мы в итоге сдвинемся на среднее направление всех этих скоростей, умноженное на dt.

Иными словами, если в точке x_t на частицу действует много возможных сил, то на бесконечно малом масштабе их можно заменить одной средней силой:

u^*(x,t)=mathbb{E}_{x_0,x_1} left[ x_1 - x_0 mid x_t=x right].

Именно это среднее поле и нужно выучить модели. Оно не пытается восстановить конкретную пару (x_0,x_1), а описывает средний поток частиц, который переносит всё распределение шума к распределению данных.

Рисунок 17. Замена множества возможных микротолчков в точке x в момент времени t на малом масштабе одним средним сдвигом u^*(x, t)

Рисунок 17. Замена множества возможных микротолчков в точке x в момент времени t на малом масштабе одним средним сдвигом u^*(x, t)
Механическая аналогия

Эту идею можно представить с помощью механической аналогии. Пусть каждая пара (x_0, x_1) задаёт один «удар по мячу»: футболист стоит в точке x_0, ворота находятся в точке x_1, и мяч летит от футболиста к воротам по прямой траектории:

x_t=(1-t)x_0 + t x_1.

Скорость такого мяча постоянна и равна:

x_1 - x_0.

Причём футболисты не знают, в какие конкретно ворота они должны попадать, поэтому каждый бьёт в каждые ворота.

Теперь посмотрим на некоторую точку x в момент времени t. Через неё могут пролетать мячи, запущенные из разных x_0 в разные x_1. Поэтому в одной и той же точке x мы можем увидеть много различных скоростей: один мяч летит чуть левее, другой — чуть правее, третий — почти прямо.

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

Рисунок 18. Механическая аналогия: разные удары проходят через одну область, а заграждение сдвигается в направлении среднего толчка

Рисунок 18. Механическая аналогия: разные удары проходят через одну область, а заграждение сдвигается в направлении среднего толчка

Но если за это маленькое время через точку проходит много мячей с разными скоростями, то суммарный эффект определяется не одной конкретной скоростью, а средним направлением всех этих микротолчков. То есть заграждение будет двигаться так, как если бы на него действовала средняя скорость всех мячей, пролетающих через точку x в момент t (если бы «запульнули» все мячи сразу).

Именно эта средняя скорость и задаёт поле Flow Matching’а:

u^*(x,t)=mathbb{E}_{x_0,x_1} left[ x_1 - x_0 mid x_t=x right].

То есть модель не пытается понять, какой именно мяч сейчас пролетел через точку x. Вместо этого она учится предсказывать средний толчок от всех возможных мячей, которые могли оказаться в точке в данный момент времени. Такое усреднённое движение и переносит всё распределение шума к распределению данных.

Алгоритм

На каждой итерации мы берём случайный шум x_0, случайный объект из данных x_1 и случайное время t. Затем строим промежуточную точку x_t на прямой между шумом и объектом и просим модель предсказать её скорость. На практике всё то же самое делается сразу для целого батча объектов, поэтому алгоритм ниже записан в батчевой форме.

  • ⚙️ Алгоритм: Обучение Flow Matching-модели

  • Вход: Семплы из распределения p_{text{data}} (возможность семплировать), модель векторного поля u_theta, оптимизатор для параметров theta, количество итераций N, размер батча B

  • Результат: Обученное поле u_theta approx u^*

  • Для i=1 до N:

    • Семплируем батч шумов, объектов и времён: x_0^{(b)} sim mathcal{N}(0,I), qquad x_1^{(b)} sim p_{text{data}}, qquad t^{(b)} sim mathcal{U}[0,1], qquad b=1,ldots,B

    • Строим промежуточные точки: x_t^{(b)} gets (1-t^{(b)})x_0^{(b)} + t^{(b)}x_1^{(b)}

    • Считаем целевые скорости прямых траекторий: v^{(b)} gets x_1^{(b)} - x_0^{(b)}

    • Считаем среднюю ошибку предсказания скорости по батчу: mathcal{L} gets frac{1}{B} sum_{b=1}^{B} left| u_theta(x_t^{(b)},t^{(b)}) - v^{(b)} right|_2^2

    • Обновляем параметры theta шагом оптимизатора по mathcal{L}

После обучения у нас есть модель u_{theta}(x,t), которая приближает среднее поле скоростей. Теперь мы можем использовать её для генерации новых объектов.

Идея семплирования следующая — мы стартуем из случайного шума:

x_0 sim mathcal{N}(0,I)

И постепенно двигаем точку по выученному векторному полю от времени t=0 до времени t=1. Если разбить отрезок [0,1] на K_{text{ode}} маленьких шагов, то один шаг движения можно записать так:

x_{t+Delta t}=x_t + Delta t cdot u_theta(x_t,t),

Где:

Delta t=frac{1}{K_{text{ode}}}.

То есть на каждом шаге модель говорит нам, в каком направлении нужно немного сдвинуть текущую точку. После K_{text{ode}} таких шагов мы получаем финальный объект x_1, который должен быть похож на данные. Если нужно сгенерировать множество объектов, тот же цикл обычно выполняется параллельно для батча начальных шумов.

  • ⚙️ Алгоритм: Семплирование с помощью Flow Matching’а

  • Вход: Обученное поле скоростей u_theta, количество шагов семплирования K_{text{ode}}

  • Результат: Сгенерированный объект hat{x}

  • Семплируем начальный шум: x sim mathcal{N}(0,I)

  • Задаём размер шага по времени: Delta t gets frac{1}{K_{text{ode}}}

  • Для k=0 до K_{text{ode}}-1

    • Текущее время: t_k gets kDelta t

    • Предсказываем скорость в текущей точке: v_k gets u_theta(x,t_k)

    • Делаем маленький шаг по полю: x gets x + Delta t cdot v_k

  • hat{x} gets x

Здесь используется самый простой численный метод — метод Эйлера. Он буквально повторяет нашу интуицию: посмотреть на текущую скорость, сделать маленький шаг в этом направлении, затем снова посмотреть на скорость и снова сделать шаг. Чем больше K, тем точнее мы следуем выученному полю, но тем медленнее становится генерация.

В следующем скрытом блоке представлен формальный вывод Flow Matching’а через пробные функции, а также разбор уравнения непрерывности. Это позволит перейти от интуиции к математическому пониманию метода.

Формальный вывод

Итак, давайте более формально с помощью пробных функций покажем, почему Flow Matching учит именно поле:

u^*(x,t)=mathbb{E}_{x_0simmathcal{N}(0,I),,x_1sim p_{text{data}}} left[ x_1 - x_0 mid x_t=x right],

И почему, начиная с x_0 и двигаясь по u^*(x_t, t), мы получим то же самое распределение у x_t, как если бы мы просто семплировали его как x_t=(1-t)x_0 + t x_1.

Рассмотрим случайный процесс:

x_t=(1-t)x_0 + t x_1,

Где:

x_0 sim mathcal{N}(0,I), qquad x_1 sim p_{text{data}}.

Для каждой конкретной пары (x_0,x_1) скорость вдоль прямой траектории равна:

frac{d x_t}{dt}=x_1 - x_0.

Пусть p_t(x) — плотность распределения промежуточных точек x в момент времени t. Мы хотим понять, как она меняется со временем.

Возьмём произвольную гладкую пробную функцию с компактным носителем:

varphi : mathbb{R}^d to mathbb{R}

Тогда для фиксированного t:

mathbb{E}_{x_0simmathcal{N}(0,I),,x_1sim p_{text{data}}} left[ varphi(x_t) right]=int varphi(x),p_t(x),dx.

Продифференцируем это ожидание по времени:

frac{d}{dt} mathbb{E}_{x_0simmathcal{N}(0,I),,x_1sim p_{text{data}}} left[ varphi(x_t) right]=mathbb{E}_{x_0simmathcal{N}(0,I),,x_1sim p_{text{data}}} left[ frac{d}{dt}varphi(x_t) right].

По правилу цепочки:

frac{d}{dt}varphi(x_t)=nabla varphi(x_t)cdot frac{d x_t}{dt}=nabla varphi(x_t)cdot (x_1-x_0).

Значит:

frac{d}{dt} mathbb{E}_{x_0simmathcal{N}(0,I),,x_1sim p_{text{data}}} left[ varphi(x_t) right]=mathbb{E}_{x_0simmathcal{N}(0,I),,x_1sim p_{text{data}}} left[ nabla varphi(x_t)cdot (x_1-x_0) right].

Теперь используем условное математическое ожидание по текущей точке x_t:

mathbb{E}_{x_0simmathcal{N}(0,I),,x_1sim p_{text{data}}} left[ nabla varphi(x_t)cdot (x_1-x_0) right]=mathbb{E}_{x_t} left[ mathbb{E}_{x_0simmathcal{N}(0,I),,x_1sim p_{text{data}}} left[ nabla varphi(x_t)cdot (x_1-x_0) mid x_t right] right].

Так как nablavarphi(x_t) зависит только от x_t, её можно вынести из внутреннего условного ожидания:

mathbb{E}_{x_0simmathcal{N}(0,I),,x_1sim p_{text{data}}} left[ nabla varphi(x_t)cdot (x_1-x_0) mid x_t right]=nabla varphi(x_t) cdot mathbb{E}_{x_0simmathcal{N}(0,I),,x_1sim p_{text{data}}} left[ x_1-x_0 mid x_t right].

Введём среднее поле скоростей:

u^*(x,t)=mathbb{E}_{x_0simmathcal{N}(0,I),,x_1sim p_{text{data}}} left[ x_1-x_0 mid x_t=x right].

Тогда:

frac{d}{dt} mathbb{E}_{x_0simmathcal{N}(0,I),,x_1sim p_{text{data}}} left[ varphi(x_t) right]=mathbb{E}_{x_t} left[ nabla varphi(x_t)cdot u^*(x_t,t) right].

Переходим к интегралу по плотности p_t и получаем:

frac{d}{dt} mathbb{E}_{x_0simmathcal{N}(0,I),,x_1sim p_{text{data}}} left[ varphi(x_t) right]=int nabla varphi(x)cdot u^*(x,t),p_t(x),dx.

Теперь проинтегрируем по частям. Так как varphi имеет компактный носитель, граничные члены исчезают (на бесконечности они равны 0):

int nabla varphi(x)cdot u^*(x,t),p_t(x),dx=- int varphi(x), nablacdot left( p_t(x)u^*(x,t) right) dx.

С другой стороны:

frac{d}{dt} mathbb{E}_{x_0simmathcal{N}(0,I),,x_1sim p_{text{data}}} left[ varphi(x_t) right]=frac{d}{dt} int varphi(x)p_t(x),dx=int varphi(x) frac{partial p_t(x)}{partial t} dx.

Значит:

int varphi(x) frac{partial p_t(x)}{partial t} dx=- int varphi(x) nablacdot left( p_t(x)u^*(x,t) right) dx.

Или:

int varphi(x) left[ frac{partial p_t(x)}{partial t} + nablacdot left( p_t(x)u^*(x,t) right) right] dx=0.

Это верно для любой пробной функции varphi, следовательно, мы получаем уравнение непрерывности для плотности p_t:

frac{partial p_t}{partial t} + nablacdot left( p_t u^* right)=0.

Итак, мы показали — плотности p_t из прямых траекторий:

x_t=(1-t)x_0 + t x_1,

Удовлетворяют уравнению непрерывности с полем u^*(x,t).

Но почему можно двигаться по среднему полю?

Рассмотрим новый процесс генерации. Возьмём начальную точку:

tilde{x}_0 sim mathcal{N}(0,I)

И будем двигать её по ODE:

frac{d tilde{x}_t}{dt}=u^*(tilde{x}_t,t).

Пусть q_t(x) — плотность распределения случайной величины tilde{x}_t.

Теперь отдельно покажем, что эта плотность q_t тоже удовлетворяет уравнению непрерывности.

Снова возьмем гладкую пробную функцию varphi:

mathbb{E}_{tilde{x}_0simmathcal{N}(0,I)} left[ varphi(tilde{x}_t) right]=int varphi(x)q_t(x),dx.

Продифференцируем это ожидание по времени:

frac{d}{dt} mathbb{E}_{tilde{x}_0simmathcal{N}(0,I)} left[ varphi(tilde{x}_t) right]=mathbb{E}_{tilde{x}_0simmathcal{N}(0,I)} left[ frac{d}{dt}varphi(tilde{x}_t) right].

По правилу цепочки:

frac{d}{dt}varphi(tilde{x}_t)=nablavarphi(tilde{x}_t)cdot frac{d tilde{x}_t}{dt}.

Так как tilde{x}_t движется по ODE:

frac{d tilde{x}_t}{dt}=u^*(tilde{x}_t,t),

Получаем:

frac{d}{dt} mathbb{E}_{tilde{x}_0simmathcal{N}(0,I)} left[ varphi(tilde{x}_t) right]=mathbb{E}_{tilde{x}_0simmathcal{N}(0,I)} left[ nablavarphi(tilde{x}_t)cdot u^*(tilde{x}_t,t) right].

Теперь перепишем это ожидание как интеграл по плотности q_t:

mathbb{E}_{tilde{x}_0simmathcal{N}(0,I)} left[ nablavarphi(tilde{x}_t)cdot u^*(tilde{x}_t,t) right]=int nablavarphi(x)cdot u^*(x,t)q_t(x),dx.

Интегрируем по частям:

int nablavarphi(x)cdot u^*(x,t)q_t(x),dx=- int varphi(x) nablacdot left( q_t(x)u^*(x,t) right) dx.

С другой стороны:

frac{d}{dt} mathbb{E}_{tilde{x}_0simmathcal{N}(0,I)} left[ varphi(tilde{x}_t) right]=frac{d}{dt} int varphi(x)q_t(x),dx=int varphi(x) frac{partial q_t(x)}{partial t} dx.

Значит:

int varphi(x) frac{partial q_t(x)}{partial t} dx=- int varphi(x) nablacdot left( q_t(x)u^*(x,t) right) dx.

Или:

int varphi(x) left[ frac{partial q_t(x)}{partial t} + nablacdot left( q_t(x)u^*(x,t) right) right] dx=0.

Так как это верно для любой пробной функции varphi, получаем уравнение непрерывности для плотности q_t:

frac{partial q_t}{partial t} + nablacdot left( q_t u^* right)=0.

Теперь сравним два процесса. Для исходных прямых траекторий мы получили:

frac{partial p_t}{partial t} + nablacdot left( p_t u^* right)=0, qquad p_0=mathcal{N}(0,I).

А для движения по среднему полю:

frac{partial q_t}{partial t} + nablacdot left( q_t u^* right)=0, qquad q_0=mathcal{N}(0,I).

Это одно и то же уравнение с одинаковым начальным условием. При этом, в уравнении явно показывается изменение плотности frac{partial p_t}{partial t} по текущим p_t и u^*. Значит, если его решение единственно, то:

q_t=p_t

для всех tin[0,1].

Зачем нужна единственность решения уравнения

Рассмотрим одномерное поле frac{dx}{dt}=2sqrt{|x|}. Из точки x(0)=0 существуют разные решения. Можно навсегда остаться в нуле x(t)=0, а можно начать двигаться x(t)=t^2.

Более того, можно некоторое время стоять, а затем начать движение

x(t)=begin{cases} 0, & 0 le tle tau,\ (t-tau)^2, & t>tau. end{cases}, tau ge 0

Все эти траектории имеют одну и ту же начальную точку x(0)=0 и удовлетворяют одинаковому полю. Причина в том, что функция u(x)=2sqrt{|x|} не является локально липшицевой в окрестности точки x=0.

На практике наш сетап удовлетворяет условию единственности (можно вывести из теоремы Коши — Липшица), примем это за факт не будем к нему возвращаться далее в статье 🙂

Именно это и объясняет, почему можно двигаться по среднему полю. Мы не просто «на глаз» заменили разные скорости на среднюю. Мы показали, что исходный процесс с прямыми траекториями и новый процесс, который движется по ODE, порождают одну и ту же эволюцию плотности:

frac{d tilde{x}_t}{dt}=u^*(tilde{x}_t,t),

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

p_t(x)u^*(x,t).

Поэтому Flow Matching учит макроскопическое поле скоростей u^*(x,t), а не отдельные траектории между конкретными парами (x_0,x_1). Двигаясь по этому среднему полю из шума, мы получаем ту же эволюцию распределений, что и при движении всех обучающих пар по прямым траекториям.

Итак, первая часть про обучение Flow Matching завершена. Мы показали, что суть Flow Matching’а — выучивание средних скоростей. Самое время обновить чай и перейти ко второй части — дистилляции в одношаговый генератор! 😇

Мотивация

Итак, мы разобрали, что такое Flow Matching. Давайте посмотрим, почему его не очень удобно использовать на практике 😢

После обучения Flow Matching-модели генерация требует нескольких шагов. Мы стартуем из шума:

x_0 sim mathcal{N}(0,I)

И численно движемся по выученному полю u_{theta}(x,t) от t=0 до t=1:

x_{t+Delta t}=x_t + Delta t cdot u_theta(x_t,t).

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

Delta t=0.01,

То на отрезке [0,1] получится 100 шагов. То есть 100 обращений к нейросети для генерации одного изображения. А это очень много.

Здесь возникают две основные проблемы:

  1. Вычислительная стоимость. Каждый шаг требует отдельного запуска нейросети. Даже если один проход модели работает быстро — десятки или сотни проходов могут сделать генерацию слишком медленной для практического использования.

  2. Ошибка дискретизации. В идеале мы хотели бы двигаться по непрерывной траектории. Но на практике мы делаем конечные шаги размера Delta t. Если шаг слишком большой, траектория может заметно отклоняться от той, которую задаёт непрерывное поле. Это может ухудшать качество сгенерированного изображения.

Рисунок 19. Отклонение дискретной траектории от непрерывной при крупных шагах и приход в другую конечную точку hat{x}_1. Здесь мы считаем пространство одномерным: по вертикальной оси — время, по горизонтальной — координата

Рисунок 19. Отклонение дискретной траектории от непрерывной при крупных шагах и приход в другую конечную точку hat{x}_1. Здесь мы считаем пространство одномерным: по вертикальной оси — время, по горизонтальной — координата

Дистилляция

Отсюда возникает идея — научить отдельную модель, которая будет сразу выучивать генерацию картинок. То есть мы хотим обучить генератор G_phi сразу строить финальный объект:

hat{x}_1=G_phi(z), qquad z sim mathcal{N}(0,I_m).

Такой подход называется дистилляцией: мы используем обученную многошаговую Flow Matching-модель как учителя и пытаемся передать её поведение более быстрой одношаговой модели-генератору.

Обычно под дистилляцией понимают обучение маленькой модели по выходам большой, но мы имеем в виду обучение одношагового генератора (сама модель-генератор может быть больше модели учителя).

Рисунок 20. Дистилляция учит одношаговый генератор G_phi делать распределение p_G ближе к распределению данных p_1, которое получается многошаговым учителем

Рисунок 20. Дистилляция учит одношаговый генератор G_phi делать распределение p_G ближе к распределению данных p_1, которое получается многошаговым учителем

Важное уточнение — мы не хотим выучить отображение из x_0 в x_1, которое бы повторяло предсказание Flow Matching’а. У нас G_phi — отдельный генератор со своим латентным пространством z, не связанным с x_0 (может иметь даже другую размерность). На самом деле мы сделаем что-то похожее на GAN, дискриминатор которого имеет определённый вид и некоторую стохастичность из-за x_0.

Есть и другие методы, например, consistency models, которые ставят своей задачей выучить генератор G: x_0 → x_1 повторять hat{x}_1, полученный по x_0, двигаясь по u_theta(x_t, t) или u^*(x_t, t), но это уже тема для отдельного поста 🙂.

Построение поля по сгенерированным картинкам

Теперь мы хотим заменить многошаговое движение по Flow Matching-полю одним прямым переходом от шума к изображению.

Для этого построим генератор:

G_phi : mathbb{R}^m to mathbb{R}^d,

Он получает на вход латентный шум (его размерность может отличаться от размерности исходного пространства x_0 и x_1):

z sim mathcal{N}(0,I_m)

И сразу выдаёт объект в пространстве данных:

hat{x}_1=G_phi(z) in mathbb{R}^d.

Здесь m — размерность латентного шума генератора, а d — размерность данных. Например, если мы генерируем изображение размером 32times32times3, то d=3072, а размерность m можно выбрать отдельно.

Во время генерации нам нужен только один запуск генератора: мы семплируем z, считаем G_phi(z) и сразу получаем изображение. Но для его обучения нужно построить вспомогательное поле, как и для Flow Matching’а.

Идея следующая — траектории поля, полученного по генератору, должны быть похожи на настоящие Flow Matching-траектории (которые мы получили по полю, построенному из реальных данных). Тогда генерации и исходные картинки совпадут.

Рисунок 21. Обучение генератора с использованием двух независимых шумов: z задаёт финальную точку hat{x}1=Gphi(z), а x_0 — начало траектории в пространстве данных

Рисунок 21. Обучение генератора с использованием двух независимых шумов: z задаёт финальную точку hat{x}1=Gphi(z), а x_0 — начало траектории в пространстве данных

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

x_0 sim mathcal{N}(0,I_d), qquad z sim mathcal{N}(0,I_m).

Шум z подаётся в генератор и задаёт сгенерированный объект:

hat{x}_1=G_phi(z) in mathbb{R}^d.

А независимый шум x_0 используется как начальная точка траектории в пространстве данных. Между x_0 и hat{x}_1 строим прямую:

hat{x}_t=(1-t)x_0 + that{x}_1, qquad tin[0,1].

Так как и x_0, и hat{x}_1 лежат в mathbb{R}^d, вся траектория hat{x}_t тоже находится в пространстве данных:

hat{x}_t in mathbb{R}^d.

Скорость вдоль этой прямой равна:

frac{partial hat{x}_t}{partial t}=hat{x}_1 - x_0=G_phi(z)-x_0.

💡 Важно: x_0 и z — разные шумы.

Шум z живёт в латентном пространстве mathbb{R}^m и нужен генератору для финальной точки hat{x}_1=G_phi(z). Шум x_0 живёт в пространстве данных mathbb{R}^d и нужен для Flow Matching-траектории от стандартного шума к распределению генератора.

Поэтому при t=0 мы имеем обычный шум в пространстве данных:

hat{x}_0=x_0 sim mathcal{N}(0,I_d),

А при t=1 получаем объект из генератора:

hat{x}_1=G_phi(z).

Теперь у сгенерированных траекторий тоже есть своё среднее поле скоростей. В одной и той же точке x в момент времени t могут проходить разные траектории — они построены из различных пар (x_0,z). Поэтому определим поле генератора как условное среднее:

u^phi(x,t)=mathbb{E}_{x_0simmathcal{N}(0,I_d),,zsimmathcal{N}(0,I_m)} left[ G_phi(z)-x_0 mid hat{x}_t=x right].

Оно описывает, как в среднем движутся точки, если мы берём независимый шум x_0, сгенерированный объект G_phi(z), а затем соединяем их прямой траекторией.

То есть мы строим поле Flow Matching’а, но не на реальных данных, а на сгенерированных.

Рисунок 22. Поле генератора u^phi(x,t) — средняя скорость траекторий, которые соединяют шум x_0 со сгенерированными объектами G_phi(z)

Рисунок 22. Поле генератора u^phi(x,t) — средняя скорость траекторий, которые соединяют шум x_0 со сгенерированными объектами G_phi(z)

Лосс дистилляции

С другой стороны, у нас уже есть поле учителя u^*(x,t), заранее обученное с помощью Flow Matching’а. Оно описывает правильное движение от шума к настоящим данным.

Поэтому цель у нас — подобрать генератор G_phi так, чтобы поле его траекторий u^phi(x,t) совпадало с полем учителя u^*(x,t). Формально это можно записать так:

min_{G_phi} mathbb{E}_{t,,x_0simmathcal{N}(0,I_d),,zsimmathcal{N}(0,I_m)} left[ left| u^*(hat{x}_t,t) - u^phi(hat{x}_t,t) right|_2^2 right],

Где:

hat{x}_t=(1-t)x_0 + tG_phi(z).

То есть мы семплируем hat{x}_t по сгенерированным изображениям и учим приближать поля (скорости) в точке. При этом реальные картинки нам не нужны, достаточно предобученной модели Flow Matching’а — предобученного поля.

Рисунок 23. Приближение u^* и u^phi семплированием из p_z и p_0; p_1 показаны для наглядности

Рисунок 23. Приближение u^* и u^phi семплированием из p_z и p_0; p_1 показаны для наглядности

Мы хотим обучить генератор так, чтобы его поле совпадало с полем на реальных данных. Еще мы знаем, что распределения в начальный момент времени x_0 и hat{x}_0 совпадают (это просто гауссианы). Тогда совпадут и траектории, по которым мы движемся, и сгенерированные по ним изображения (обсудили в формальном выводе Flow Matching).

Но это пока идеальная математическая цель. На практике поле u^phi(x,t) неизвестно явно, потому что оно само является условным средним по траекториям генератора. Поэтому теперь нам нужно придумать, как приблизить цель и получить удобный алгоритм обучения! 😉

Как вычислять лосс дистилляции (делать tractable)

В предыдущем разделе мы получили идеальную цель дистилляции:

min_{G_phi} mathbb{E}_{t,,x_0simmathcal{N}(0,I_d),,zsimmathcal{N}(0,I_m)} left[ left| u^*(hat{x}_t,t) - u^phi(hat{x}_t,t) right|_2^2 right],

Где:

hat{x}_1=G_phi(z), qquad hat{x}_t=(1-t)x_0 + that{x}_1,

А поле генератора определяется как условное среднее:

u^phi(x,t)=mathbb{E}_{x_0simmathcal{N}(0,I_d),,zsimmathcal{N}(0,I_m)} left[ hat{x}_1 - x_0 mid hat{x}_t=x right].

Проблема в том, что u^phi(x,t) нельзя посчитать напрямую. Это условное математическое ожидание скорости по всем парам (x_0,z) или, эквивалентно, (x_0, hat{x}_1), траектории через точку x в момент времени t. Поэтому нам нужен способ приблизить это поле.

Для этого введём дополнительную Flow Matching-модель u_psi(x,t). Её задача — учиться предсказывать скорость траекторий генератора hat{x}_1 - x_0, чтобы использовать эту модель для вычисления лосса.

То есть u_psi обучается, как и Flow Matching-модель, только вместо настоящих данных x_1sim p_{text{data}} мы используем сгенерированные объекты:

hat{x}_1=G_phi(z).

При фиксированном генераторе G_phi модель u_psi можно обучать обычной MSE-регрессией:

min_{u_psi} mathbb{E}_{t,,x_0simmathcal{N}(0,I_d),,zsimmathcal{N}(0,I_m)} left[ left| u_psi(hat{x}_t,t) - (hat{x}_1-x_0) right|_2^2 right].

Как и раньше, минимум MSE соответствует условному среднему, поэтому в оптимуме получаем:

u_psi(x,t) approx u^phi(x,t).

Такую модель u_psi(x,t) мы будем называть фейк-моделью — она обучена на сгенерированных (фейковых) данных.

💡 Здесь важно не перепутать три разных поля:

u^*(x,t)

— это фиксированное поле учителя, полученное из настоящих данных (мы выучили его заранее с помощью Flow Matching-лосса);

u^phi(x,t)

— истинное условное среднее поле траекторий генератора (его невозможно вычислить, потому что нам нужно посчитать матож по бесконечному числу семплов);

u_psi(x,t)

— обучаемая фейк-модель, которая приближает u^phi(x,t).

Выпишем отдельно текущие обозначения. К ним можно будет возвращаться при необходимости 😇

Обозначение

Смысл

x_0

шум в пространстве данных, x_0simmathcal{N}(0,I_d)

x_1

настоящий объект из данных, x_1sim p_{text{data}}

x_t

промежуточная точка между x_0 и x_1, x_t=(1-t)x_0 + t x_1

u^*(x,t)

поле реальных данных или поле учителя, предобученное с помощью Flow Matching’а (u_theta(x_t, t))

z

латентный шум генератора, zsimmathcal{N}(0,I_m)

hat{x}_1=G_phi(z)

объект, сгенерированный одношаговым генератором

hat{x}_t

промежуточная точка между x_0 и hat{x}_1, hat{x}_t=(1 - t) x_0 + t hat{x}_1

u^phi(x,t)

истинное среднее поле траекторий генератора (матож, который нельзя просто получить)

u_psi(x,t)

вспомогательная фейк-модель, приближающая u^phi(x,t)

Теперь возникает идея: если u_psi приближает поле генератора — можно сравнивать его с полем учителя u^*. То есть мы хотим, чтобы поле генератора стало ближе к полю учителя. Поэтому генератор должен менять свои выходы так, чтобы траектории от x_0 к G_phi(z) имели такое же среднее поле скоростей, как у учителя.

Тут и появляется игра u_psi и G_phi — первый учит поле, созданное вторым, а второй меняется в зависимости от того, что выучил первый. Можно записать игру между генератором G_phi и вспомогательным полем u_psi:

min_{G_phi}max_{u_psi} mathbb{E}_{t,,x_0simmathcal{N}(0,I_d),,zsimmathcal{N}(0,I_m)} left[ left| u^*(hat{x}_t,t) - (hat{x}_1-x_0) right|_2^2 - left| u_psi(hat{x}_t,t) - (hat{x}_1-x_0) right|_2^2 right].

Рисунок 24. Схема метода. Справа в лоссе — не разность u и v, а разность x_t + (1 - t) u и x_1 для наглядности (так как x_1=x_t + (1 - t) (x_1 - x_0), то с точностью до веса это одно и то же)

Рисунок 24. Схема метода. Справа в лоссе — не разность u и v, а разность x_t + (1 – t) u и x_1 для наглядности (так как x_1 = x_t + (1 – t) (x_1 – x_0), то с точностью до веса это одно и то же)

То есть сначала фейк-модель u_psi максимизирует значение внутри, и как только получается оптимум, мы делаем небольшой шаг оптимизации по G_phi. То есть G_phi должен оптимизировать некоторый максимум по u_psi.

Мы видим, что в новой игре u_psi также учится по лоссу Flow Matching’а на сгенерированных данных:

min_{u_psi} mathbb{E}_{t,,x_0simmathcal{N}(0,I_d),,zsimmathcal{N}(0,I_m)} left[ left| u_psi(hat{x}_t,t) - (hat{x}_1-x_0) right|_2^2 right].

Потому что первый аргумент не зависит от u_psi и max -x=min x.

В полученной игре первое слагаемое сравнивает текущую скорость траектории генератора с полем учителя u^*. Второе слагаемое вычитает ошибку лучшего поля, которое может быть выучено на траекториях самого генератора.

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

Рисунок 25. Сравнение минимаксным лоссом поле учителя u^*, поле генератора u^phi и вспомогательного поля u_psi, чтобы приблизить средний поток генератора к учителю

Рисунок 25. Сравнение минимаксным лоссом поле учителя u^*, поле генератора u^phi и вспомогательного поля u_psi, чтобы приблизить средний поток генератора к учителю
Сравнение минимаксной целью средних полей

Если вспомогательная фейк-модель u_psi идеально выучила поле генератора — минимаксная цель действительно превращается в «выучивание»:

u^phi(x,t) approx u^*(x,t).

Зафиксируем точку hat{x}_t и время t. Обозначим скорость траектории генератора через:

v=hat{x}_1 - x_0.

Тогда поле генератора — условное среднее этой скорости:

u^phi(hat{x}_t,t)=mathbb{E}_{x_0,z} left[ v mid hat{x}_t,t right].

Для любого вектора a(hat{x}_t,t) верно стандартное разложение MSE (мы уже делали это, когда показывали, почему при использовании MSE мы выучиваем условное среднее):

mathbb{E}_{x_0,z} left[ left| a(hat{x}_t,t)-v right|_2^2 mid hat{x}_t,t right]=left| a(hat{x}_t,t) - u^phi(hat{x}_t,t) right|_2^2 + operatorname{tr} operatorname{Cov}_{x_0,z} left( v mid hat{x}_t,t right).

Теперь подставим вместо a поле учителя u^*:

begin{aligned} & mathbb{E}_{x_0,z} left[ left| u^*(hat{x}_t,t)-v right|_2^2 mid hat{x}_t,t right] \ &=left| u^*(hat{x}_t,t) - u^phi(hat{x}_t,t) right|_2^2 + operatorname{tr} operatorname{Cov}_{x_0,z} left( v mid hat{x}_t,t right). end{aligned}

А теперь — вместо a идеальное вспомогательное поле. Если u_psi обучено идеально:

u_psi(hat{x}_t,t)=u^phi(hat{x}_t,t).

Поэтому:

begin{aligned} & mathbb{E}_{x_0,z} left[ left| u_psi(hat{x}_t,t)-v right|_2^2 mid hat{x}_t,t right] \ &=operatorname{tr} operatorname{Cov}_{x_0,z} left( v mid hat{x}_t,t right). end{aligned}

Вычтем второе равенство из первого. Дисперсионные члены сократятся:

begin{aligned} & mathbb{E}_{x_0,z} left[ left| u^*(hat{x}_t,t)-v right|_2^2 - left| u_psi(hat{x}_t,t)-v right|_2^2 mid hat{x}_t,t right] \ &=left| u^*(hat{x}_t,t) - u^phi(hat{x}_t,t) right|_2^2. end{aligned}

Теперь возьмём ожидание по всем t, x_0 и z. Получим:

begin{aligned} & mathbb{E}_{t,hat{x}_t} mathbb{E}_{x_0,z | t, hat{x}_t} left[ left| u^*(hat{x}_t,t)-(hat{x}_1-x_0) right|_2^2 - left| u_psi(hat{x}_t,t)-(hat{x}_1-x_0) right|_2^2 right] \ &=mathbb{E}_{t,hat{x}_t} mathbb{E}_{x_0,z | t, hat{x}_t} left[ left| u^*(hat{x}_t,t) - u^phi(hat{x}_t,t) right|_2^2 right] \ &=mathbb{E}_{t,,x_0,z} left[ left| u^*(hat{x}_t,t) - u^phi(hat{x}_t,t) right|_2^2 right]. end{aligned}

Поэтому такая разность двух MSE-ошибок является удобным способом обучать генератор. Она заставляет каждую отдельную скорость hat{x}_1-x_0 совпадать не u^*, а условные средние поля:

u^phi(x,t) approx u^*(x,t).

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

Рисунок 26. Минимаксная цель убирает разброс отдельных траекторий и оставляет сравнение средних полей

Рисунок 26. Минимаксная цель убирает разброс отдельных траекторий и оставляет сравнение средних полей
Линеаризация

Теперь покажем другой способ получить тот же минимаксный лосс. Этот вывод основан на приёме линеаризации (linearization trick), который используется в нашей статье RealUID.

Напомним, что идеальная цель дистилляции состоит в согласовании двух средних полей:

min_{G_phi} mathbb{E}_{t,,x_0simmathcal{N}(0,I_d),,zsimmathcal{N}(0,I_m)} left[ left| u^*(hat{x}_t,t) - u^phi(hat{x}_t,t) right|_2^2 right],

Где:

hat{x}_1=G_phi(z), qquad hat{x}_t=(1-t)x_0 + that{x}_1,

А поле генератора равно:

u^phi(x,t)=mathbb{E}_{x_0simmathcal{N}(0,I_d),,zsimmathcal{N}(0,I_m)} left[ hat{x}_1-x_0 mid hat{x}_t=x right].

Проблема в том, что здесь стоит разность с условным математическим ожиданием u^phi. Напрямую считать и дифференцировать такое выражение нельзя — нужно бесконечное число раз семплировать.

Запишем разность полей в другой форме. Так как u^*(hat{x}_t,t) при фиксированных hat{x}_t и t является константой, имеем:

begin{aligned} u^*(hat{x}_t,t)-u^phi(hat{x}_t,t) &=u^*(hat{x}_t,t) - mathbb{E}_{x_0,z} left[ hat{x}_1-x_0 mid hat{x}_t,t right] \ &=mathbb{E}_{x_0,z} left[ u^*(hat{x}_t,t)-(hat{x}_1-x_0) mid hat{x}_t,t right]. end{aligned}

Обозначим:

zeta=u^*(hat{x}_t,t)-(hat{x}_1-x_0).

Тогда разность средних полей можно записать как:

u^*(hat{x}_t,t)-u^phi(hat{x}_t,t)=mathbb{E}_{x_0,z} left[ zeta mid hat{x}_t,t right].

Значит, внутри идеальной цели стоит выражение вида:

left| mathbb{E}_{x_0,z} left[ zeta mid hat{x}_t,t right] right|_2^2.

Теперь используем простое тождество (приём линеаризации):

|a|_2^2=max_s left[ -|s|_2^2 + 2langle s,arangle right].

Максимум достигается при s=a, поэтому и слева, и справа производная по a равна 2a (2s=2a для выражения справа). Равенство значений и производных по a показывают, что мы можем эквивалентно подставлять выражение справа вместо выражения слева при подсчете градиентов.

Применим это тождество к:

a=mathbb{E}_{x_0,z} left[ zeta mid hat{x}_t,t right].

Тогда:

left| mathbb{E}_{x_0,z} left[ zeta mid hat{x}_t,t right] right|_2^2=max_s left[ -|s(hat{x}_t,t)|_2^2 + 2 leftlangle s(hat{x}_t,t), mathbb{E}_{x_0,z} left[ zeta mid hat{x}_t,t right] rightrangle right].

Так как s(hat{x}_t,t) зависит только от hat{x}_t и t, его можно внести внутрь условного ожидания:

begin{aligned} & -|s(hat{x}_t,t)|_2^2 + 2 leftlangle s(hat{x}_t,t), mathbb{E}_{x_0,z} left[ zeta mid hat{x}_t,t right] rightrangle \ &=mathbb{E}_{x_0,z} left[ -|s(hat{x}_t,t)|_2^2 + 2 leftlangle s(hat{x}_t,t), zeta rightrangle mid hat{x}_t,t right]. end{aligned}

Теперь в выражении больше нет нормы от условного математического ожидания. Мы заменили её на максимум по вспомогательной функции s, а внутри ожидания осталось линейное выражение по zeta. Именно поэтому этот шаг называется приёмом линеаризации.

Параметризуем вспомогательную функцию через дополнительную нейросеть u_psi:

s_psi(x,t)=u^*(x,t)-u_psi(x,t).

Тогда в оптимуме s_psi должен приближать:

u^*(x,t)-u^phi(x,t),

Значит, u_psi должен приближать поле генератора u^phi.

Подставим:

s_psi(hat{x}_t,t)=u^*(hat{x}_t,t)-u_psi(hat{x}_t,t)

И:

zeta=u^*(hat{x}_t,t)-(hat{x}_1-x_0)

В максимизируемое выражение. Получаем:

begin{aligned} & -left| u^*(hat{x}_t,t)-u_psi(hat{x}_t,t) right|_2^2 \ &quad + 2 leftlangle u^*(hat{x}_t,t)-u_psi(hat{x}_t,t), u^*(hat{x}_t,t)-(hat{x}_1-x_0) rightrangle. end{aligned}

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

a=u^*(hat{x}_t,t), qquad b=u_psi(hat{x}_t,t), qquad v=hat{x}_1-x_0.

Тогда:

-|a-b|_2^2 + 2langle a-b,a-vrangle=|a-v|_2^2 - |b-v|_2^2.

Возвращаемся к исходным обозначениям, получаем:

begin{aligned} & -left| u^*(hat{x}_t,t)-u_psi(hat{x}_t,t) right|_2^2 \ &quad + 2 leftlangle u^*(hat{x}_t,t)-u_psi(hat{x}_t,t), u^*(hat{x}_t,t)-(hat{x}_1-x_0) rightrangle \ &=left| u^*(hat{x}_t,t)-(hat{x}_1-x_0) right|_2^2 - left| u_psi(hat{x}_t,t)-(hat{x}_1-x_0) right|_2^2. end{aligned}

В итоге вместо:

min_{G_phi} mathbb{E}_{t,,x_0simmathcal{N}(0,I_d),,zsimmathcal{N}(0,I_m)} left[ left| u^*(hat{x}_t,t) - u^phi(hat{x}_t,t) right|_2^2 right],

Получаем минимаксную задачу:

min_{G_phi}max_{u_psi} mathbb{E}_{t,,x_0simmathcal{N}(0,I_d),,zsimmathcal{N}(0,I_m)} left[ left| u^*(hat{x}_t,t)-(hat{x}_1-x_0) right|_2^2 - left| u_psi(hat{x}_t,t)-(hat{x}_1-x_0) right|_2^2 right].

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

hat{x}_1-x_0.

Роль u_psi следующая: он выучивает среднее поле текущего генератора. После этого генератор меняется так, чтобы оно приблизилось к полю учителя u^*.

Рисунок 27. Линеаризация — превращение трудной цели со средним полем в удобную игру генератора и вспомогательного поля

Рисунок 27. Линеаризация — превращение трудной цели со средним полем в удобную игру генератора и вспомогательного поля

Алгоритм обучения

Теперь соберём всё в практический алгоритм обучения.

У нас есть три модели:

  1. Предобученное поле учителя u^*(x,t), полученное обычным Flow Matching’ом;

  2. Одношаговый генератор G_phi(z), который мы хотим обучить;

  3. Вспомогательное поле u_psi(x,t), которое приближает среднее поле текущего генератора.

Шаг 1: обновляем вспомогательное поле

Здесь мы хотим обновлять только параметры u_psi. Генератор G_phi используется только для получения текущих сгенерированных объектов, но сам генератор на этом шаге не обновляется.

Семплируем:

x_0 sim mathcal{N}(0,I_d), qquad z sim mathcal{N}(0,I_m), qquad t sim mathcal{U}[0,1].

Затем строим:

hat{x}_1=operatorname{stopgrad}(G_phi(z)), qquad hat{x}_t=(1-t)x_0 + that{x}_1.

Скорость прямой траектории генератора равна:

v=hat{x}_1 - x_0.

Теперь обучаем u_psi предсказывать эту скорость:

mathcal{L}_psi=mathbb{E}_{t,,x_0simmathcal{N}(0,I_d),,zsimmathcal{N}(0,I_m)} left[ left| u_psi(hat{x}_t,t) - v right|_2^2 right].

Это MSE-регрессия. Поэтому при достаточно хорошем обучении фейк-модель u_psi приближает условное среднее поле генератора:

u_psi(x,t) approx u^phi(x,t).

Делаем такие шаги обновления несколько раз.

Шаг 2: обновляем генератор

Здесь мы хотим обновлять только параметры генератора G_phi. Поля u^* и u_psi используются как фиксированные функции.

Семплируем:

x_0 sim mathcal{N}(0,I_d), qquad z sim mathcal{N}(0,I_m), qquad t sim mathcal{U}[0,1].

Теперь считаем выход генератора без stop-gradient:

hat{x}_1=G_phi(z), qquad hat{x}_t=(1-t)x_0 + that{x}_1, qquad v=hat{x}_1 - x_0.

После этого минимизируем лосс генератора:

mathcal{L}_G=mathbb{E}_{t,,x_0simmathcal{N}(0,I_d),,zsimmathcal{N}(0,I_m)} left[ left| u^*(hat{x}_t,t) - v right|_2^2 - left| u_psi(hat{x}_t,t) - v right|_2^2 right].

Интуитивно этот шаг меняет генератор так, чтобы среднее поле его траекторий стало ближе к полю учителя u^*.

Теперь посмотрим на псевдокод (в батчевой форме), чтобы ещё лучше закрепить материал 😎

  • ⚙️ Алгоритм: Дистилляция Flow Matching’а в одношаговый генератор

  • Вход: Предобученное поле учителя u^*, генератор G_phi, вспомогательное поле u_psi, оптимизаторы для phi и psi, количество итераций N, количество шагов обновления u_psi на один шаг генератора K_psi, размер батча B

  • Результат: Одношаговый генератор G_phi

  • Для i=1 до N

    • Шаг 1: обновляем вспомогательное поле u_psi

    • На этом шаге обновляем только u_psi

    • Для k=1 до K_psi

      • Семплируем батч: x_0^{(b)} sim mathcal{N}(0,I_d), qquad z^{(b)} sim mathcal{N}(0,I_m), qquad t^{(b)} sim mathcal{U}[0,1], qquad b=1,ldots,B

      • Считаем выход генератора со stop-gradient: hat{x}_1^{(b)} gets operatorname{stopgrad}(G_phi(z^{(b)})).

      • Строим промежуточные точки: hat{x}_t^{(b)} gets (1-t^{(b)})x_0^{(b)} + t^{(b)}hat{x}_1^{(b)}.

      • Целевые скорости: v^{(b)} gets hat{x}_1^{(b)} - x_0^{(b)}.

      • Считаем средний лосс для u_psi по батчу: mathcal{L}_psi gets frac{1}{B} sum_{b=1}^{B} left| u_psi(hat{x}_t^{(b)},t^{(b)}) - v^{(b)} right|_2^2.

      • Обновляем psi шагом оптимизатора по mathcal{L}_psi

    • Шаг 2: обновляем генератор G_phi

    • На этом шаге обновляем только G_phi Поля u^* и u_psi используются как фиксированные функции Важно: не используем torch.no_grad() для u^*(hat{x}_t^{(b)},t^{(b)}) и u_psi(hat{x}_t^{(b)},t^{(b)}) или .detach()для hat{x}_t^{(b)}.

    • Семплируем батч: x_0^{(b)} sim mathcal{N}(0,I_d), qquad z^{(b)} sim mathcal{N}(0,I_m), qquad t^{(b)} sim mathcal{U}[0,1], qquad b=1,ldots,B

    • Считаем выход генератора без stop-gradient: hat{x}_1^{(b)} gets G_phi(z^{(b)}).

    • Строим промежуточные точки: hat{x}_t^{(b)} gets (1-t^{(b)})x_0^{(b)} + t^{(b)}hat{x}_1^{(b)}.

    • Целевые скорости текущих траекторий генератора: v^{(b)} gets hat{x}_1^{(b)} - x_0^{(b)}.

    • Считаем средний лосс генератора по батчу: mathcal{L}_G gets frac{1}{B} sum_{b=1}^{B} Big[ left| u^*(hat{x}_t^{(b)},t^{(b)}) - v^{(b)} right|_2^2 - left| u_psi(hat{x}_t^{(b)},t^{(b)}) - v^{(b)} right|_2^2 Big].

    • Обновляем phi шагом оптимизатора по mathcal{L}_G

После обучения для генерации больше не нужны ни поле учителя u^*, ни вспомогательное поле u_psi. Мы просто семплируем латентный шум:

z sim mathcal{N}(0,I_m)

И один раз применяем генератор:

hat{x}_1=G_phi(z).

Так многошаговая Flow Matching-генерация заменяется одним шагом.

Следующий скрытый блок — подробный практический алгоритм дистилляции.

Обучение на практике

Рассмотрим нюансы при обучении.

Поле учителя u^* во время дистилляции не обучается. Оно уже знает, как правильно переносить шум в распределение данных. Поэтому мы один раз переводим его в режим eval и замораживаем параметры для скорости:

texttt{u_star.eval()},  qquad texttt{u_star.requires_grad_(False)}.

Однако здесь есть важная тонкость. Заморозить параметры модели — не то же самое, что запретить градиент через её вход. При обновлении генератора параметры u^* не должны меняться, но градиент должен проходить через значение u^*(hat{x}_t,t) к точке hat{x}_t, потому что hat{x}_t зависит от выхода генератора G_phi(z).

Поэтому при обновлении генератора нельзя считать u^*(hat{x}_t,t) внутри torch.no_grad() и применять .detach()для hat{x}_t.

Рисунок 28. Заморозка параметров u^* и u_psi при обновлении генератора, но градиент должен пройти через их вход hat{x}t обратно к Gphi

Рисунок 28. Заморозка параметров u^* и u_psi при обновлении генератора, но градиент должен пройти через их вход hat{x}t обратно к Gphi

💡 Частая ошибка: при обновлении генератора нельзя использовать torch.no_grad() для вычисления u^*(hat{x}_t,t) и u_psi(hat{x}_t,t) или применять .detach()для hat{x}_t. Параметры этих моделей заморожены, но градиент должен проходить через их вход hat{x}_t к генератору G_phi.

В практическом коде величины ниже обычно являются батчами. Например, для изображений x_0, hat{x}_1, hat{x}_t и v имеют форму [B, C, H, W], а время t удобно хранить в форме [B, 1, 1, 1], чтобы оно автоматически broadcast’илось по каналам и пикселям.

Ещё раз приведем алгоритм обучения, но с более подробным техническим описанием.

Шаг 1: обновляем вспомогательное поле

Здесь мы хотим обновлять параметры u_psi. Генератор G_phi используется только для получения текущих сгенерированных объектов, но сам генератор на этом шаге не обновляется.

Поэтому:

  • u_psi переводим в режим train и разрешаем градиенты по его параметрам;

  • G_phi переводим в режим eval и замораживаем параметры;

  • выход генератора hat{x}_1=G_phi(z) считаем со stop-gradient, например, через torch.no_grad() или .detach();

  • u^* на этом шаге не используем.

Семплируем:

x_0 sim mathcal{N}(0,I_d), qquad z sim mathcal{N}(0,I_m), qquad t sim mathcal{U}[0,1].

Затем строим:

hat{x}_1=operatorname{stopgrad}(G_phi(z)), qquad hat{x}_t=(1-t)x_0 + that{x}_1.

Скорость прямой траектории генератора равна:

v=hat{x}_1 - x_0.

Теперь обучаем u_psi предсказывать скорость:

mathcal{L}_psi=mathbb{E}_{t,,x_0simmathcal{N}(0,I_d),,zsimmathcal{N}(0,I_m)} left[ left| u_psi(hat{x}_t,t) - v right|_2^2 right].

Это обычная MSE-регрессия. Поэтому при достаточно хорошем обучении u_psi приближает условное среднее поле генератора:

u_psi(x,t) approx u^phi(x,t).

Шаг 2: обновляем генератор

Здесь мы хотим обновлять только параметры генератора G_phi. Поля u^* и u_psi используются как фиксированные функции.

Поэтому:

  • G_phi переводим в режим train и разрешаем градиенты по его параметрам;

  • u^* оставляем в режиме eval с замороженными параметрами;

  • u_psi переводим в режим eval и замораживаем его параметры;

  • не используем torch.no_grad() вокруг u^*(hat{x}_t,t) и u_psi(hat{x}_t,t);

  • не делаем hat{x}_t.text{detach}();

  • проводим градиент через u^*(hat{x}_t,t), u_psi(hat{x}_t,t), hat{x}_t, hat{x}_1 и далее в параметры G_phi.

Семплируем:

x_0 sim mathcal{N}(0,I_d), qquad z sim mathcal{N}(0,I_m), qquad t sim mathcal{U}[0,1].

Теперь считаем выход генератора без stop-gradient:

hat{x}_1=G_phi(z), qquad hat{x}_t=(1-t)x_0 + that{x}_1, qquad v=hat{x}_1 - x_0.

После этого минимизируем лосс генератора:

mathcal{L}_G=mathbb{E}_{t,,x_0simmathcal{N}(0,I_d),,zsimmathcal{N}(0,I_m)} left[ left| u^*(hat{x}_t,t) - v right|_2^2 - left| u_psi(hat{x}_t,t) - v right|_2^2 right].

Интуитивно этот шаг меняет генератор так, чтобы среднее поле его траекторий стало ближе к полю учителя u^*.

  • ⚙️ Алгоритм: Дистилляция Flow Matching’а в одношаговый генератор

  • Вход: Предобученное поле учителя u^*, генератор G_phi, вспомогательное поле u_psi, оптимизаторы для phi и psi, количество итераций N, количество шагов обновления u_psi на один шаг генератора K_psi, размер батча B

  • Результат: Одношаговый генератор G_phi

  • Переводим u^* в режим eval

    Замораживаем параметры u^*: requires_grad_(False)

  • Для i=1 до N

    • Шаг 1: обновляем вспомогательное поле u_psi

    • На этом шаге обновляем только u_psi Переводим u_psi в режим train

      Разрешаем градиенты для u_psi: requires_grad_(True)

      Переводим G_phi в режим eval

      Замораживаем параметры G_phi: requires_grad_(False)

    • Для k=1 до K_psi

      • Семплируем батч: x_0^{(b)} sim mathcal{N}(0,I_d), qquad z^{(b)} sim mathcal{N}(0,I_m), qquad t^{(b)} sim mathcal{U}[0,1], qquad b=1,ldots,B

      • Считаем выход генератора со stop-gradient: hat{x}_1^{(b)} gets operatorname{stopgrad}(G_phi(z^{(b)})).

      • Строим промежуточные точки: hat{x}_t^{(b)} gets (1-t^{(b)})x_0^{(b)} + t^{(b)}hat{x}_1^{(b)}.

      • Целевые скорости: v^{(b)} gets hat{x}_1^{(b)} - x_0^{(b)}.

      • Считаем средний лосс для u_psi по батчу: mathcal{L}_psi gets frac{1}{B} sum_{b=1}^{B} left| u_psi(hat{x}_t^{(b)},t^{(b)}) - v^{(b)} right|_2^2.

      • Обновляем psi шагом оптимизатора по mathcal{L}_psi

    • Шаг 2: обновляем генератор G_phi

    • На этом шаге обновляем только G_phi Поля u^* и u_psi используются как фиксированные функции Важно: не используем torch.no_grad() для u^*(hat{x}_t^{(b)},t^{(b)}) и u_psi(hat{x}_t^{(b)},t^{(b)})

    • Переводим G_phi в режим train

      Разрешаем градиенты для G_phi: requires_grad_(True) Переводим u_psi в режим eval Замораживаем параметры u_psi: requires_grad_(False) u^* остаётся в режиме eval и с замороженными параметрами

    • Семплируем батч: x_0^{(b)} sim mathcal{N}(0,I_d), qquad z^{(b)} sim mathcal{N}(0,I_m), qquad t^{(b)} sim mathcal{U}[0,1], qquad b=1,ldots,B

    • Считаем выход генератора без stop-gradient: hat{x}_1^{(b)} gets G_phi(z^{(b)}).

    • Строим промежуточные точки: hat{x}_t^{(b)} gets (1-t^{(b)})x_0^{(b)} + t^{(b)}hat{x}_1^{(b)}.

    • Целевые скорости текущих траекторий генератора: v^{(b)} gets hat{x}_1^{(b)} - x_0^{(b)}.

    • Считаем средний лосс генератора по батчу: mathcal{L}_G gets frac{1}{B} sum_{b=1}^{B} Big[ left| u^*(hat{x}_t^{(b)},t^{(b)}) - v^{(b)} right|_2^2 - left| u_psi(hat{x}_t^{(b)},t^{(b)}) - v^{(b)} right|_2^2 Big].

    • Обновляем phi шагом оптимизатора по mathcal{L}_G

После обучения для генерации больше не нужны ни поле учителя u^*, ни вспомогательное поле u_psi. Мы просто семплируем латентный шум:

z sim mathcal{N}(0,I_m)

И один раз применяем генератор:

hat{x}_1=G_phi(z).

Так, многошаговая Flow Matching-генерация заменяется одним шагом.

Ура, мы на финишной прямой! Мы рассмотрели, что такое Flow Matching и как его дистиллировать. Осталось только обсудить детали обучения и посмотреть на результаты! 😎

Пара деталей

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

Что применимо и для обучения, и для дистилляции Flow Matching’а:

  • EMA

На практике мы не просто учим модель, мы сохраняем EMA-версию и уже используем её.

EMA-версия модели — скользящее среднее по параметрам модели. Из-за того, что обучение происходит нестабильно, мы берём среднее последних весов модели.

Это стандартная практика в связке с Adam/AdamW, которая дат стабильность и сильный прирост к качеству. Для начала можно взять EMA = 0.999-0.9999.

  • Gradient clipping

Ещё одна техника — обрезка значений градиентов. Это заметно улучшает стабильность обучения и является стандартом. Норму градиента для начала можно ограничить единицей.

  • Warmup

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

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

Warmup можно поставить, например, на 5000 шагов обучения. У основной модели базовый lr = 2e-4.

  • Dropout

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

Позже мы обсудим, почему это может не сработать при оптимизации u_psi.

Dropout для начала можно поставить 0.0-0.2.

  • Реализация U-Net

Модель U-Net (которую мы использовали в статье RealUID) можно взять из torchcfm. В этом репозитории понятный код, есть веса предобученных моделей для разных задач.

Подробнее про U-Net можно почитать тут.

Специфичные детали для дистилляции:

  • Параметризация генератора

Если латентный шум имеет ту же размерность, что и данные, то есть m=d, генератор удобно параметризовать как G_phi(z)=z + g_phi(z,0) или G_phi(z)=z + g_phi(z,tau), где g_phi — нейросеть той же архитектуры, что и поле Flow Matching’а u_theta, и tau может быть любым. В этом случае g_phi инициализируем весами предобученной Flow Matching-модели.

Если же m neq d — такая остаточная форма z+g_phi(z,0) уже не подходит напрямую, потому что z и G_phi(z) живут в пространствах разной размерности. В этом случае генератор — отображение G_phi : mathbb{R}^m to mathbb{R}^d.

  • Инициализация вспомогательного поля

Вспомогательное поле u_psi удобно инициализировать весами предобученного поля учителя u^*. В начале обучения генератор ещё «плохой», поэтому такая инициализация может сделать оптимизацию более стабильной.

Генератор также можно инициализировать весами u^*.

  • Оптимизаторы

Для обеих сетей можно использовать Adam/AdamW. Для вспомогательного поля u_psi часто полезно отключить первый момент, то есть взять Adam с beta_1=0. Для генератора G_phi — применять Adam с beta_1 ge 0. Такая настройка используется в минимаксных задачах, где одна модель играет роль критика или вспомогательного поля.

Интуиция следующая: beta_1 — momentum, то есть некоторая инерция при оптимизации. И так как внешняя игра очень сильно меняет внутреннюю, то эта инерция будет мешать, то есть модель Flow Matching не успеет за генератором.

Также рекомендуется для начала отключать dropout для вспомогательного поля u_psi — это замедляет обучение модели.

  • Несколько шагов u_psi на один шаг генератора

На практике полезно делать несколько обновлений u_psi на один шаг генератора. Рабочее соотношение: K=5. То есть сначала 5 раз обновляем u_psi, а затем один раз обновляем G_phi. Это помогает вспомогательному полю лучше отслеживать текущее распределение траекторий генератора.

Также вместе с K можно адаптировать lr у u_psi. Вместе с K это позволит u_psi успевать учиться за генератором.

  • Низкая скорость обучения генератора

Стоит ставить lr у генератора поменьше (например, 3e-5), иначе он будет учиться нестабильно из-за сложности лосса.

При этом lr у фейк-модели можно брать как 3e-5, так и 2e-4. Остальные гиперпараметры можно взять стандартные.

  • Заморозка моделей и stop-gradient

Очень важно правильно управлять градиентами. Когда мы обновляем u_psi, генератор G_phi используется только для получения текущих сгенерированных объектов. Поэтому на этом шаге параметры генератора не обновляются: G_phi.eval(), G_phi.requires_grad_(False). Выход генератора можно считать со stop-gradient: hat{x}_1=operatorname{stopgrad}(G_phi(z)).

Например, в PyTorch это можно сделать через torch.no_grad() или .detach(). Когда мы обновляем генератор G_phi, наоборот, параметры u_psi и u^* замораживаются: u_psi.eval(), u_psi.requires_grad_(False), u_star.eval(), u_star.requires_grad_(False).

Но здесь нельзя использовать torch.no_grad() вокруг u^*(hat{x}_t,t) и u_psi(hat{x}_t,t) и делать x_t_hat.detach(). Причина следующая: параметры u^* и u_psi не должны обновляться, но градиенту нужно проходить через их вход hat{x}_t, потому что hat{x}_t=(1-t)x_0 + tG_phi(z) зависит от генератора. Если остановить этот градиент, генератор не получит правильный обучающий сигнал.

  • Постоянная фиксация учителя

Предобученное поле u^* во время всей дистилляции остаётся в режиме eval и не обновляется. Оно играет роль фиксированного учителя, который задаёт правильное направление движения от шума к данным.

Результаты обучения

Рассмотрим результаты обучения и дистилляции Flow Matching-модели на датасете CIFAR-10 — картинки размером 32 х 32 пикселя. Всего есть 10 разных классов — самолёты, машины, птицы и др.

Чтобы проверить, насколько дистилляция работает хорошо, мы будем использовать метрику FID (Fréchet Inception Distance) — она сравнивает набор фото с набором сгенерированных картинок. Чем наборы картинок ближе, тем лучше.

FID-метрика

Сначала по каждому изображению x_i из предобученной нейросети (например, Inception-v3) извлекаются числовые признаки h_i. Затем мы смотрим на эти признаки как на семплированные из многомерного нормального распределения. При этом, признаки реальных и сгенерированных картинок семплируются из двух разных нормальных распределений: mathcal{N(mu, Sigma)} и mathcal{N(hat{mu}, hat{Sigma})}.

Для двух нормальных распределений мы знаем метрику: d_Fleft(mathcal{N}(mu,Sigma),mathcal{N}(hat{mu},hat{Sigma})right)^2=left|mu-hat{mu}right|_2^2+operatorname{tr}left(Sigma+hat{Sigma}-2left(Sigmahat{Sigma}right)^{frac{1}{2}}right).

Среднее и матрицу ковариаций распределения можно оценить по семплам: mu — среднее семплов, матрица ковариации — через Sigma=frac{1}{N-1}(X-mathbf{1}mu)^T(X-mathbf{1}mu).

В статье RealUID FID для Flow Matching’а мы получили равным 3.57 на 100 шагах генерации, при этом после дистилляции FID уменьшился до 2.58. То есть мы не только в несколько раз ускоряем генерацию (1 шаг вместо 100), но и улучшаем её качество.

Мы сгенерировали несколько примеров из дистиллированной модели на CIFAR-10, чтобы увидеть результат:

Рисунок 29. Примеры генерации дистиллированной модели, выученной на CIFAR-10

Рисунок 29. Примеры генерации дистиллированной модели, выученной на CIFAR-10

Также приведём примеры генерации лиц на датасете CelebA размером 64 х 64:

Рисунок 30. Примеры генерации дистиллированной модели, выученной на CelebA

Рисунок 30. Примеры генерации дистиллированной модели, выученной на CelebA

Заключение

Итак, в этом статье мы разобрали, как устроен Flow Matching — переход от шума к данным с помощью векторного поля скоростей.

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

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

Полезные ссылки

  1. Статья про Flow Matching.

  2. Статья про ошибку, которую минимизирует Flow Matching.

  3. TorchCFM — реализация Flow Matching’а.

  4. SiD, DMD, FGM — методы дистилляции.

  5. GAN как min-max игра.

  6. RealUID, code — наша статья про универсальный метод дистилляции с реальными данными + код к ней.

  7. Consistency models, FACM, DuMo, MeanFlow, π-Flow, LADD — другие методы дистилляции и генерации в один или несколько шагов.

  8. Rectified Flows — иной способ ускорить генерацию Flow Matching’а.

  9. Diffusion Meets Flow Matching — пост про связь диффузионных моделей и flow matching’а.

  10. DDPM, Score-Based Generative Modeling through Stochastic Differential Equations — про диффузионные модели (сильно связаны с Flow Matching’ом).

  11. VAE, LVAE, VDVAE, NCSN, Deep Unsupervised Learning using Nonequilibrium Thermodynamics, Variational Diffusion Models, Latent Diffusion, NFDM — для лучшего понимания диффузионных моделей (три взгляда на них — variational, score и flow).

  12. Статья про continuous normalizing flows (может помочь лучше понять Flow Matching).

  13. Лекции BayesGroup по continuous normalizing flows и диффузионным моделям.

  14. The Annotated Transformer — для лучшего понимания устройства используемых моделей.

  15. nn.labml.ai — для понимания реализации различных методов машинного обучения, в том числе — U-Net, Stable Diffusion.

Полезные ссылки на посты DeepSchool

  1. Пост 1 и пост 2 про дистиляцию диффузии

  2. Consistency models

  3. Про Rectified Flows и InstaFlow

  4. Введение в диффузионные модели

  5. Про модели Genie

Flow Matching: обучение и дистилляция - 716

Автор: tixonmavrin

Источник