Mesterséges intelligencia, magyarul.
Az eredeti közleményekből.

A PyTorch MXFP8-támogatást épített a FlashAttention-4-be Blackwellhez

2026. szeptember 16.Forrás: PyTorch

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 lényeg röviden
  • 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.

Kapcsolódó hírek

Fejlesztőknek
Fejlesztőknek2026. október 2.

A Helion gyorsíthatja a vLLM LLM-inferenciáját NVIDIA GPU-kon

A PyTorch Helion-alapú lineáris backendet integrált a vLLM-be, hogy automatikus hangolással javítsa a nagy nyelvi modellek következtetési teljesítményét. Az NVIDIA…

Kutatás
Kutatás2026. október 1.

A PyTorch szerint a TLX gyorsabb lett a Blackwell Jagged Flash Attentionjénél

A PyTorch csapata olyan Jagged Flash Attention kernelt mutatott be NVIDIA Blackwell B200 GPU-kra, amely a vállalat mérései szerint a GEM munkaterhelésén gyorsabb volt a…

Kutatás
Kutatás2026. október 1.

A PyTorch TLX-kernellel gyorsította a Meta hirdetési modelljének figyelmét

A PyTorch olyan Jagged Flash Attention-kernelt mutatott be, amely az NVIDIA Blackwell B200 gyorsítón a Meta Generative Ads Model modelljéhez fontos alakzatokon…