Olmo-core 3 — это расширение фреймворка Olmo-core системой обучения, спроектированной для гораздо более крупных MoEMixture-of-Experts, разреженная архитектура, где на каждый токен активируется лишь часть параметров модели-моделей. Вопрос, который стоит разобрать: как именно устроена механика маршрутизации, балансировки, параллелизма и точности вычислений, если цель — масштабировать обучение в диапазон триллиона параметров, не потеряв вычислительную эффективность. Нетривиальность в том, что рост пула экспертов почти не должен отражаться на пропускной способности, а это требует согласованной работы роутера, all-to-all обмена токенами и групповых матричных умножений.
Важно сразу отделить архитектуру от инфраструктуры. OlmoE использовал MoE-архитектуру с 64 routed expertsэксперты, между которыми роутер распределяет токены, а Olmo 3, напротив, использовал dense-архитектуруплотная архитектура, где почти вся модель активна для каждого токена. Olmo-core 3 не меняет число экспертов, top-kчисло экспертов, выбираемых для обработки каждого токена или активные параметры относительно OlmoE — он расширяет фреймворк системой обучения для гораздо более крупных MoE-моделей. Та же инфраструктура протестирована на более чем одном триллионе общих параметров.
Как задаётся конфигурация MoE
MoEConfigкласс конфигурации MoE-слоя: задаёт число экспертов, размер их скрытого слоя и общие модули, а по переданной размерности модели умеет посчитать полное и активное число параметров. MoERouterConfigV2класс конфигурации роутера MoE: хранит размерность модели, число экспертов и top-k как обязательные атрибуты и умеет посчитать свои параметры. MoEConfig описывает слой целиком и опирается на конфигурацию роутера, а MoERouterConfigV2 отвечает только за маршрутизацию.
Число экспертов в MoEConfig задаётся полем num_experts со значением по умолчанию 1, а размер скрытого слоя экспертов — полем hidden_size со значением по умолчанию 256. Размерность модели d_model передаётся отдельно — в методы подсчёта параметров и в сборку. В MoERouterConfigV2 поля d_model, num_experts и top_k задаются напрямую как обязательные атрибуты конфигурации роутера.
Подсчёт параметров устроен так, что общая ёмкость и число активных параметров считаются по разным формулам. Латентная размерность равна d_model, если латентный MoE отсутствует. Общее число параметров складывается из параметров роутера, слагаемого 3 * latent_dim * hidden_size * num_experts и параметров общих модулей, если они заданы. Число активных параметров получается вычитанием всех экспертных параметров и добавлением только доли top-k:
self.num_params(d_model)
- (3 * latent_dim * self.hidden_size * self.num_experts)
+ (3 * latent_dim * self.hidden_size * self.router.top_k)Для роутера число параметров — это d_model * num_experts плюс num_experts, если включён bias. Для маршрутизируемых экспертов число параметров на одного эксперта равно (up_factor + 1) * d_model * hidden_size, где up_factor равен 2 для swiglu и gpt_oss_swiglu и 1 иначе, плюс bias-параметры; всё умножается на num_experts, а для активных параметров — на top_k.
Латентная проекция экспертов
LatentMoEConfig заставляет маршрутизируемую ветвь MoE работать в латентном измерении: latent_dim задаёт входное и выходное измерение routed expert payloads и expert MLPполносвязные слои внутри каждого эксперта, которые обрабатывают назначенные ему токены. Метод resolved_up_proj_input_norm возвращает заданную норму либо, если она не задана, bias-free RMSNorm.
Механика меняет форму данных на входе и выходе экспертов. При наличии латентного MoE создаются понижающая проекция из d_model в latent_dim, опциональная нормализация и повышающая проекция обратно из latent_dim в d_model; при отсутствии — все три модуля равны None. Вход сначала проходит через prepare_routed_input: без понижающей проекции возвращается исходный вход, иначе применяется проекция, и форма входа экспертов меняется с (*, d_model) на (*, latent_dim). Выход экспертов затем проходит restore_routed_output: без повышающей проекции возвращается как есть, иначе при наличии нормы к нему применяется нормализация, результат приводится к dtype веса повышающей проекции и умножается на неё, возвращая форму (*, d_model). Эксперты при этом инициализируются с d_model=latent_dim, то есть работают в латентном измерении, а не в исходном.
Маршрутизация: как роутер выбирает экспертов
Вход формы (batch_size, seq_len, d_model) сначала проходит jitterслучайное возмущение входа роутера при обучении, помогающее избежать закрепления одних и тех же экспертов за одними и теми же токенами, затем при необходимости приводится к fp32. Логиты экспертов вычисляются линейным слоем по весам формы (num_experts, d_model), давая форму (batch_size, seq_len, num_experts). Дальше поведение зависит от функции гейтинга — способа превратить логиты в веса экспертов.
Функции гейтинга различаются по смыслу так. softmax нормирует логиты всех экспертов в распределение вероятностей и использует их как веса. topk_softmax делает то же самое, но только среди выбранных top-k экспертов, поэтому веса выбранных экспертов пересчитываются заново. sigmoid считает каждый эксперт независимо, без конкуренции за общую сумму весов.
Для softmax и topk_softmax scores получаются как softmax логитов по последней оси. Для sigmoid scores — это sigmoid логитов плюс sigmoid_stability_epsilon, если он задан: добавка нужна, чтобы избежать NaN в лоссе балансировки, когда все логиты токена сильно отрицательны и sigmoid даёт нули для всех экспертов. Для topk_softmax индексы выбираются по логитам плюс score_biasобучаемое смещение к логитам экспертов, которым управляет механизм балансировки нагрузки: эксперт, получающий слишком мало токенов, постепенно получает положительный сдвиг и начинает выбираться чаще, если он задан, а веса берутся как softmax от логитов выбранных экспертов. В остальных случаях top-k выбирается по scores (или scores плюс score_bias), а веса собираются gather по исходным scores. При top_k == 1 индексы берутся через max, иначе через topk.
Нормализация весов выполняется делением на норму порядка normalize_expert_weights по последней оси. Затем возможны масштабирования: при restore_weight_scaleмасштабирование весов обратно к исходному порядку величины после нормализации — умножение на top_k, веса умножаются на top_k, при expert_weight_scaleобщий множитель для весов экспертов, задаваемый в конфигурации роутера, — на него, а при original_top_kисходное значение top_k, относительно которого пересчитывается масштаб весов при изменении текущего top_k, отличном от top_k, — на (original_top_k / top_k) ** 0.5.
Разрешение равенства логитов и квантизация
Функция _break_ties добавляет к логитам смещение, пропорциональное индексу эксперта, чтобы при равных значениях выбор top-k был детерминированным и отдавал предпочтение эксперту с меньшим индексом. Функция _quantize_scores округляет логиты до сетки с шагом 1/q, где q = 2**14, что уменьшает чувствительность выбора к малым различиям логитов. Обе применяются только внутри ветки, управляемой флагом UES_QUANT_SCORES, который в коде установлен в False: сначала scores квантизуются, затем разрешается равенство, после чего по обработанным scores выбираются индексы, а веса берутся из исходных scores. То есть при близких или равных логитах они могли бы менять выбор экспертов, но в текущем коде эта ветка отключена.
Jitter при обучении
Шум применяется к входу роутера в самом начале forward. Функция возвращает вход без изменений, если jitter_eps не задан или роутер не в режиме обучения; иначе вычисляет границы 1 ± jitter_eps, генерирует равномерный шум той же формы, что и вход, и поэлементно умножает вход на случайный множитель из этого диапазона. Шум добавляется до вычисления логитов, поэтому влияет на выбор экспертов через scores.
Балансировка нагрузки и вспомогательные штрафы
Счётчик токенов на эксперта формируется в forward: индексы выбранных экспертов гистограммируются, затем суммируются по длине последовательности и по батчу, давая вектор размера num_experts. Этот вектор передаётся в лосс балансировкиштраф, который выравнивает число токенов, назначаемых каждому эксперту, чтобы нагрузка распределялась равномерно вместе с нормированными scores. При сигмоидной гейтинге scores предварительно нормируются делением на сумму по экспертам. Агрегация по процессам: при распределённом режиме счётчики суммируются через all_reduce по группе. Метрика load balancing lossштраф, который выравнивает число токенов, назначаемых каждому эксперту, чтобы нагрузка распределялась равномерно агрегируется с ReduceType.mean, а load imbalance — с ReduceType.max. Накопление идёт прибавлением локального значения, сброс — обнулением.
Помимо балансировки есть два дополнительных штрафа. z_lossштраф, который ограничивает рост логитов роутера, вычисляется из логитов экспертов и делителя лосса аккумулируется и добавляется к общему aux_loss с весом z_loss_weight. orth_lossштраф ортогональности весов роутера, который поощряет независимость направлений весов и не зависит от данных — метод вычисления в базовом классе не реализован и должен быть реализован в подклассах; он масштабируется множителем 1 / loss_div_factor и добавляется к aux_loss с весом orth_loss_weight. Все три штрафа вычисляются только при обучении и включённом градиенте и суммируются в один aux_loss. Различие в назначении: балансировка выравнивает число токенов на эксперта, z_loss штрафует логиты роутера, а orth_loss считается только на весах роутера и не зависит от данных.
Счётчики и bias-корректировка
score_bias_batch_size_per_expert — накопитель числа токенов, назначенных каждому эксперту; он существует только при включённой bias-корректировке и сбрасывается после каждого шага. batch_size_per_expert — отдельный счётчик назначений, который читается для метрик и обнуляется при сбросе. global_batch_size_per_expert существует только при глобальной балансировке и служит для метрики глобального дисбаланса.
Корректировка работает так: при обучении и наличии bias_gamma накопитель при распределённом режиме сначала суммируется по группе, затем вычисляется идеальный средний размер на эксперта, и сдвиг bias_gamma * sign(ideal - actual) добавляется к score_bias, после чего накопитель обнуляется. При наличии score_bias выбор экспертов идёт по scores плюс bias, а веса берутся из исходных scores. Метрика load imbalance — это максимум по счётчику, делённый на среднее, с агрегацией ReduceType.max; глобальная метрика — то же отношение для глобального счётчика с агрегацией ReduceType.mean.
Параллелизм: EP, TP, CP и FSDP
Expert parallelрежим, при котором разные эксперты MoE размещаются на разных устройствах, а токены пересылаются к нужным экспертам у маршрутизируемых экспертов шардирует экспертов по нулевой размерности. Для этого используется подмеш, собранный из двух измерений: одно реплицирует экспертов между устройствами, второе распределяет их между устройствами. Размер второго измерения задаёт число частей, на которые делится пул экспертов, и число экспертов должно делиться на эту величину; частное — это число локальных экспертов на устройстве. Параметры экспертов пересоздаются с ведущей размерностью, равной числу локальных экспертов.
Tensor parallel у MoE делает sequence parallel: вход приводится к Shard(1), роутер получает измерение устройств, между которыми распределяются вычисления тензорного параллелизма, латентные проекции оборачиваются в SequenceParallel с локальным выводом, затем параллелятся эксперты и общий MLP, а выход приводится к Shard(1). Context parallel просто передаёт в роутер измерение устройств, между которыми распределяется обработка длинных последовательностей, — роутер сохраняет его. При включённой expert parallel в forward выбирается parallel_forward_once вместо обычного forward_once.
Подготовка к FSDP и DDP
Методы подготовки к обёртыванию не выполняют обёртывание сами — они делегируют вызов внутреннему модулю. prepare_experts_for_fsdp предназначен для вызова перед обёртыванием модуля в FSDP2, prepare_experts_for_ddp — перед обёртыванием в DDP2. В MoEBase те же методы перенаправляют вызов экспертам. Expert parallel включается отдельно: apply_ep вызывает соответствующий метод у экспертов и выставляет флаг включения. При включённой expert parallel parallel_forward_once реализует ту же математику, что обычный forward, но с expert model parallelism: токены сначала локально переставляются по экспертам, затем переставляются между expert parallel устройствами, затем снова локально, после чего считаются линейные слои и три шага повторяются в обратном порядке.
Прогрев кэша
В базовом классе warmup_cache ничего не делает. В рабочей реализации он сохраняет размер микро-батча, вычисляет ёмкость эксперта и локальную ёмкость и заранее строит и кэширует индексы и бины для параллельной перестановки. Ключи кэша формируются по обеим ёмкостям; если оба тензора уже есть в кэше для устройства, они возвращаются без пересчёта. Индексы строятся через остаток от деления диапазона по числу экспертов и степени шардирования скрытого измерения на число локальных экспертов, повторяются нужное число раз и сортируются; бины — кумулятивная сумма, где каждый локальный эксперт получает ровно ёмкость эксперта. Построенное сохраняется в кэш по ключам.
Пересылка токенов между экспертами
Сначала индексы экспертов приводятся к целому и сортируются, что даёт порядок токенов по экспертам; бины получаются кумулятивной суммой числа токенов на эксперта. Затем выполняется локальная перестановка так, чтобы токены для каждого устройства лежали непрерывно. После этого идёт межdevice-обмен, запускаемый асинхронно; перед обменом при степени шардирования скрытого измерения больше единицы токены повторяются, чтобы устройства, совместно владеющие экспертами, получили все назначенные им токены. После ожидания выполняется ещё одна локальная перестановка с top-k, равным единице, чтобы сгруппировать полученные от разных устройств токены по экспертам. Обратный обмен выполняет all-to-all в другую сторону, а восстановление исходного порядка делает scatter по сохранённым индексам; при шардировании скрытого измерения результат предварительно суммируется по этому измерению.
Ёмкость эксперта и переполнение
Ёмкость эксперта вычисляется как произведение степени expert parallel на допустимое число локальных входов на эксперта, где допустимое число — это округлённое вверх до кратного восьми произведение фактора ёмкости на идеальное число входов; идеальное число входов — это top_k * local_batch_size / num_experts. При переполнении каждый локальный эксперт получает ровно локальную ёмкость токенов, потому что индексы повторяются локальную ёмкость раз, а размеры батчей на эксперта заданы как локальная ёмкость для каждого локального эксперта. Ограничение применяется на этапе локальной перестановки при сборе токенов.
Вычисления экспертов и точность
Локальные вычисления экспертов принимают уже переставленные по устройствам токены и метаданные локальной перестановки, группируют токены по экспертам, выполняют MLP и возвращают результат на исходные позиции. Групповое матричное умножение при включённом torch grouped_mm вычисляет смещения как кумулятивную сумму размеров батчей в int32 и вызывает F.grouped_mm; при trans_b=True правая матрица предварительно транспонируется. Обёртки для квантизованных путей лишь собирают аргументы и делегируют в соответствующие реализации.
Возврат BF16 независимо от dtype параметров и входа объясняется прямо: grouped-gemm эксперты возвращают BF16 вне зависимости от dtype своих параметров и входа, поэтому перед применением плотной повышающей проекции результат приводится к dtype её веса.
Rowwise FP8 и MXFP8
Путь rowwise FP8 строится вокруг двух хранилищ весовобъектов, которые держат веса экспертов в формате FP8 вместе с их масштабами, создаваемых с анкорамипривязанными исходными весами в обычной точности, от которых хранилище обновляет своё содержимое на веса экспертов; оптимизатор для них включается только при режиме fp8-onlyрежиме, в котором параметры хранятся и обучаются только в FP8, без постоянной копии в обычной точности. Обновление кэша сначала синхронизирует анкоры, логическую форму и флаг оптимизатора у обоих хранилищ и при fp8-only снимает требование градиента с весов. Если режим fp8-only и хранилище анкора уже освобождено, кэш не пересчитывается из bf16, а берётся из уже готовых предквантизованных представленийзаранее переведённых в FP8 весов вместе с их масштабами, которые можно сразу подать в матричное умножение, иначе выбрасывается ошибка; в этом случае версии весов сбрасываются. В обычном случае оба хранилища обновляются из логических весов, предквантизованные представления сужаются до нужного типа, а версии весов запоминаются. Использование rowwise FP8 отключается, если режим выключен, конфигурация отсутствует или выключена, вход не на CUDA или его dtype не bf16/fp16/fp32. Сам forward требует swiglu и отсутствия экспертных bias, берёт предквантизованные представления, строит смещения как кумулятивную сумму размеров батчей, при наличии предквантизованного входа формирует соответствующий тип и при fp8-only направляет градиент весов в хранилище. Синхронизация между шагами обеспечивается инвалидацией кэша, которая обнуляет все предквантизованные поля и версии весов и инвалидирует оба хранилища.
SwiGLU с квантизацией
Обёртка квантизации строк вызывает autograd-функцию и помечена как отключающая компилятор. В forward вычисляются три тензора — активация и её квантизованные представления; вход сохраняется только если он требует градиента. В backward градиенты квантизованных представлений отбрасываются; если градиент активации отсутствует, возвращаются нули, иначе по сохранённому входу вычисляется градиент и оборачивается в соответствующий тип. Функция обратного прохода выбирает скомпилированную реализацию, если оба тензора на CUDA и переменная окружения компиляции установлена в одно из значений {"1","true","yes","on"}, иначе — обычную. В обычной реализации последняя размерность делится пополам на up и gate, в float32 считаются sigmoid, произведение gate на sigmoid, производная и два градиента, которые конкатенируются обратно в исходный dtype. Отдельное ограничение: torch grouped_mm не имеет аргумента trans_b и ожидает правую матрицу формы (num_groups, K, N), поэтому при эмуляции trans_b=True правая матрица транспонируется.
Управление состоянием между шагами
Инициализация обнуляет счётчик токенов на эксперта, при глобальной балансировке — глобальный счётчик, обнуляет score_bias, создаёт скрытые аккумуляторы потерь и инициализирует вес и смещение усечённым нормальным распределением. Сброс метрик обнуляет все накопленные значения — счётчики, лосс балансировки, z_loss и orth_loss, — чтобы начать новый интервал измерения; сбор метрик при соответствующем флаге вызывает сброс после сбора значений. Пост-батч вызывается после финального backward полного батча, но до шага оптимизатора: он берёт накопитель, при распределённом обучении суммирует его по группе, вычисляет идеальный средний размер, сдвигает score_bias и обнуляет накопитель.
Мутабельные метрики скрыты от обёрток, чтобы те не регистрировали, не перемещали и не редуцировали их как буферы — это исторический обходной путь для FSDP и composable-DDP; кроме того, torch.compile всё ещё плохо относится к подобной мутации кэша буферов.
Группа процессов и ортогональный лосс
Метод установки группы процессов для балансировки сохраняет переданную группу в поле, и эта группа используется при суммировании глобальных счётчиков токенов на эксперта, когда включена глобальная балансировка. Вычисление ортогонального лосса — заглушка, которая должна быть реализована в подклассах; она вызывается при подсчёте вспомогательного лосса, умножается на множитель 1 / loss_div_factor и добавляется к aux_loss с соответствующим весом. Считать её предлагается только на последнем микро-батче, потому что она вычисляется на весах роутера и не меняется от данных, но при этом считается столько раз, сколько микро-батчей в глобальном батче, из-за чего её приходится масштабировать вниз.
Как сравнивали и что получилось
Замеры проводились на GPU NVIDIA B300. Предварительный тест 47-миллиардной MoE-модели шёл на восьми GPU. Эффект MXFP8 измеряли в контролируемом бенчмарке на четырёх GPU с равномерным распределением работы по экспертам. Масштабный тест модели на 1,2 трлн параметров с 58,36 млрд активных параметров на токен выполнялся на 512 GPU. Эксперимент с DeepEP v2 достигал конфигурации с 2,38 трлн суммарных параметров. Полная конфигурация EP/TP/CP/FSDP и версии PyTorch/CUDA не указаны; ограничения ресурсов не описаны, а приведённая команда запуска задаёт только восемь процессов на узел.
Результаты по сценариям:
| Сценарий | Стенд | Результат |
|---|---|---|
| Рост пула экспертов с 8 до 128 при выборе четырёх на токен | — | активные параметры на токен ~3.2B, общая ёмкость выросла с 4.6B до 47B, пропускная способность упала менее чем на 5% |
| Новый стек против старой реализации MoE на FSDP | 8 GPU NVIDIA B300, модель 47 млрд параметров | 52 000 токенов в секунду на GPU против 19 400 — примерно в 2,7 раза выше |
| MXFP8 против BF16 | 4 GPU NVIDIA B300, равномерная нагрузка по экспертам | пропускная способность примерно на 21% выше, пиковая активная память упала с 103 ГиБ до 95 ГиБ |
| Масштабирование | 512 GPU, модель 1,2 трлн параметров, 58,36 млрд активных на токен | максимальная наблюдённая пропускная способность 858 TFLOP/s/GPU |
| DeepEP v2 | — | конфигурация с 2,38 трлн суммарных параметров |
Оговорки авторов существенны. Тесты в триллионном диапазоне использовали случайный роутинг, чтобы измерить производительность системы, а не качество обученной модели. Эксперимент с DeepEP v2 был кратким тестом ёмкости, а не полным прогоном обучения, то есть показывает достижимый масштаб, а не устойчивую производительность. Замер MXFP8 привязан к конфигурации с четырьмя GPU и равномерной нагрузкой. Авторы также предупреждают, что ускорение одной части обучения создаёт издержки в другом месте: более быстрое вычисление может потребовать большего перемещения данных, а перемещение меньшего числа битов может не помочь, если конвертация занимает слишком много времени. Наконец, сравнения производительности требуют совпадения входных значений, а не только форм матриц.
Инфраструктура обучения и воспроизводимость
Официальные тренировочные скрипты для выпущенных моделей лежат в каталоге src/scripts/official/ и запускаются через torchrun или через Beaker launch CLI при наличии доступа к Beaker. Обучение запускается, например, так:
torchrun --nproc-per-node=8 src/scripts/official/OLMo2/OLMo-2-0325-32B-train.py \
--save-folder=/path/to/save/checkpointsБольшинство опций конфигурации можно переопределять из командной строки, в частности learning rate. Для продолжения annealing с чекпоинта используется отдельный скрипт, которому передаются --save-folder и --checkpoint.
Установка возможна из PyPI командой pip install ai2-olmo-core либо из исходников через pip install -e .[all]. Для отдельных функций нужны опциональные зависимости: flash-attn, ring-flash-attn и TransformerEngine — для соответствующих attention-бэкендов; Liger-Kernel — для low-memory fused-linear реализации лосса; torchao — для float8-обучения; grouped_gemm — для dropless MoE-моделей (возможно, потребуется сборка из исходников до релиза PR #21 после v0.1.6); QuACK — для некоторых CuTe-based ядер. Для transformers требуется версия не ниже 4.57.0, для vLLM — не ниже 0.11.0. Опубликованные Docker-образы содержат все основные и опциональные зависимости, но не включают сам пакет Olmo-core и могут не работать на другом оборудовании или версиях драйвера/CUDA.
Инференс и интеграция
Для загрузки OLMo через transformers нужно установить transformers не ниже 4.57.0, затем загрузить модель и токенизатор по имени allenai/Olmo-3-1125-32B, подготовить входной текст и вызвать генерацию с параметрами max_new_tokens=100, do_sample=True, temperature=1.0, top_p=0.7, после чего декодировать результат. Альтернативно можно использовать абстракцию pipeline. Для vLLM требуется версия не ниже 0.11.0: создаётся объект LLM по тому же имени, задаются параметры сэмплирования temperature=1.0, top_p=0.7, вызывается генерация и выводится текст первого выхода.
Бета-режим генерации реализован как поддержка авторегрессионной генерации напрямую в Olmo-core, на её основе предоставляется демонстрационный чат-цикл. Запуск — командой python -m olmo_core.generate.chat с URL чекпоинта и параметром --max-new-tokens; в примере чекпоинт указывает на step11921, а --max-new-tokens установлен в 512.
Что из этого следует на практике
Разделение общей ёмкости и активных параметров — не косметика, а рабочий инструмент: рост пула экспертов с 8 до 128 при фиксированном top-k увеличивает ёмкость с 4.6B до 47B, тогда как пропускная способность падает менее чем на 5%. Значит, масштабировать модель можно через число экспертов, а не через число активных на токен.
Балансировка нагрузки многослойна, и её части решают разные задачи: счётчики и bias-корректировка выравнивают распределение токенов, z_loss штрафует логиты, orth_loss работает только на весах роутера и не зависит от данных. Поскольку orth_loss считается на каждом микро-батче, но не меняется от данных, его приходится масштабировать вниз — отсюда и предложение считать его только на последнем микро-батче.
Параллелизм требует согласования: expert parallel шардирует экспертов по нулевой размерности и требует делимости их числа на степень шардирования, sequence parallel приводит вход и выход к шардированию по измерению последовательности, а all-to-all обмен токенами выполняется в несколько перестановок с ограничением ёмкости эксперта. Ёмкость эксперта — это порог, за которым лишние токены отбрасываются на этапе локальной перестановки, поэтому фактор ёмкости напрямую влияет на то, сколько назначений реально дойдёт до экспертов.
Точность вычислений неоднородна: grouped-gemm эксперты возвращают BF16 независимо от dtype параметров и входа, поэтому перед плотной проекцией результат приводится к нужному dtype. Пути rowwise FP8 и MXFP8 требуют явной синхронизации анкоров и кэшей между шагами, а при режиме fp8-only анкоры освобождаются и кэш обновляется из предквантизованных представлений, а не из bf16.
Наконец, заявленные цифры описывают производительность системы, а не качество модели: тесты в триллионном диапазоне использовали случайный роутинг, а тест DeepEP v2 был кратким тестом ёмкости. Ускорение одной части обучения может обернуться издержками в перемещении данных, поэтому сравнения требуют совпадения входных значений, а не только форм матриц.
Где смотреть в коде
- router.py: set_load_balancing_process_group
- parallel_mlp.py: _get_parallel_indices_and_bins
- router.py: reset_metrics
- moe.py: prepare_routed_input
- parallel_mlp.py: apply_tp
- routed_experts.py: _use_rowwise_fp8
- parallel_mlp.py: reverse_all_to_all
- router.py: get_top_k