В эпоху, когда размер моделей искусственного интеллекта измеряется сотнями миллиардов параметров, архитектура Mixture of Experts (MoE) стала не просто опцией, а необходимым стандартом для достижения передовых результатов. Такие модели, как DeepSeek-V3, Qwen и Mixtral, демонстрируют, что условные вычисления могут превосходить плотные модели по качеству при значительно меньших затратах на обучение. Однако за кажущейся элегантностью этой архитектуры скрываются сложные инженерные вызовы: от динамической маршрутизации токенов до управления неравномерной нагрузкой на графические процессоры. В этой статье мы подробно разберем, как интеграция библиотеки NVIDIA Transformer Engine с фреймворком JAX решает эти проблемы, превращая теоретическую эффективность MoE в реальную вычислительную мощь.
Ключевым прорывом, описанным в последних исследованиях NVIDIA, является переход от традиционных методов маршрутизации к подходу "Dropless MoE" (без отбрасывания токенов). Этот метод сохраняет качество модели, обрабатывая каждый токен без необходимости его отбрасывания или искусственного дополнения (padding). Благодаря оптимизированным ядрам grouped GEMM и ускоренным операциям экспертного параллелизма (Expert Parallelism), производительность обучения модели DeepSeek-V3 на GPU NVIDIA GB200 выросла с 103 до 1068 TFLOPS/GPU — это улучшение в 10,4 раза. Мы рассмотрим технические детали, стоящие за этим достижением, и объясним, как инженеры смогли преодолеть фундаментальные ограничения традиционных библиотек машинного обучения.
01Почему обучение MoE сложнее, чем обучение плотных моделей?
Традиционные плотные модели (Dense Models) обрабатывают все токены через одну и ту же сеть прямого распространения (FFN). Это создает предсказуемую, прямоугольную структуру данных, с которой отлично справляются оптимизированные библиотеки линейной алгебры. В отличие от них, MoE заменяет единую FFN на множество небольших "экспертных" сетей и обучаемый маршрутизатор, который решает, какие Top-K экспертов активировать для каждого токена.
Проблема возникает из-за того, что маршрутизатор является обучаемым компонентом. В процессе обучения распределение токенов между экспертами может стать сильно несбалансированным. Маршрутизатор начинает отдавать предпочтение некоторым экспертам, игнорируя другие. В результате ни одна партия данных (batch) не создает одинаковой нагрузки, и даже внутри одной партии один эксперт может получить в разы больше токенов, чем другой. Это приводит к тому, что каждый эксперт обрабатывает разное количество токенов, что делает невозможным использование стандартных прямоугольных матричных умножений (GEMM). Вместо этого возникают так называемые "рваные" тензоры (ragged tensors).

Рваные тензоры — это структуры данных, где количество элементов в разных измерениях варьируется. Большинство библиотек машинного обучения оптимизированы для однородных, прямоугольных структур, поэтому работа с рваными тензорами требует специализированных подходов. Если не оптимизировать путь маршрутизации и сбора данных, коммуникация между GPU становится доминирующим фактором, заставляя процессоры простаивать в ожидании данных. NVIDIA Transformer Engine был разработан именно для решения этой проблемы, предоставляя ядра, которые нативно поддерживают рваные структуры данных.
02Dropless MoE против Capacity-Based MoE: выбор стратегии
Существует два основных подхода к обработке маршрутизации токенов в MoE: Capacity-Based (на основе емкости) и Dropless (без отбрасывания). Понимание различий между ними критически важно для выбора архитектуры обучения.
В подходе Capacity-Based MoE система ограничивает динамическую маршрутизацию, назначая каждому эксперту фиксированный бюджет токенов. Если эксперт получает больше токенов, чем его лимит, лишние токены либо отбрасываются (drop), либо дополняются до нужного размера (padding). Этот метод сохраняет регулярность вычислений, что удобно для оборудования, но создает компромисс: отбрасывание токенов приводит к обучению на неполных данных, снижая качество модели, а padding ведет к пустой трате вычислительных ресурсов и памяти.

