Mixture-of-Experts
Реализация:
llm/src/llm/core/moe.py· классMoE, функцияload_balancing_lossГде используется: Mixtral
Что вы узнаете
Заголовок раздела «Что вы узнаете»- Что такое условные вычисления (conditional computation) и почему смесь экспертов позволяет увеличить число параметров модели почти без роста стоимости одного токена.
- Как устроен слой Mixture-of-Experts (MoE): роутер, выбор top-k экспертов, взвешенная сумма их выходов.
- Почему «softmax по выбранным k» (Mixtral) и «softmax по всем, затем перенормировка» (HuggingFace) — одно и то же, с доказательством.
- Как считать общее и активное число параметров и FLOPs слоя MoE.
- Что такое коллапс роутера, как с ним борется load-balancing loss и почему его градиент идёт через вероятности, а не через доли токенов.
- Как MoE реализован в этом репозитории: алгоритм dispatch/combine циклом по экспертам, softmax роутера во float32, вспомогательный loss в
Trainer.
Предварительные знания
Заголовок раздела «Предварительные знания»- Feed-forward сеть и SwiGLU — feed-forward.md: каждый эксперт в этой главе — обычный SwiGLU-блок.
- Softmax и его производная — notation.md.
- Блок декодера и residual-связи — language-modeling.md.
- Функция потерь и градиентный спуск — training.md (достаточно общего представления).
Зачем условные вычисления
Заголовок раздела «Зачем условные вычисления»В плотном (dense) трансформере каждый токен проходит через все параметры каждого слоя. Стоимость прямого прохода на один токен поэтому пропорциональна числу параметров: грубо, операций с плавающей точкой (FLOPs) на токен для модели с параметрами (одно умножение и одно сложение на каждый вес матрицы). Хотите модель в 8 раз больше — платите в 8 раз больше за каждый токен и при обучении, и при генерации.
Большая часть параметров трансформера сосредоточена в feed-forward сетях (FFN). У модели с и SwiGLU со скрытым размером FFN одного слоя содержит млн параметров, а attention с GQA — около 42 млн (см. подсчёт в mixtral.md).
Условные вычисления (conditional computation) разрывают связь «параметры = стоимость»: сеть содержит много параметров, но для каждого входа включается только их часть, выбранная в зависимости от самого входа. Mixture-of-Experts (смесь экспертов) — самый успешный способ сделать это в трансформерах: вместо одного FFN слой содержит параллельных FFN — экспертов (experts), — а маленькая сеть-роутер (router, gating network) для каждого токена выбирает, к каким из них его отправить.
Интуиция: разные токены требуют разной «обработки» — код, математика, разговорная речь, служебные слова. Один большой FFN должен уметь всё сразу; набор экспертов позволяет разделить эту работу, и каждый токен платит только за тех, кто его действительно обрабатывает.
История: от смеси экспертов к разреженным трансформерам
Заголовок раздела «История: от смеси экспертов к разреженным трансформерам»- 1991 — смесь экспертов. Jacobs, Jordan, Nowlan, Hinton, Adaptive Mixtures of Local Experts (Neural Computation, 3(1), 1991). Несколько сетей-экспертов и gating-сеть, которая выдаёт softmax-распределение по экспертам; выход — взвешенная сумма выходов всех экспертов. Идея — разделение задачи между специалистами; вычисления ещё плотные: считаются все эксперты.
- 2017 — разреженный MoE-слой. Shazeer et al., Outrageously Large Neural Networks: The Sparsely-Gated Mixture-of-Experts Layer (arXiv:1701.06538). Gating становится разреженным: оставляются только наибольших значений (noisy top-k gating — к логитам роутера добавляется гауссов шум, амплитуда которого тоже вычисляется обучаемым слоем), остальные эксперты не вычисляются вовсе. MoE-слой вставлялся между слоями LSTM; модели достигали 137 млрд параметров. Для равномерной загрузки экспертов предложены вспомогательные потери (importance loss и load loss).
- 2020 — GShard. Lepikhin et al., GShard: Scaling Giant Models with Conditional Computation and Automatic Sharding (arXiv:2006.16668). MoE-слой заменяет FFN в каждом втором блоке трансформера; top-2 роутинг, ограниченная ёмкость эксперта (expert capacity), вспомогательный loss балансировки и распределение экспертов по тысячам устройств.
- 2021 — Switch Transformer. Fedus, Zoph, Shazeer, Switch Transformers: Scaling to Trillion Parameter Models with Simple and Efficient Sparsity (arXiv:2101.03961). Упрощение до (каждый токен — ровно к одному эксперту), простая дифференцируемая формула load-balancing loss (разд. 2.2), capacity factor, вычисления роутера во float32 (разд. 2.4, «selective precision»).
- 2024 — Mixtral 8x7B. Jiang et al., Mixtral of Experts (arXiv:2401.04088). Каждый FFN Mistral 7B заменён на 8 SwiGLU-экспертов с top-2 роутингом; все токены обрабатываются (без отбрасывания). Подробно — в mixtral.md.
Слой MoE: определение
Заголовок раздела «Слой MoE: определение»Пусть — вектор одного токена на входе FFN-подслоя (после нормализации). Слой MoE вычисляет
где:
- — выход слоя для этого токена (той же размерности, что вход, — он прибавляется к residual-потоку);
- — число экспертов, — сколько из них выбирается на токен ();
- — матрица роутера; — логиты роутера (router logits), по одному числу на эксперта;
- — множество индексов экспертов с наибольшими логитами;
- — вес (gate) эксперта , ; как именно он считается — в следующем разделе;
- — -я FFN-сеть. В Mixtral и в этом репозитории это SwiGLU (feed-forward.md):
где , — собственные веса эксперта ; — поэлементное произведение.
Всё это применяется независимо к каждому токену: роутер смотрит только на вектор этого токена, соседние позиции не участвуют. Для матрицы из токенов (далее в главе без индекса — всегда число токенов; число параметров обозначается с индексом, например ) слой — это независимых вычислений по формуле выше, и разные токены одного предложения могут попасть к разным экспертам.
Интуиция: если и веса считаются softmax-ом по всем экспертам, получается «плотная» смесь 1991 года — дорого, но гладко. Если , каждый токен целиком отдаётся одному эксперту — дёшево, но выбор жёсткий. Mixtral берёт , : каждый токен смешивает мнения двух специалистов из восьми.
Роутер — один линейный слой (в оригинале без bias). Логит — это скалярное произведение вектора токена со столбцом матрицы : столбец можно понимать как «обучаемый ключ» эксперта , и токен идёт к экспертам, чьи ключи на него больше похожи.
Роутер очень дешёвый: параметров (для Mixtral на слой — в пять тысяч раз меньше одного эксперта). Обучается он вместе со всей сетью обычным градиентным спуском — отдельной разметки «какой эксперт для чего» нет.
В коде: self._router = nn.Linear(emb_size, num_experts, bias=bias) в MoE.__init__, вызов router_logits = self._router(x_flat) в MoE.forward (core/moe.py).
Top-k gating
Заголовок раздела «Top-k gating»Две записи весов
Заголовок раздела «Две записи весов»Встречаются две записи одного и того же правила.
Форма A — softmax по выбранным k (статья Mixtral, эталонный код Mistral, этот репозиторий):
В статье Mixtral это записано как , где оставляет наибольших логитов, а остальные заменяет на (после экспоненты они дают ноль).
Форма B — softmax по всем, затем перенормировка (HuggingFace MixtralSparseMoeBlock):
где — распределение роутера по всем экспертам, а в форме B выбирается как top-k по , а не по .
Эквивалентность
Заголовок раздела «Эквивалентность»Утверждение. Формы A и B выбирают одно и то же множество и дают одинаковые веса .
Доказательство
Шаг 1 — одинаковое множество. Обозначим . Тогда . Функция строго возрастает, поэтому
Порядок экспертов по совпадает с порядком по , и наибольших элементов у них одни и те же (при точных вычислениях; о совпадающих значениях — в «Тонкостях»).
Шаг 2 — одинаковые веса. Подставим в форму B:
Общий знаменатель — сумма по невыбранным экспертам в том числе — сокращается. ∎
Интуиция: softmax задаёт веса с точностью до общего множителя; перенормировка выбрасывает этот множитель. Поэтому неважно, считали ли вы экспоненты невыбранных экспертов: в веса они не попадают.
Эквивалентность нужна на практике: веса HF Mixtral загружаются в модель этого репозитория без изменений роутера, и выходы совпадают (тест llm/tests/models/test_mistral_mixtral_hf_parity.py). Для load-balancing loss (ниже) всё же нужно полное распределение по всем экспертам — форма B естественна там.
Не путайте с вариантом без перенормировки, (Switch Transformer при и часть других моделей; в GShard веса двух выбранных экспертов, наоборот, перенормируются): там сумма весов меньше 1, и выход слоя по норме меньше. Это другая модель, и веса между этими вариантами не переносятся.
Численный пример: E = 4, k = 2
Заголовок раздела «Численный пример: E = 4, k = 2»Логиты роутера для одного токена: .
Top-2 — эксперты 0 и 1. Форма A:
Форма B: , откуда
p = (0.6095, 0.2242, 0.1360, 0.0303) сумма = 1p_0 + p_1 = 0.8337w_0 = 0.6095 / 0.8337 = 0.7311w_1 = 0.2242 / 0.8337 = 0.2689— те же веса. Пусть эксперты на этом токене выдают и (здесь ). Тогда
Эксперты 2 и 3 для этого токена не вычисляются вовсе. Проверка на коде репозитория:
import torchfrom llm.core.moe import MoE
torch.manual_seed(0)moe = MoE(emb_size=2, num_experts=4, top_k_experts=2, dropout=0.0, bias=False)with torch.no_grad(): # логиты роутера для x = (1, 0) — первый столбец матрицы роутера moe._router.weight[:] = torch.tensor([[2.0, 0.0], [1.0, 0.0], [0.5, 0.0], [-1.0, 0.0]])
x = torch.tensor([[[1.0, 0.0]]]) # [B=1, T=1, d=2]y = moe(x)print(moe.router_logits) # tensor([[ 2.0000, 1.0000, 0.5000, -1.0000]], ...)e0, e1 = moe._experts[0](x), moe._experts[1](x)print(torch.allclose(y, 0.7311 * e0 + 0.2689 * e1, atol=1e-4)) # True(nn.Linear хранит матрицу транспонированной, weight имеет форму [E, d], поэтому логиты задаются строками.)
Как через top-k проходит градиент
Заголовок раздела «Как через top-k проходит градиент»— выбор индексов, у него нет производной. Градиент loss языковой модели доходит до роутера только через веса выбранных экспертов:
где — символ Кронекера (1 при , иначе 0). Роутер учится перераспределять вес между уже выбранными экспертами: если эксперт 0 дал более полезный выход, чем эксперт 1, вес растёт. Логиты невыбранных экспертов от основного loss градиента не получают.
Отсюда важное следствие для : в форме A единственный вес — константа, и роутер вообще не обучается от loss языковой модели. Поэтому Switch Transformer при использует (softmax по всем экспертам, без перенормировки) — тогда вес зависит от всех логитов. Shazeer et al. (2017) предполагали, что для обучения роутера нужно , чтобы было что сравнивать; Switch показал, что работает, если вес — это .
Разреженность: параметры и FLOPs
Заголовок раздела «Разреженность: параметры и FLOPs»Общее и активное число параметров
Заголовок раздела «Общее и активное число параметров»Слой MoE хранит всех экспертов, но токен проходит только через из них. Отсюда два числа:
где:
- — параметры роутера (он вычисляется всегда);
- — параметры одного SwiGLU-эксперта без bias (три матрицы );
- общее (total) число — сколько весов нужно хранить в памяти;
- активное (active) число — через сколько весов проходит один токен, то есть что определяет стоимость вычислений.
Во всей модели остальные части — эмбеддинги, attention, нормализации, выходная проекция — общие и активны всегда. Для Mixtral 8x7B (, , , , слоя) расчёт в mixtral.md даёт:
| параметров | |
|---|---|
| всего | млрд |
| активных на токен | млрд |
Это совпадает с цифрами статьи (Jiang et al., 2024): «47B параметров, из которых на токен используется 13B». Название «8x7B» вводит в заблуждение: модель не млрд, потому что размножен только FFN, а attention и эмбеддинги у экспертов общие.
Прямой проход через линейный слой стоит около FLOPs на токен (умножение и сложение на каждый вес). Для FFN-части одного слоя:
(плюс поэлементные SiLU и произведение — , ими пренебрегаем). Для Mixtral : FFN-часть стоит как два плотных FFN, а параметров в ней как в восьми. Вся модель на токен — около GFLOPs (без квадратичной части attention), тогда как плотная модель на 47 млрд параметров стоила бы около 93 GFLOPs.
Чего MoE не экономит
Заголовок раздела «Чего MoE не экономит»- Память. Хранить нужно все экспертов: для инференса Mixtral 8x7B нужна память под 47 млрд параметров, хотя считает он как 13-миллиардная модель.
- Пропускную способность памяти при генерации по одному токену. На маленьком батче разные токены выбирают разных экспертов, и с каждого шага читаются веса почти всех экспертов.
- Коммуникацию. При обучении на многих устройствах эксперты разнесены по ним, и токены приходится пересылать туда, где лежит их эксперт (all-to-all обмен в GShard/Switch).
Схема слоя
Заголовок раздела «Схема слоя»%%{init: {"flowchart": {"rankSpacing": 28, "nodeSpacing": 28}}}%%
flowchart TB
X(["x · один токен"]):::io --> Router["Router<br/>Linear(d → E)"]:::gray
Router --> TopK["top-k логитов"]:::gray
TopK --> W["softmax по выбранным k<br/>→ веса w"]:::purple
TopK -- "индексы экспертов" --> Disp["dispatch:<br/>x → выбранные эксперты"]:::gray
X --> Disp
subgraph Experts[" "]
direction LR
E1["Expert 0<br/>(выбран)"]:::blue
E2["Expert 1"]:::dim
Ed["⋯"]:::dim
En["Expert E−1<br/>(выбран)"]:::blue
end
Disp --> E1
Disp --> En
E1 --> Sum["combine:<br/>взвешенная сумма"]:::gold
En --> Sum
W --> Sum
Sum --> Out(["y"]):::io
style Experts fill:transparent,stroke:#6c8ebf,stroke-dasharray:4 3
classDef io fill:#ffffff,stroke:#999999,color:#1a1a1a;
classDef blue fill:#dae8fc,stroke:#6c8ebf,color:#1a1a1a;
classDef purple fill:#e1d5e7,stroke:#9673a6,color:#1a1a1a;
classDef gray fill:#f5f5f5,stroke:#666666,color:#1a1a1a;
classDef gold fill:#fff2cc,stroke:#d6b656,color:#1a1a1a;
classDef dim fill:#f5f5f5,stroke:#bbbbbb,color:#999999,stroke-dasharray:4 3;
Коллапс роутера и load-balancing loss
Заголовок раздела «Коллапс роутера и load-balancing loss»Проблема
Заголовок раздела «Проблема»Роутер и эксперты учатся одновременно, и в этом есть положительная обратная связь. Пусть в начале обучения роутер случайно чуть чаще выбирает эксперта 0. Эксперт 0 получает больше токенов → больше градиентных шагов → быстрее становится полезным → роутер (градиент через ) ещё сильнее его предпочитает. Остальные эксперты получают мало токенов и почти не учатся. В пределе MoE вырождается в «любимых» экспертов — по сути, плотный FFN, а остальные экспертов — мёртвый груз в памяти. Это называют коллапсом роутера (router collapse); о нём предупреждали уже Shazeer et al. (2017).
При распределённом обучении неравномерность плоха и сама по себе: устройство с перегруженным экспертом считает дольше всех, и остальные его ждут.
Формула
Заголовок раздела «Формула»Стандартное средство — вспомогательный loss балансировки (auxiliary load-balancing loss), который прибавляется к loss языковой модели. Формула Switch Transformer (разд. 2.2) для :
где:
- — число токенов в батче, — номер токена;
- — распределение роутера по всем экспертам для токена ;
- — доля токенов, отправленных к эксперту (фактическая загрузка); ;
- — средняя вероятность эксперта по батчу; ;
- — индикатор (1, если условие верно, иначе 0);
- — коэффициент (в Switch ); множитель делает значение при равномерной загрузке равным 1 независимо от .
Для HuggingFace (load_balancing_loss_func в Mixtral) и этот репозиторий считают долю отдельно для каждой позиции top-k:
где — позиция в top-k (1 — эксперт с наибольшей вероятностью). Для каждого . Удобно обозначить — долю токенов, у которых эксперт вообще попал в top-k; , и
Коэффициент в коде вынесен наружу: load_balancing_loss возвращает без него, а Mixtral.auxiliary_loss() умножает на router_aux_loss_coef.
Значение при равномерной загрузке: k, а не 1
Заголовок раздела «Значение при равномерной загрузке: k, а не 1»Если загрузка равномерна, для всех , и
Заметьте: здесь даже не нужно, чтобы были равны, — достаточно равномерных . Аналогично, если все , то при любой загрузке. У Switch () это значение равно 1, у формулы HF/репозитория с — 2. Некоторые реализации дополнительно делят на ; при сравнении значений aux loss между кодовыми базами это надо учитывать.
Проверка на коде (4 токена, , ; каждый эксперт выбран ровно двумя токенами):
import torchfrom llm.core.moe import load_balancing_loss
balanced = torch.tensor([[2., 1., 0., 0.], [0., 2., 1., 0.], [0., 0., 2., 1.], [1., 0., 0., 2.]])collapsed = torch.tensor([[2., 1., 0., 0.]] * 4)print(load_balancing_loss([balanced], num_experts=4, top_k=2)) # tensor(2.0000)print(load_balancing_loss([collapsed], num_experts=4, top_k=2)) # tensor(3.3392)Во втором случае все токены выбирают экспертов 0 и 1: , , и . В пределе полного коллапса () значение стремится к — максимуму для , .
Почему минимум — при равномерном распределении
Заголовок раздела «Почему минимум — при равномерном распределении»Строго говоря, — функция двух связанных, но разных величин: дискретных долей и гладких вероятностей . Утверждение «минимум при равномерной загрузке» делается для согласованного случая, когда роутер «честен»: фактическая загрузка совпадает с вероятностями. Для это .
Вывод
Пусть для всех , . Тогда
По неравенству Коши — Буняковского для векторов и :
Равенство достигается, только когда векторы пропорциональны, то есть все равны: . Значит, с минимумом ровно при равномерном распределении. Максимум — когда вся масса у одного эксперта ().
Для top-k при (каждый эксперт попадает в top-k пропорционально своей вероятности) то же рассуждение даёт .
Без условия согласованности значение может опуститься ниже . Пример (, ): два токена с идут к эксперту 0, один токен с — к эксперту 1. Тогда , и . Поэтому точнее думать о load-balancing loss не как о функции с минимумом, а как о направлении градиента, которое он задаёт роутеру.
Почему градиент идёт через P, а не через f
Заголовок раздела «Почему градиент идёт через P, а не через f» — это счётчик: доля токенов, у которых или выбрал эксперта . Малое изменение логитов либо не меняет выбора (производная 0), либо меняет его скачком (производная не определена). Поэтому — кусочно-постоянная функция логитов, и в коде она вычисляется из torch.topk и F.one_hot, через которые градиент не проходит. же — среднее гладких softmax-вероятностей, и она дифференцируема.
При обратном проходе ведут себя как константы-«веса» при :
Чем сильнее загружен эксперт, тем сильнее loss наказывает его вероятность. Дойдём до логитов токена . Так как и :
Величина — средняя загрузка экспертов «с точки зрения» токена . Градиентный спуск меняет логит на , то есть уменьшает логиты экспертов с загрузкой выше и увеличивает логиты недогруженных. На следующих шагах top-k начинает чаще выбирать недогруженных экспертов — и так меняется косвенно, через .
Численно для «схлопнувшегося» батча выше (, , ): , и градиент по логитам одного токена
∂L/∂g = (0.6103·(1 − 0.8348), 0.2245·(1 − 0.8348), 0.0826·(0 − 0.8348), 0.0826·(0 − 0.8348)) = (0.1008, 0.0371, −0.0690, −0.0690)— логиты перегруженных экспертов 0 и 1 пойдут вниз, недогруженных 2 и 3 — вверх. Это же значение выдаёт autograd для load_balancing_loss.
Если бы вместо стояла «гладкая» загрузка (то есть loss ), он поощрял бы равномерность вероятностей, но не фактического распределения токенов: роутер мог бы сделать все почти равными, а top-k всё равно выбирал бы одних и тех же экспертов. Множитель привязывает штраф к реальной загрузке.
Статистика по слоям и паддинг
Заголовок раздела «Статистика по слоям и паддинг»В HF и в этом репозитории логиты роутера всех слоёв MoE склеиваются в один тензор [L·N, E], и , считаются по этим строкам сразу. Значит, это не сумма потерь по слоям, а одна общая статистика: при равномерной загрузке она равна , а не . Побочный эффект: перекосы разных слоёв в противоположные стороны (слой 1 перегружает эксперта 0, слой 2 — эксперта 1) частично компенсируют друг друга в общей статистике.
Паддинг в статистику не должен входить: иначе одинаковые pad-токены, которые все идут к одному эксперту, выглядят как перекос. Для этого у load_balancing_loss есть аргумент token_mask: при нём и — средние только по настоящим токенам.
Ёмкость эксперта и отбрасывание токенов
Заголовок раздела «Ёмкость эксперта и отбрасывание токенов»В GShard и Switch Transformer вычисления распределены по устройствам, и каждому эксперту заранее выделяется буфер фиксированного размера — ёмкость эксперта (expert capacity):
где — число токенов в батче (или в группе), — сколько токенов досталось бы эксперту при идеально равномерной загрузке, — capacity factor (коэффициент запаса; в экспериментах Switch — от 1.0 до 2.0). Фиксированный размер нужен, потому что на ускорителях (TPU) формы тензоров должны быть известны заранее.
Если к эксперту пришло больше токенов, лишние отбрасываются (token dropping): для них этот эксперт не вычисляется, и токен проходит слой только по residual-связи (у Switch с — выход FFN для него просто ноль). Пример: , , , → . Если 4 токена выбрали эксперта 0, два из них отброшены.
Большой уменьшает отбрасывание, но тратит память и вычисления на пустые слоты; малый экономит, но теряет токены. Load-balancing loss снижает и эту проблему.
В этом репозитории ёмкости нет, как и в Mixtral (HF MixtralSparseMoeBlock, эталонный код Mistral): реализация без отбрасывания (dropless) — каждый эксперт обрабатывает все выбравшие его токены, сколько бы их ни было. Это возможно, потому что эксперт вызывается на тензоре переменной длины (ниже). Цена — неравномерная загрузка по времени при распределённом обучении, но на результат вычислений она не влияет.
Алгоритм: dispatch и combine
Заголовок раздела «Алгоритм: dispatch и combine»Формула слоя записана для одного токена. Наивная реализация — цикл по токенам: для каждого роутер, top-k и вызовов экспертов на векторе длины . Это медленно: матричные умножения на одном векторе не используют параллелизм.
MoE.forward делает наоборот — цикл по экспертам, и каждый эксперт обрабатывает сразу все свои токены одним вызовом:
X = x.reshape(N, d) # N = B · T токенов; батч и позиция неважныG = X @ W_r # [N, E] логиты роутераtopk_logits, topk_idx = topk(G, k) # [N, k] k лучших экспертов на токенw = softmax(float32(topk_logits)).to(dtype) # [N, k] веса, сумма по k равна 1Y = zeros(N, d)for e in 0 … E−1: tok, slot = where(topk_idx == e) # токены, выбравшие e, и место e в их top-k if tok пуст: continue # эксперт никем не выбран — не считается Y.index_add_(0, tok, w[tok, slot, None] · Expert_e(X[tok]))return dropout(Y).reshape(B, T, d)- Dispatch (рассылка) —
where(topk_idx == e): номера токенов, у которыхeпопал в top-k, и позицияeв их top-k (по ней берётся вес). В top-k одного токена эксперты не повторяются, поэтому токен встречается у эксперта не больше одного раза.X[tok]собирает эти токены в плотную матрицу[n_e, d], где — сколько токенов у эксперта . - Combine (сборка) —
index_add_(0, tok, …): взвешенный выход эксперта прибавляется в строки его токенов. Каждый токен получает ровно слагаемых — от своих экспертов, — и сумма даёт формулу слоя. - Стоимость. : каждый эксперт считает ровно столько строк, сколько токенов его выбрало, итого от «все эксперты на все токены». Python-цикл — итераций на слой, а не .
- Без отбрасывания. Размер
X[tok]определяется во время выполнения, поэтому ёмкость не нужна. Эксперт без токенов пропускается целиком (continue), и его параметры в этом проходе не получают градиента.
Так же устроены MixtralSparseMoeBlock в HuggingFace и MoeLayer в эталонном коде Mistral. Корректность проверяет тест test_matches_naive_per_token_reference в llm/tests/core/test_moe.py: результат совпадает с наивным циклом по токенам.
Пример dispatch на 3 токенах
, , , пусть
topk_idx = [[0, 1], # токен 0 → эксперты 0 и 1 [1, 3], # токен 1 → 1 и 3 [0, 3]] # токен 2 → 0 и 3эксперт e | where(topk_idx == e) → (tok, slot) | вызов эксперта |
|---|---|---|
| 0 | ([0, 2], [0, 1]) | на 2 строках X[[0, 2]] |
| 1 | ([0, 1], [1, 0]) | на 2 строках X[[0, 1]] |
| 2 | ([], []) | не вызывается |
| 3 | ([1, 2], [1, 1]) | на 2 строках X[[1, 2]] |
Всего 6 строк . Токен 0 получает вклады от экспертов 0 (вес w[0, 0]) и 1 (вес w[0, 1]) и т. д.
Softmax роутера во float32
Заголовок раздела «Softmax роутера во float32»Модели часто обучают и запускают в bfloat16 или float16. У bfloat16 всего 8 бит мантиссы (относительная точность около ), и близкие веса экспертов в нём различаются грубо. Switch Transformer (разд. 2.4, selective precision) показал, что вычисления роутера в bfloat16 делают обучение нестабильным, и предложил приводить вход роутера к float32 и вести во float32 всё вычисление внутри роутера (логиты, softmax, выбор экспертов), а результат — возвращать в bfloat16; остальная модель остаётся в bfloat16.
В этом репозитории, как в HF Mixtral, во float32 переводится только softmax весов (логиты роутера считаются в dtype входа):
topk_weights = F.softmax(topk_logits.float(), dim=-1).to(x.dtype) # [N, top_k]В HF Mixtral это softmax(router_logits, dtype=torch.float), эталонный код Mistral поступает так же. Встроенный softmax PyTorch и так накапливает сумму во float32 для половинной точности, так что, например, на CPU результат не меняется; явное приведение делает точность весов независимой от реализации на конкретном backend (тест test_router_softmax_in_float32). load_balancing_loss тоже считает всё во float32 (torch.cat(...).float()).
Реализация в репозитории
Заголовок раздела «Реализация в репозитории»Класс MoE
Заголовок раздела «Класс MoE»llm/src/llm/core/moe.py, MoE(emb_size, num_experts, top_k_experts, dropout=0.1, hidden_dim=None, bias=True):
| Параметр | Смысл |
|---|---|
emb_size | |
num_experts | ; меньше 1 — ValueError |
top_k_experts | ; вне 1 … num_experts — ValueError (при слой молча возвращал бы нули) |
dropout | вероятность единственного dropout — на выходе слоя; эксперты создаются с dropout=0.0, чтобы выход не прореживался дважды |
hidden_dim | каждого эксперта; по умолчанию 4 * emb_size (у Mixtral 8x7B — 14336) |
bias | bias у роутера и всех матриц экспертов; в Mixtral его нет (bias=False) |
Атрибуты: _router (nn.Linear(emb_size, num_experts)), _experts (nn.ModuleList из num_experts блоков SwiGLU с матрицами _gate, _up, _down), _dropout.
Соответствие формулам в MoE.forward:
| Формула | Код |
|---|---|
x_flat = x.reshape(-1, emb_size) | |
router_logits = self._router(x_flat) | |
topk_logits, topk_indices = torch.topk(router_logits, k=self._top_k_experts, dim=-1) | |
| , форма A | topk_weights = F.softmax(topk_logits.float(), dim=-1).to(x.dtype) |
| на своих токенах | self._experts[expert_id](x_flat[token_idx].unsqueeze(0)).squeeze(0) — SwiGLU ждёт [batch, seq, emb], выбранные токены становятся одной «последовательностью» |
output.index_add_(0, token_idx, weights * expert_output) |
Кроме выхода, forward сохраняет self.router_logits ([N, E]) последнего прохода — для load-balancing loss. Логиты хранятся вместе с графом вычислений, поэтому градиент aux loss доходит до роутера.
Функция load_balancing_loss
Заголовок раздела «Функция load_balancing_loss»load_balancing_loss(router_logits, num_experts, top_k, token_mask=None) в том же файле:
logits = torch.cat(list(router_logits), dim=0).float() # [L·N, E] все слои MoE вместеprobs = F.softmax(logits, dim=-1) # p_t — по всем E экспертам_, selected = torch.topk(probs, top_k, dim=-1) # [L·N, k]expert_mask = F.one_hot(selected, num_experts).float() # [L·N, k, E] без градиентаtokens_per_expert = expert_mask.mean(dim=0) # f_{s,i}: [k, E]prob_per_expert = probs.mean(dim=0) # P_i: [E]return num_experts * (tokens_per_expert * prob_per_expert.unsqueeze(0)).sum()(с token_mask средние взвешиваются маской настоящих токенов, повторённой на все слои). Формула совпадает с HF load_balancing_loss_func до точности float — это проверяет llm/tests/core/test_load_balancing_loss.py. Здесь top-k берётся по вероятностям, а в MoE.forward — по логитам; по доказательству выше это одно и то же множество.
Подключение к обучению
Заголовок раздела «Подключение к обучению»Mixtral.auxiliary_loss()(models/mixtral/mixtral.py) собираетdecoder._ff.router_logitsвсех слоёв, вызываетload_balancing_lossс маской настоящих токенов из последнегоforwardи умножает наrouter_aux_loss_coef. При коэффициенте 0 (по умолчанию) возвращаетNone.BaseModel.auxiliary_loss()по умолчанию возвращаетNone— у плотных моделей вспомогательного loss нет.Trainer(training/trainer.py) послеcompute_lm_lossприбавляетself.model.auxiliary_loss(), если он неNone, и делаетbackwardот суммы. При оценке (evaluate) — только loss языковой модели. Так же поступаетHFGPTAdapterвhf-proxy(aux loss только приself.training).
import torchfrom llm.models.mixtral import Mixtral
model = Mixtral({"vocab_size": 1000, "embed_dim": 256, "num_q_heads": 4, "num_kv_heads": 2, "num_layers": 2, "max_position_embeddings": 64, "num_experts": 8, "top_k_experts": 2, "dropout": 0.0, "router_aux_loss_coef": 0.01})ids = torch.randint(0, 1000, (2, 16))logits, _ = model(ids)aux = model.auxiliary_loss() # ≈ 0.01 · 2 в начале обучения: загрузка почти равномернаяТипичные ошибки и тонкости
Заголовок раздела «Типичные ошибки и тонкости»- Совпадающие логиты. При равных логитах
torch.topkвыбирает по своему правилу (обычно меньший индекс). HF берёт top-k по вероятностям во float32, здесь — по логитам в dtype входа; в bfloat16 два близких логита могут округлиться в одно число, и выбор эксперта в редких случаях разойдётся. На случайных весах это не проявляется, но при сравнении в половинной точности иногда видно. - с softmax по выбранным. Вес всегда 1, роутер не получает градиента от основного loss — только от aux loss (см. выше). В этом репозитории
top_k_experts: 1допустим, но для обучения роутера нуженrouter_aux_loss_coef > 0или другая формула весов. - Aux loss не выключен при оценке. Если прибавлять его к loss валидации, перплексия будет завышена.
Trainer.evaluateего не прибавляет. - Сравнение значений aux loss. Равномерная загрузка даёт (HF, здесь) или 1 (Switch, либо реализации с делением на ); коэффициенты 0.01 и 0.001 в разных статьях относятся к разным нормировкам.
- Паддинг в статистике. Без маски pad-токены искажают и .
Mixtral.forwardзапоминает маску изattention_mask;Trainerпередаёт её из батча (датасетыllm/datasetsеё возвращают). Безattention_maskучитываются все токены. - MoE — не ансамбль. Эксперты не обучаются отдельно на разных подзадачах; их «специализация» возникает сама и часто не совпадает с интуитивными темами (см. анализ роутинга в mixtral.md).
- Dropout. В этой реализации один dropout на выходе MoE; в Mixtral 8x7B dropout нет вовсе — для воспроизведения оригинала
dropout: 0.
- MoE заменяет один FFN на экспертов и роутер; каждый токен обрабатывают только экспертов с наибольшими логитами, выход — их взвешенная сумма.
- Softmax по выбранным и softmax по всем с перенормировкой дают одинаковые веса: общий знаменатель сокращается.
- Общее число параметров растёт как , стоимость токена — как : Mixtral 8x7B хранит 46.7 млрд параметров, а считает 12.9 млрд на токен. Память MoE не экономит.
- Без балансировки роутер схлопывается на немногих экспертов. Load-balancing loss равен при равномерной загрузке; его градиент идёт через гладкие и понижает логиты перегруженных экспертов.
- Switch/GShard ограничивают ёмкость эксперта и отбрасывают лишние токены; Mixtral и этот репозиторий обрабатывают все токены (dropless) циклом по экспертам с
index_add_. - Softmax весов роутера считается во float32 (как в HF Mixtral и эталонном коде Mistral).
Вопросы и упражнения
Заголовок раздела «Вопросы и упражнения»-
Логиты роутера , , . Какие эксперты выбраны и с какими весами? Что изменится при ?
Ответ
Выбраны эксперты 1 и 3 (логиты по 3), веса каждый. При добавляется эксперт 2: веса , сумма , то есть для экспертов 1, 3, 2.
-
Докажите, что прибавление одной и той же константы ко всем логитам роутера не меняет ни выбранных экспертов, ни весов. Следует ли отсюда, что bias роутера (одинаковый для всех токенов, но разный для экспертов) бесполезен?
Ответ
Порядок тот же, что у , а — множитель сокращается. Bias роутера — это разные константы для разных экспертов; они меняют порядок и веса (например, заранее сдвигают предпочтение к эксперту с большим ), так что сам по себе он не бесполезен. В Mixtral bias у роутера нет.
-
Посчитайте общее и активное число параметров слоя MoE учебного конфига Mixtral из
mixtral_train.json: , , , , bias включён.Ответ
Эксперт SwiGLU с bias: . Роутер: . Всего: . Активных: — в 4 раза меньше.
-
Почему при и весах «softmax по выбранным» роутер не обучается от loss языковой модели? Что делает Switch Transformer, чтобы этого избежать?
Ответ
Softmax от одного числа всегда равен 1: независимо от логитов, производная по логитам нулевая, а не дифференцируем. Switch использует вес — вероятность эксперта из softmax по всем логитам, без перенормировки; он зависит от всех логитов, и градиент доходит до роутера.
-
В батче из 4 токенов, , : все токены выбрали эксперта 0, . Посчитайте и знак градиента по логитам и токена с .
Ответ
, . . , . Градиентный спуск уменьшит логит перегруженного эксперта 0 и увеличит логит эксперта 1.
-
Батч из токенов, , , capacity factor . Какова ёмкость эксперта? Сколько токенов будет отброшено, если загрузка экспертов ? А в реализации этого репозитория?
Ответ
. У эксперта 0 лишних токена — они отброшены (проходят слой только по residual). Остальные в пределах ёмкости. В этом репозитории ёмкости нет: эксперт 0 обработает все 6 токенов, отбрасывания не происходит.
-
(Код.) Почему в
MoE.forwardможно использоватьindex_add_, а не присваиваниеoutput[token_idx] = ...? Что сломается при присваивании?Ответ
Каждый токен получает вкладов от разных экспертов в разных итерациях цикла; их надо сложить. Присваивание перезаписало бы вклад предыдущего эксперта, и при у токена остался бы только выход последнего по номеру эксперта с весом . Внутри одной итерации индексы
token_idxуникальны (эксперт не повторяется в top-k токена), так что конфликтов записи нет. -
(Исследование.) Обучите учебный Mixtral из
mixtral_train.jsonдважды: сrouter_aux_loss_coef: 0и0.01. После обучения прогоните валидационный текст и постройте гистограмму выбора экспертов по слоям (используйтеdecoder._ff.router_logitsиtorch.topk). Как меняется равномерность и loss языковой модели?
Литература
Заголовок раздела «Литература»- Jacobs, Jordan, Nowlan, Hinton. Adaptive Mixtures of Local Experts. Neural Computation, 3(1), 1991 (на arXiv не публиковалась).
- Shazeer et al. Outrageously Large Neural Networks: The Sparsely-Gated Mixture-of-Experts Layer. 2017. arXiv:1701.06538
- Lepikhin et al. GShard: Scaling Giant Models with Conditional Computation and Automatic Sharding. 2020. arXiv:2006.16668
- Fedus, Zoph, Shazeer. Switch Transformers: Scaling to Trillion Parameter Models with Simple and Efficient Sparsity. 2021. arXiv:2101.03961 — load-balancing loss (разд. 2.2), selective precision (разд. 2.4), capacity factor
- Jiang et al. Mixtral of Experts. 2024. arXiv:2401.04088
- Shazeer. GLU Variants Improve Transformer. 2020. arXiv:2002.05202 — SwiGLU, из которого сделан каждый эксперт