MMD есть мера различия между двумя распределениями P, Q. Задается некоторый положительно определенный кернел k, и MMD определяется как:
k(P, Q) =
\mathbb{E}_{x,x'\sim P}[k(x,x')]
+ \mathbb{E}_{y,y'\sim Q}[k(y,y')]
- 2\mathbb{E}_{x\sim P,\,y\sim Q}[k(x,y)]Кернел k может быть линейным, RBF или чем-то еще.
Далее вопрос — что мы хотим туда подставлять и как?
Предлагается использовать признаки замороженной диффузионной модели. Берутся состояния из чистых сэмплов на вход, ибо использование зашумленных не дает профита. Кроме того, в кернел подаются потокенные признаки, а не усредненные, в зависимости от постановки:
- 🌐 Для безусловной генерации в непрерывном пространстве подаются все позиции.
- 💬 Для условной — только позиции ответа.
- 🎭 Для маскированной диффузии — только с позиций маскированных токенов.
И все попарно усредняется по всей последовательности.
MMD можно разбить на три члена:
- 🤝 real–generated — способствует тому, что признаки из распределения P похожи на Q.
- 🎲 generated–generated — способствует разнообразию внутри распределения P.
- 🗑️ real–real — нахрен не сдался, и его можно выкинуть.
Для дискретных диффузионных моделей дифференцировать через категорическое распределение не представляется возможным. Потому используется REINFORCE для оптимизации. Чтобы уменьшить дисперсию градиента, сэмплируется группа MMD-батчей для конкретного примера и оптимизируется policy-gradient суррогат.
Для непрерывных диффузионных моделей важно делать self-conditioning (обусловливание на прошлые выходы модели). В процессе обучения важно подавать такие примеры, backprop делается только через последний шаг генерации. Чтобы уменьшить число шагов на инференсе, дополнительно предлагается iterative refinement distillation (IRD), где студент в один шаг пытается воспроизвести несколькошагового MMD-учителя.
🧪 Эксперименты
Метод валидируется на предобученных MDLM для дискретных моделей и ELF поверх разных энкодеров в непрерывном случае.
• MDLM-MMD достигает лучшего качества, измеряемого по gPPL / Энтропия, против DiDi и IDLM бейзлайнов.
• ELF-MMD / ELF-MMD-IRD также оказывается лучше ELF\*, FMLM\* на разном числе шагов.
Кроме того, выводы обобщаются и на GSM8k для случая условной генерации.
Подход масштабируется и на более-менее серьезные модельки, в частности, DMax поверх LLada-2.0-Mini. MMD-модель на наборе математических и кодовых бенчей не уступает базовой модели и DMax, при этом выдавая больше tokens-per-forward — т. е. обещает более быструю генерацию. И учить недолго — порядка 2 H100 ГПУ-часов.
❤️ Выводы
Мораль сей басни такова — любите MMD, друзья.