Подход Dropless MoE, напротив, гарантирует, что каждый токен обрабатывается выбранным экспертом, независимо от неравномерности нагрузки. Это привлекательно для качества модели, но требует от системы высокой гибкости. Библиотека MegaBlocks, например, реформатирует вычисления экспертов как блочно-разреженное матричное умножение. Это позволяет каждому эксперту работать с разным количеством токенов без отбрасывания или дополнения. Однако это требует новых специализированных GPU-ядер, оптимизированных grouped GEMM и примитивов маршрутизации, созданных специально для переменных размеров токенов.

03Специализированные оптимизации для Dropless MoE
Переход к Dropless MoE означает, что стек обучения больше не может полагаться на фиксированные формы экспертов. Каждое ядро, участвующее в вычислениях, должно эффективно обрабатывать переменные количества токенов. Более того, поскольку количество токенов для каждого эксперта зависит от данных и динамически меняется, ядра должны принимать динамические формы, которые могут быть недоступны на CPU, что позволяет использовать CUDA Graphs и избегать повторной компиляции.
NVIDIA Transformer Engine предоставляет три ключевых строительных блока для реализации этого подхода в JAX:
- Групповое квантование MXFP8 (Group-aware MXFP8 quantization).
- Grouped GEMM на базе MXFP8 для матричных умножений экспертов.
- Оптимизированные операции экспертного параллелизма (EP) для этапов dispatch и combine.

Рассмотрим слой MoE с экспертным параллелизмом, распределенный между двумя GPU. Маршрутизатор назначает каждый токен эксперту, этап dispatch перемещает токены на GPU выбранного эксперта. Затем выполняется группированное многослойное перцептрон (grouped MLP), который запускает два grouped GEMM для групп переменной длины. Наконец, этап combine возвращает токены в исходный порядок на их родные GPU.
Оптимизация 1: Grouped GEMM
В плотной FFN каждый токен проходит через одну и ту же матрицу весов. В MoE маршрутизатор распределяет токены неравномерно, поэтому каждый эксперт получает разное количество токенов на каждом шаге, что разрушает регулярную форму GEMM, для которой оптимизированы типичные ядра.
Предыдущие подходы включали циклы ядер GEMM или батченные GEMM. Циклы требовали копирования данных с устройства на хост (Device-to-Host) для получения количества токенов. Это находилось на критическом пути, вызывая задержки передачи и ломая CUDA Graphs. Батченные GEMM вычисляли худший случай (максимальную емкость), даже если использовалось меньше токенов, так как они дополнялись до фиксированного размера, что приводило к лишним вычислениям.
Grouped GEMM решает эту проблему, обрабатывая все матричные умножения экспертов в одном вызове ядра, используя фактическое количество токенов для каждого. Он вычисляет только области с действительными токенами, что значительно повышает производительность. Transformer Engine реализует это через grouped_gemm / ragged_dot, используя cuBLAS и cuBLASLt, что обеспечивает полное использование Tensor Core даже при нерегулярных формах экспертов. На GPU NVIDIA Blackwell этот путь также открывает доступ к квантованию MXFP8 с блочным масштабированием.
Оптимизация 2: Интеграция Dispatch и Combine через NCCL EP
После того как ядра маршрутизатора назначили токены экспертам, модель должна физически переместить эти токены на правильные устройства, обработать их и вернуть результаты. Этот процесс делится на два этапа: Dispatch (отправка) и Combine (сборка).
- Dispatch: Токены перемутуются и отправляются через GPU на свои назначенные эксперты. Этот шаг включает как локальную переупорядочивание, так и меж-GPU коммуникацию.
- Combine: Обработанные токены маршрутизируются обратно на их исходные GPU, а результаты от разных экспертов суммируются.
В наивной реализации эти этапы выполняются как последовательная цепочка отдельных операций, что приводит к простоям GPU, многократному чтению/записи данных в память и простоям коммуникационных каналов. Реализация Transformer Engine EP интегрирует Dispatch и Combine в тесно связанное ядро. Эта интеграция поддерживается NCCL EP — бэкендом коммуникации, настроенным специально для нерегулярных, несбалансированных паттернов трафика, характерных для экспертного параллелизма.
NCCL EP также использует механизм дедупликации токенов: когда токен отправляется нескольким экспертам на одном ранге или на несколько рангов на удаленном узле InfiniBand, он проходит по сети только один раз, а репликация происходит на принимающем узле. Это экономит пропускную способность сети. EP является counterpart (противоположной стороной) для grouped GEMM: grouped GEMM обрабатывает то, что происходит внутри каждого эксперта, а EP управляет всем, что происходит вокруг него.

