A PyTorch MXFP8-támogatást épített a FlashAttention-4-be Blackwellhez
A PyTorch kiterjesztette a FlashAttention-4-et az MXFP8 formátumra, a forward és backward számításokban egyaránt. A Blackwell GPU-kra készült megoldás nagy nyelvi modellek alakzatain akár 2,85 PFLOP/s forward és 2 PFLOP/s backward teljesítményt ért el.
- A FlashAttention-4 forward és backward oldalon is MXFP8-támogatást kapott.
- A PyTorch LLM-alakzatokon 2,85 PFLOP/s forward és 2 PFLOP/s backward értéket mért.
- A belső alakzatokon a BF16-hoz képest legfeljebb 1,6-szoros gyorsulást jelentett a forward számítás.
- A fejlesztést a Meta belső GEM-tanítási rendszerében használják.
- A kód nyílt forrásúként elérhető a PyTorch által hivatkozott GitHub-tárházban.
MXFP8-támogatás a tanítás teljes folyamatában
A PyTorch szeptember 16-án ismertetett fejlesztése a FlashAttention-4, röviden FA4, alacsony pontosságú változatára épül. Az MXFP8 támogatása kiterjed az előre- és a visszafelé irányuló számításokra is, így a megoldás a modell tanításának mindkét fontos szakaszában használható.
A PyTorch mérései szerint az LLM-eknél használt alakzatokon az MXFP8-változat 2,85 PFLOP/s forward és 2 PFLOP/s backward teljesítményt ért el. A vállalat belső alakzatain 2,54 PFLOP/s forward és 1,58 PFLOP/s backward értéket mértek. Ez a BF16-hoz képest legfeljebb 1,6-szoros, illetve 1,52-szeres gyorsulást jelentett.
A PyTorch szerint a megoldás egyik első olyan, élvonalbeli MXFP8-os FlashAttention-4-implementációja, amelyet éles tanítási munkafolyamatban használnak. A fejlesztést a Meta belső GEM-tanítási rendszerében alkalmazzák, a kódot pedig nyílt forrásúvá tették a GitHubon.
A Blackwell hardveréhez igazított működés
A Blackwell architektúra tenzormagjai natívan támogatják a blokkonként skálázott mikroskálázási formátumokat, köztük az MXFP8-at, az MXFP6-ot, az MXFP4-et és az NVFP4-et. A PyTorch szerint ezek a műveletek a BF16-alapú MMA-hoz képest 2-4-szeres átviteli sebességet kínálnak.
A gyakorlati használathoz azonban a típus egyszerű lecserélése nem elegendő. A skálafaktorokat a teljesen kihasznált TMEM-ben kell elhelyezni, a kvantálást pedig a GEMM K-dimenziója mentén kell elvégezni minden operandusnál. Ez az online kiszámított köztes értékekre, például a P-re és a dS-re is vonatkozik.
A fejlesztők olyan TMEM-kiosztást alakítottak ki, amely a skálafaktorokat az 512 oszlopos tárterületbe illeszti. Emellett a dS blokkonkénti kvantálásához a Blackwell redux.sync.max.abs.f32 warp-szintű redukcióját használják. A számítási és átalakítási lépéseket többek között összevont RMSNorm+Quantize, valamint GEMM+Quantize kernelekkel hangolták össze.
Kevesebb adatmozgatás a változó hosszúságú sorozatoknál
Az Ads-modellekben fontos szerepet kapnak a változó hosszúságú, úgynevezett jagged adatok. Ezek kezelése az MXFP8 esetében azért nehéz, mert a Blackwell blokkonként skálázott MMA-műveletei 512 bájtos, átrendezett skálafaktor-blokkokkal dolgoznak, amelyek egy 128×128-as adattile-hoz tartoznak.
A PyTorch megoldása nem a teljes adathalmaz kitöltésével kezeli az eltérő hosszúságokat. Csak a jóval kisebb skálafaktor-tenzorokat igazítja 128-cal osztható címekhez, miközben az FP8-adatok tömörek maradnak. A GEMM- vagy attention-kernel a skálafaktoroknál kitöltött eltolásokat, az adatoknál pedig az eredeti jagged eltolásokat használja.
A vállalat szerint ezzel csökkenthető a kitöltésből eredő memóriaforgalom és tárhelyigény. A fejlesztés különösen az olyan tanítási munkafolyamatokban lehet fontos, ahol a változó hosszúságú sorozatok és az alacsony pontosságú számítás egyszerre vannak jelen.