Математика глубокого обучения

Mixture of Experts

Разделить параметры и вычисления — и почему главная проблема здесь статистическая

Шаг 88 из 117 · ~28 мин

Идея: разорвать связь между памятью и вычислениями

В обычной сети число параметров и объём вычислений на токен — это одно и то же число: каждый токен проходит через все веса. MoE эту связь разрывает.

y=etop-kge(x)Ee(x),g(x)=softmax(Wrx)y = \sum_{e \in \text{top-}k} \htmlData{k=gate}{g_e(x)}\, \htmlData{k=expert}{E_e(x)}, \qquad g(x) = \text{softmax}(W_r x)

Сумма идёт только по kk выбранным экспертам из EE, поэтому на токен приходится k/Ek/E доля параметров слоя. При E=8E = 8, k=2k = 2 это 25%25\%: восемь наборов веса, работы — как от двух.

EEkkактивная доля
882225%25\%
881112.5%12.5\%
16164425%25\%
6464223.1%3.1\%

Mixtral 8×7B — прямая иллюстрация: 46.746.7 млрд параметров всего, 12.912.9 млрд активны на токен. Память нужна под все, арифметика — под четверть.

Почему это трудно

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

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

Покрутите разброс логитов роутера. Уверенный роутер — несбалансированный роутер.

12345678

эксперт токенов

вместимость: 20

экспертов на токен k 2
разброс логитов 1
capacity factor 1.25
активная доля
25%
перегрузка
×1.63
balance loss
1.239
отброшено
16
Перегрузка — отношение загрузки самого занятого эксперта к средней. Единица означает идеальный баланс. Balance loss устроен так, что равен единице при равномерном распределении и растёт при перекосе — поэтому его добавляют к основной цели с небольшим весом.

Часть токенов отброшена: эксперт исчерпал вместимость, и для этих токенов слой сработал как тождественное отображение — они прошли по residual-связи без обработки. Ошибки при этом не возникает.

Два числа под графиком стоит понимать точно.

Перегрузка — отношение максимальной загрузки к средней. При идеальном балансе единица; при ×3\times 3 на восьми экспертах треть слоя простаивает, пока один считает за трёх.

Balance loss — вспомогательное слагаемое, добавляемое к основной цели:

Laux=Ee=1EfePe\mathcal{L}_{\text{aux}} = E \sum_{e=1}^{E} f_e \cdot P_e

где fef_e — доля токенов, попавших эксперту ee, а PeP_e — его средний гейт. Устроено так, что при равномерном распределении величина равна единице (каждое слагаемое 1E1E\frac1E \cdot \frac1E, сумма 1E\frac{1}{E}, умноженная на EE), и растёт при перекосе.

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

Capacity factor и отброшенные токены

Каждому эксперту заранее выделяется буфер:

вместимость=токеновkEcapacity factor\text{вместимость} = \frac{\text{токенов} \cdot k}{E} \cdot \text{capacity factor}

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

Заметьте, что происходит: это не ошибка и не исключение. Для отброшенного токена MoE-слой работает как тождественное отображение, и обнаружить это можно только по счётчику. Ещё один тихий режим отказа, и второй в этом уроке.

Capacity factor 1.251.25 означает запас в четверть. Больше — меньше потерь и больше памяти под буферы, меньше — наоборот.

Что ещё ломается

проблемапроявлениеобычное решение
коллапс роутеравсе токены к нескольким экспертамbalance loss
отброшенные токенытихий пропуск слояcapacity factor выше единицы
нестабильность обучениярасходимость на больших моделяхz-loss на логиты роутера
дисбаланс между устройствамиодно ждёт остальныебалансировка + capacity
переобучение экспертовкаждый видит мало данныхбольше данных, меньше EE

Строка про устройства стоит отдельного внимания, потому что она превращает статистику в инженерную проблему. Эксперты обычно распределены по устройствам, и раз шаг синхронный, все ждут самое загруженное. Значит перегрузка ×3\times 3 — это не «треть простаивает в среднем», а шаг втрое дольше. Баланс тут нужен не для качества, а для скорости.

Что MoE даёт и чего не даёт

обычная сетьMoE
параметровNNENэксE \cdot N_{\text{экс}}
вычислений на токенN\propto NkNэкс\propto k \cdot N_{\text{экс}}
памяти под весаN\propto Nпод все эксперты
обмен между устройстваминетall-to-all на каждом слое
качество при равных вычисленияхбазовоеобычно выше
качество при равной памятибазовоеобычно ниже

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

Итог

  • MoE разрывает связь «параметры = вычисления»: активная доля равна k/Ek/E.
  • Главная проблема — статистическая: роутер сходится к вырожденному решению, и отказ тихий.
  • Balance loss равен единице при равномерном распределении и работает через произведение недифференцируемой доли на дифференцируемый гейт.
  • Отброшенные токены — не ошибка, а пропуск слоя; capacity factor задаёт запас.
  • Выгода в вычислениях, плата — памятью и обменом между устройствами.

Источники

Проверки

0 из 2
  1. Разреженность и её цена

    Отметьте все верные утверждения о MoE-слоях.

  2. Баланс, вместимость и потери

    Реализуйте moe_facts(experts, top_k, tokens, capacity_factor, loads), где loads[e] — сколько назначений получил эксперт ee (сумма равна tokens * top_k, потому что каждый токен идёт к top_k экспертам). Верните [active_fraction, imbalance, capacity, dropped, balance]:

    • active_fraction = k/Ek/E — доля параметров слоя, работающая на один токен;
    • imbalance — загрузка самого занятого эксперта, поделённая на среднюю ˉ=tokenskE\bar{\ell} = \frac{\text{tokens} \cdot k}{E};
    • capacity = ˉcapacity factor\lceil \bar{\ell} \cdot \text{capacity factor} \rceil — размер буфера эксперта, вверх до целого;
    • dropped = emax(0, ecapacity)\sum_e \max(0,\ \ell_e - \text{capacity}) — назначения, не поместившиеся в буфер;
    • balance = Eefe2E \sum_e f_e^2, где fe=e/(tokensk)f_e = \ell_e / (\text{tokens} \cdot k) — вспомогательная функция потерь в упрощённом виде (гейт совпадает с распределением).

    Проверить себя можно двумя крайностями: при равномерной загрузке imbalance и balance обе равны единице, а при полном коллапсе на одного эксперта обе равны EE.

    функция moe_facts

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

    Ctrl/⌘ + Enter