04Дополнительные оптимизации: JAX Host Offloading и XLA Multistreaming
Помимо ядерных оптимизаций, стек JAX и Transformer Engine использует дополнительные методы для снижения узких мест памяти и коммуникации.
JAX Host Offloading
Промежуточные активации не обязательно должны сохраняться на устройстве (GPU) на протяжении всего прямого прохода. JAX предоставляет API для рематериализации, позволяющие выгружать активации в память хоста. Для обучения DSv3 это позволяет выгружать результаты проекций query и value, экономя драгоценную память HBM (High-Bandwidth Memory).
XLA Multistreaming Collectives
Хотя EP управляется через NCCL EP, оптимизированное FSDP (Fully Sharded Data Parallel) обрабатывается нативно в XLA. По умолчанию XLA запускает коммуникацию на одном потоке, поэтому коллективные операции, которые могли бы выполняться параллельно, сериализуются. Multistream collectives позволяют компилятору планировать независимые коллективные операции одновременно через отдельные CUDA потоки. Это позволяет перекрывать передачи между узлами (InfiniBand) с внутриузловой коммуникацией (NVLink), используя оба канала связи одновременно, а не ожидая завершения одного сериализованного потока.
Планировщик скрытия задержек (Latency Hiding Scheduler, LHS) определяет, какие коллективные операции безопасно перекрывать, анализируя их группы реплик и проверяя риск взаимной блокировки (deadlock). Это обеспечивает автоматическое увеличение пропускной способности памяти без необходимости ручной аннотации кода.
05Влияние на производительность обучения
Результаты тестирования на модели DeepSeek-V3 671B показывают впечатляющий прирост пропускной способности в 10 раз. Базовый стек JAX оставлял большую часть потенциала оборудования неиспользованной, где коммуникация между GPU потребляла 84% времени накопления ядер. По мере добавления оптимизаций — cuBLAS GroupedGEMM, XLA multistream collectives, MXFP8 GroupQuant, выгрузки активаций и оптимизированной реализации EP — производительность стабильно росла.
Особое внимание стоит уделить масштабированию на нескольких стойках (multirack scaling). Обучение больших моделей требует обработки триллионов токенов, и на масштабе в тысячи GPU даже небольшие узкие места в вычислениях, памяти и коммуникации могут быстро накапливаться. Обычно коммуникационные накладные расходы растут быстрее, чем вычислительная мощность. Однако с применением стека JAX MoE и Transformer Engine эта деградация остается под контролем. Система сохраняет 97% эффективности масштабирования на 1024 GPU, что свидетельствует о высокой эффективности оптимизаций коммуникации, сохраняющих пропускную способность по мере роста кластера.
06Что это значит на практике
Для исследователей и инженеров, работающих с крупными языковыми моделями, эти оптимизации означают переход от компромиссов к максимальной эффективности. Раньше выбор между качеством модели (Dropless) и скоростью обучения (Capacity-Based) был неизбежным. Теперь, благодаря NVIDIA Transformer Engine и JAX, можно получить лучшее из обоих миров: качество модели, не теряющее токены, и скорость, сопоставимую с плотными моделями.
Для тех, кто хочет воспроизвести эти результаты, NVIDIA предоставляет контейнер NGC MaxText с предварительно включенным Transformer Engine. Начать можно с базовой конфигурации, используя флаги te_moe_block: true и te_gmm_quantization: "te_mxfp8". Для точного воспроизведения результатов DeepSeek-V3 671B требуется расширенная конфигурация с настройками квантования, управления памятью и профилирования, доступная в документации MaxText.
В заключение, интеграция NVIDIA Transformer Engine с JAX не просто ускоряет обучение, она меняет парадигму проектирования больших моделей. Устраняя необходимость в отбрасывании токенов и оптимизируя коммуникацию на уровне ядра, разработчики могут создавать более точные и эффективные модели, используя доступное оборудование с максимальной отдачей. Это открывает путь к следующему поколению AI-систем, где вычислительная эффективность больше не является ограничивающим фактором для качества интеллекта.
Источник: NVIDIA Developer ↗
