Назад к блогу

Как устроен Olmo-core 3: обучение больших MoE-моделей

Как устроен Olmo-core 3: обучение больших MoE-моделей

Olmo-core 3 расширяет одноимённый фреймворк системой обучения, рассчитанной на MoE-модели масштаба триллиона параметров, при этом архитектурные параметры вроде числа экспертов и top-k остаются прежними — меняется именно инфраструктура. Разбор показывает, как устроены конфигурация слоёв и роутера, латентная проекция экспертов и маршрутизация, и почему рост пула экспертов удаётся удерживать от просадки пропускной способности. Материал будет полезен тем, кто хочет понять внутреннюю механику масштабирования MoE, а не только его внешние характеристики.

Olmo-core 3 — это расширение фреймворка Olmo-core системой обучения, спроектированной для гораздо более крупных MoE-моделей. Вопрос, который стоит разобрать: как именно устроена механика маршрутизации, балансировки, параллелизма и точности вычислений, если цель — масштабировать обучение в диапазон триллиона параметров, не потеряв вычислительную эффективность. Нетривиальность в том, что рост пула экспертов почти не должен отражаться на пропускной способности, а это требует согласованной работы роутера, all-to-all обмена токенами и групповых матричных умножений.

Важно сразу отделить архитектуру от инфраструктуры. OlmoE использовал MoE-архитектуру с 64 routed experts, а Olmo 3, напротив, использовал dense-архитектуру. Olmo-core 3 не меняет число экспертов, top-k или активные параметры относительно OlmoE — он расширяет фреймворк системой обучения для гораздо более крупных MoE-моделей. Та же инфраструктура протестирована на более чем одном триллионе общих параметров.

Как задаётся конфигурация MoE

MoEConfig. MoERouterConfigV2. 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, при expert_weight_scale, — на него, а при original_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 у маршрутизируемых экспертов шардирует экспертов по нулевой размерности. Для этого используется подмеш, собранный из двух измерений: одно реплицирует экспертов между устройствами, второе распределяет их между устройствами. Размер второго измерения задаёт число частей, на которые делится пул экспертов, и число экспертов должно делиться на эту величину; частное — это число локальных экспертов на устройстве. Параметры экспертов пересоздаются с ведущей размерностью, равной числу локальных экспертов.

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-only. Обновление кэша сначала синхронизирует анкоры, логическую форму и флаг оптимизатора у обоих хранилищ и при fp8-only снимает требование градиента с весов. Если режим fp8-only и хранилище анкора уже освобождено, кэш не пересчитывается из bf16, а берётся из уже готовых предквантизованных представлений, иначе выбрасывается ошибка; в этом случае версии весов сбрасываются. В обычном случае оба хранилища обновляются из логических весов, предквантизованные представления сужаются до нужного типа, а версии весов запоминаются. Использование 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 на FSDP8 GPU NVIDIA B300, модель 47 млрд параметров52 000 токенов в секунду на GPU против 19 400 — примерно в 2,7 раза выше
MXFP8 против BF164 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 был кратким тестом ёмкости. Ускорение одной части обучения может обернуться издержками в перемещении данных, поэтому сравнения требуют совпадения входных значений, а не только форм матриц.

Где смотреть в коде

Источники

Похожее