Az NVIDIA szerint tízszer gyorsabb lett a dropless MoE tanítása JAX-ban

Az NVIDIA Transformer Engine JAX-szal kombinálva a vállalat mérése szerint 10,4-szeres teljesítménynövekedést hozott a dropless Mixture of Experts modellek tanításában. A DeepSeek-V3 teljesítménye NVIDIA GB200 rendszeren 103-ról 1068 TFLOPS/GPU-ra nőtt.
- A DeepSeek-V3 teljesítménye 103-ról 1068 TFLOPS/GPU-ra nőtt az NVIDIA mérése szerint.
- A Transformer Engine és a JAX együtt 10,4-szeres javulást hozott.
- A dropless MoE minden tokent feldolgoz, kapacitáskorlát és tokeneldobás nélkül.
- A grouped GEMM változó tokenmennyiségekkel is egyetlen kernelhívásban dolgozik.
- A stack 1024 GPU-n 97 százalékos skálázási hatékonyságot ért el DeepSeek-V3 671B tanításakor.
A dropless MoE minden tokent feldolgoz
Az NVIDIA szeptember 14-én közzétett fejlesztői bejegyzése a nagy nyelvi modellek egyik elterjedt architektúrájának, a Mixture of Expertsnek, vagyis MoE-nek a tanítását vizsgálja. Ilyen modellekre példaként a vállalat a DeepSeek, a Qwen és a Mixtral rendszereit említi.
Az MoE-modellekben egy tanult útválasztó dönti el, hogy az egyes tokeneket a sok kisebb szakértői hálózat közül melyik, illetve melyek dolgozzák fel. Ez a feltételes számítás hatékonyabbá teheti a tanítást, mivel minden tokenhez csak a Top-K szakértők aktiválódnak.
A dropless megközelítésben minden token eljut a kiválasztott szakértőhöz, függetlenül attól, mennyire egyenletes a terhelés. Ez a modellminőség szempontjából kedvező lehet, a rendszer számára azonban komoly kihívást jelent. A kapacitásalapú MoE ezzel szemben rögzített tokenkeretet ad az egyes szakértőknek. A kereten felüli tokeneket elhagyhatja, vagy kitöltheti a rendelkezésre álló helyet, ami hiányos adatok feldolgozásához, illetve felesleges számításokhoz vezethet.
Az egyenetlen terhelés lassítja a tanítást
Az útválasztó tanulás közben eltérő preferenciákat alakíthat ki a szakértők iránt. Emiatt egyetlen kötegben is előfordulhat, hogy az egyik szakértő sokkal több tokent kap, mint a másik. Az egyes szakértők tehát eltérő mennyiségű tokent kezelnek, ami szabálytalan, úgynevezett ragged tenzorokat eredményez.
Ez nehezen illeszkedik azokhoz a könyvtárakhoz, amelyek egységes, téglalap alakú adatszerkezetekre optimalizált műveleteket használnak. A szakértők közötti párhuzamosság esetén a tokeneket a megfelelő GPU-khoz kell küldeni, majd a feldolgozás után vissza kell állítani az eredeti sorrendet. Ha a küldés és az összefésülés nincs megfelelően optimalizálva, a GPU-k adatátvitelre várnak, miközben a kommunikáció a számítás mellett kihasználatlan marad.
Az NVIDIA szerint a DeepSeek-V3 NVIDIA GB300 rendszeren futtatott, optimalizálatlan alapváltozata mindössze 103 TFLOPS/GPU teljesítményt ért el, miközben a GPU-k közötti kommunikáció az összesített kernelidő 84 százalékát tette ki.
Csoportosított mátrixszorzás és összevont adatmozgatás
Az NVIDIA Transformer Engine JAX-környezetben több, kifejezetten a változó szakértői terhelések kezelésére szolgáló építőelemet kínál. Ezek között szerepel a csoportokat figyelembe vevő MXFP8-kvantálás, az expert matmul műveletekhez használt MXFP8 grouped GEMM, valamint a dispatch és combine műveletek optimalizált kezelése.
A grouped GEMM az összes szakértő mátrixszorzását egyetlen kernelhívásban kezeli, miközben mindegyiknél a tényleges tokenmennyiséggel számol. Így nem kell a legnagyobb lehetséges kapacitásra méretezett, kitöltött számításokat elvégezni. A Transformer Engine grouped_gemm és ragged_dot megoldásai a cuBLAS és a cuBLASLt könyvtárakra épülnek. Blackwell GPU-kon ez a folyamat MXFP8 blokkos skálázást is lehetővé tesz a szakértői mátrixszorzásoknál.
A dispatch során a tokeneket a hozzájuk rendelt GPU-khoz küldik, a combine pedig visszajuttatja és összegyűjti az eredményeket. Az NVIDIA Transformer Engine ezeket a lépéseket összevont kernelútvonalon kezeli. Az ehhez használt NCCL EP kommunikációs háttér az MoE-re jellemző szabálytalan és egyenetlen forgalomhoz igazodik, és token-duplikációs mechanizmust is alkalmaz.
A vállalat szerint nagyobb rendszereken is skálázódik
Az NVIDIA közlése szerint a JAX és a Transformer Engine célzott kerneloptimalizálása a DeepSeek-V3 esetében 103-ról 1068 TFLOPS/GPU-ra emelte a teljesítményt NVIDIA GB200 hardveren, ami 10,4-szeres javulás. A teljes szoftveres és hardveres stack a DeepSeek-V3 671B tanításakor 1024 GPU-n 97 százalékos skálázási hatékonyságot tartott fenn NVIDIA GB300 NVL72 hardveren.
A fejlesztők az optimalizált JAX MoE-útvonal reprodukálásához kipróbálhatják a Transformer Engine-t engedélyező NVIDIA NGC MaxText-konténert. Az NVIDIA ehhez a MaxText MoE Configuration útmutatót és a Transformer Engine dokumentációját ajánlja. A bejegyzés alapján a megoldás azoknak a fejlesztőknek lehet fontos, akik dropless MoE-modelleket tanítanak, és a változó szakértői terhelés miatt a számítási, memória- vagy kommunikációs szűk keresztmetszeteket próbálják csökkenteni.
NVIDIA Developer: Accelerating Dropless MoE Training in JAX with NVIDIA Transformer Engine


