A Mixture of Experts (MoE) architektúra fontos szerepet kapott a nagyléptékű AI‑modellek képzésében: DeepSeek, Qwen és Mixtral is MoE megoldások, amelyek sűrű (dense) modelleknél jobb vagy összevethető teljesítményt érnek el kevesebb számítási költséggel. A dropless MoE megközelítés minden tokent a kiválasztott szakértőhöz (expert) továbbít a túlterhelés esetén sem ejt tokeneket, ami modellezési minőség szempontjából kedvező, de rendszer‑oldali optimalizációt követel meg.
Kihívások a MoE képzésben
A MoE‑képzés különböző szűk keresztmetszeteket hoz be, amelyek a sűrű hálózatoknál nem jelentkeznek: dinamikus token‑routing, az expertekhez tartozó dispatch és combine műveletek, mindenfelé (all‑to‑all) kommunikáció és a „ragged” (nem téglalap alakú) GEMM‑ek kezelése. A router tanulása során az expertterhelés erősen ferde lehet: minden batch és minden tokeneloszlás különböző, így az egyes expertekhez érkező tokenek száma változó, ezért nincs egységes, egyszerűen batchelhető GEMM. Ez ragged tensorokat eredményez, amelyekre a legtöbb könyvtár nem optimális.
Ha a dispatch/combine útvonal nincs optimalizálva, a kommunikáció dominálja a kernel időt, a GPU‑k alulhasználttá válnak, és egy rosszul megvalósított all‑to‑all miatt a GPU‑k várakoznak, mielőtt hathatós számítást végeznének.
Mi a különbség dropless és capacity‑based MoE között?
- Dropless MoE: minden token feldolgozásra kerül a kiválasztott expertnél, függetlenül a terhelés egyenlőtlenségétől. Ez jobb modellminőséget enged, de rugalmas, változó token‑számokat megkövetelő kernelmegoldásokat igényel. A MegaBlocks megközelítés például blokkszegény (block‑sparse) mátrixszorzásként fogalmazza újra az expertszámítást, hogy elkerülje a tokenek dobását vagy paddingjét.
- Capacity‑based MoE: fix tokenkapacitást rendelnek minden experthez; túlcsordulás esetén tokent dobnak vagy padolnak. Ez egyszerűsíti a hardverre optimalizált végrehajtást, de kompromisszumot jelent a minőség és a hatékonyság között.
Transformer Engine optimalizációk JAX‑ben a dropless MoE támogatására
A dropless MoE bevezetése megköveteli, hogy minden, az expert‑számítást érintő kernel hatékonyan kezelje a változó token‑számokat, és képes legyen ezekre a formátumokra anélkül, hogy a CPU‑ról folyamatos visszahívásokra lenne szükség (így lehetséges CUDA grafok és elkerülhető a recompilation). A Transformer Engine JAX integrációja a következő építőelemeket nyújtja:
- csoportszintű (group-aware) MXFP8 kvantizáció
- MXFP8 grouped GEMM az expert‑matmulekhez
- optimalizált EP (expert parallelism) műveletek a dispatch és combine részekhez
Grupposított GEMM (Grouped GEMM)
A standard FFN‑ben minden token ugyanazon súlymátrixon megy át, de MoE‑ban az expertekhez érkező tokenek száma változó, így a hagyományos GEMM‑alakok felbomlanak. Korábbi megoldások vagy sorozatos GEMM loopokat használtak (amelyek Device→Host számlálást igényeltek és megtörik a CUDA graph‑ot), vagy worst‑case paddinget alkalmaztak, ami fölösleges számítást jelentett.
A grouped GEMM egyetlen kernelhívásban kezeli az összes expert matmulját az aktuális token‑számokkal, csak a valódi tokenekre számolva. A Transformer Engine grouped_gemm / ragged_dot megvalósítása cuBLAS és cuBLASLt fölé épít, így a Tensor Core‑ok teljes kihasználtsága mellett fut szabálytalan expert‑alakok esetén is. NVIDIA Blackwell architektúrán továbbá engedi az MXFP8 block scaling alkalmazását az expert matmulekhez.
Expert parallelism: Dispatch és Combine integrálása
A router token‑hozzárendelése után a tokeneket fizikailag át kell vinni a megfelelő eszközökre, kiszámolni és visszaállítani az eredeti sorendet. Ezt két fázisra bontjuk:
- Dispatch: tokenek permutálása és átküldése a GPU‑k között az expert hozzárendelés szerint (helyi átrendezés + több‑GPU kommunikáció).
- Combine: a feldolgozott eredmények visszatérése az eredeti GPU‑kra és az expertek eredményeinek összesítése.
A Transformer Engine EP implementációja szorosan összeolvasztja a Dispatch és Combine lépéseket egy fuzionált kernelútvonalba, amit az NCCL EP kommunikációs backend támogat, ami kifejezetten az expert‑paralel routing által generált szabálytalan, kiegyensúlyozatlan hálózati forgalomra van hangolva. Az NCCL EP emellett token‑deduplikációt is alkalmaz: ha egy tokenet több expertnek kell küldeni ugyanazon rangon vagy több rangra egy távoli IB node‑on, a hálózaton csak egyszer megy át, majd a fogadó oldalon replikálódik, ezzel csökkentve a sávszélesség‑igényt.
A grouped GEMM az expert belső számítását kezeli; az EP pedig a külső mozgatást és cserét.
További optimalizációk
- JAX host offloading: köztes aktivációk egy részét host memóriába lehet offloadolni a Heurisztikus rematerializációval, például a query és value projekciók esetén DSv3 edzésben a HBM megtakarítására.
- XLA multistreaming collectives: alapértelmezésben az XLA egyetlen streamen futtatja a kommunikációt, amely sorosítja a párhuzamos collectives‑eket. A multistreaming lehetővé teszi független collectives párhuzamos végrehajtását külön CUDA streameken, átfedésbe hozva InfiniBand és NVLink átviteleket. A Latency Hiding Scheduler (LHS) elemzi a replica‑csoportokat és eldönti, mely collectives biztonságosan fedhetők át, ezáltal csökkentve a kritikus útba kerülő collectives arányát.
Mért teljesítményhatások
- Egy nem optimalizált JAX baseline DeepSeek‑V3 edzés NVIDIA GB200‑on 103 TFLOPS/GPU‑t ért el, és az inter‑GPU kommunikáció a kernel idők 84%-át tette ki.
- Az átfogó Transformer Engine és JAX optimalizációs csomaggal ez 1,068 TFLOPS/GPU‑re nőtt, ami 10.4× javulást jelent a baseline‑hoz képest.
- End‑to‑end, a szerzők által megadott mérések szerint körülbelül 10× throughput javulást értek el a DeepSeek‑V3 671B modellel.
- Nagy‑skálázásnál a rendszer 97% hatékonyságot tartott fenn 1,024 GPU esetén, ami azt mutatja, hogy a kommunikációs optimalizációk sikeresen mérsékelték a multirack skálázódás okozta degradációt.
A szerzők további fejlesztéseket terveznek, például NVFP4 támogatást, kvantizáció összeolvasztását a GEMM‑mel és all‑to‑all átfedést (A2A overlap) a jövőbeni Transformer Engine JAX bindingokban.
Reprodukálás és gyakorlati lépések
Az optimalizációk az NVIDIA NGC MaxText tárolójában érhetők el Transformer Engine‑nel. A szerzők javaslata a következő lépések követése:
- Használja a MaxText referencia konfigurációt és ellenőrizze a helyességet egy kisebb MoE modellnél.
- Fokozatosan skálázza fel, miközben mér lépési időt, TFLOPS/GPU, MFU‑t, grouped GEMM latenciát és a MoE dispatch/combine késleltetését.
A cikk konkrét konfigurációs jelzéseket is ad: a MaxText nemzetközi tárolókból a 2026‑09‑09‑es (ghcr.io/nvidia/jax:maxtext-2026-09-09) vagy újabb konténer használata ajánlott. A példa DeepSeek‑V3 reprodukciós beállítások között szerepelnek modellparaméterek (deepseek3‑671b, max_target_length 4096), tréning beállítások, TE MoEBlock és MXFP8 grouped GEMM opciók, node/parallelism felosztás 128–2,048 GPU‑ra, XLA flag‑ek és környezeti változók (például XLA_PYTHON_CLIENT_MEM_FRACTION: 0.88) — pontos részletek a MaxText MoE configuration guide‑ban találhatók.
Mire érdemes figyelni
- Dropless MoE jobb modellminőséget tart meg, de megköveteli az end‑to‑end stack átalakítását: grouped GEMM‑ek, EP dispatch/combine fuzionálás, kvantizáció és host offload kombinációját.
- Az olyan optimalizációk, mint az NCCL EP token deduplikációja és az XLA multistream collectives, különösen fontosak multirack környezetben.
Köszönetnyilvánítás
A munkához hozzájárult többek között Abhinav Goel, MD Fahim Faysal Khan, Jane Liu, Terry Sun, Tj Xu, Ming Huang, Chase Roberts és Oleg Goncharov (MoE engedélyezés és optimalizáció JAX, XLA, Transformer Engine területen), valamint Artem Polyakov, Ke Wen és Subhadeep Bhattacharya (NCCL EP), illetve Igor Safanov (cuBLASLt) hozzájárulásai.



