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

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

2026. október 2. 00:26Forrás: PyTorch

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 felülmúlja a FlashAttention-4-et. A Tritonra épülő, TLX-szel vezérelt megoldás rövidebb és a fejlesztők szerint könnyebben bővíthető a kézzel írt kernelimplementációknál.

A lényeg röviden
  • A TLX-es JFA-kernel a Meta GEM modelljének figyelmi műveleteit célozza.
  • A PyTorch szerint az előreirányú teljesítmény körülbelül 13, a visszafelé számításé 50 százalékkal jobb az FA4-nél.
  • A kernel körülbelül 3,2 ezer soros, szemben az FA4 nagyjából 10 ezer soros CuteDSL-kódjával.
  • A méréseket bfloat16 formátumban, NVIDIA B200 gyorsítón végezték.
  • A forráskód elérhető a Meta ads_model_kernel_library GitHub-tárházában.

A változó hosszúságú sorozatok kezelése

A PyTorch október 1-jén közzétett blogbejegyzése a Jagged Flash Attention, vagyis a JFA optimalizálásáról szól. Ez az a figyelmi kernel, amely a Meta Generative Ads Model, röviden GEM, működésében vesz részt. A Meta hirdetési modelljei, köztük a GEM és a Kunlun architektúra, változó hosszúságú felhasználói sorozatokon futtatják a figyelmi műveleteket.

A hagyományos megközelítésben ezeket a sorozatokat azonos hosszúságúra egészítik ki. A PyTorch szerint ez akár a számítási kapacitás 50 százalékának pazarlásával is járhat. A GEM ehelyett folytonosan egymás mellé csomagolja a sorozatokat, a határaikat pedig egy offsets tenzorban tárolja. A JFA közvetlenül ezen a csomagolt Q, K és V adatokon, valamint az eltolásokat tartalmazó tenzoron futtatja a FlashAttention algoritmust, így nincs szükség kitöltött tokenek létrehozására.

A figyelem a GEM leglassabb kernelje. A Blackwell teljesítményének kihasználásához a globális memória betöltését, a softmax műveletet és a mátrixszorzásokat szorosan egymással átfedve kell végrehajtani.

A TLX nagyobb kontrollt ad a Triton fölött

A PyTorch kiindulópontja egy érett, Tritonra épülő JFA-kernel volt. Ez algoritmikusan helyes, de az adatmozgatás és az ütemezés jelentős részét a fordítóra hagyja. A fejlesztők így nem szabályozhatják közvetlenül a megosztott memória kiosztását, a futószalag mélységét, a szinkronizációt vagy a warpok szerepkörét.

A TLX, vagyis a Triton Low-level Extensions ezeket a hardverhez igazodó vezérlési lehetőségeket teszi elérhetővé a Triton magasabb szintű, csempézett programozási modellje mellett. A megoldás explicit SMEM- és TMEM-kiosztást, aszinkron feladatokat, warp-specializációt, akadályokat, TMA- és MMA-műveleteket, valamint Cluster Launch Controlt használ.

A CTA-kban külön warpok végzik a TMA-betöltéseket, a tenzormagos mátrixszorzásokat, a softmaxhoz kapcsolódó számításokat és a kimenet tárolását. A visszafelé számításban külön warp foglalkozik a dQ redukciójával is. A kernel tartósan futó szerkezetet használ, amelyben SM-enként egy CTA több csempén dolgozik. A K és V adatokhoz például háromszoros pufferelést alkalmaznak, a különböző adatfázisok közötti átadást pedig explicit producer-consumer akadályokkal kezelik.

Kevesebb kód, nagyobb teljesítmény

A PyTorch szerint a TLX-es figyelmi kernel körülbelül 3,2 ezer sor tömör Triton-kódból áll. Ez nagyjából háromszor kevesebb a FlashAttention-4, vagyis FA4 2026 májusi változatának állapotuk szerint legkorszerűbb, körülbelül 10 ezer soros CuteDSL-kerneljeinél.

A teljesítményméréseket bfloat16 formátumban, NVIDIA B200 gyorsítón végezték, az összehasonlítás alapja az FA4 volt. A GEM szempontjából fontos, változó hosszúságú alakzatokon a TLX-kernel az előreirányú számításban körülbelül 13 százalékkal, a visszafelé számításban pedig körülbelül 50 százalékkal teljesített jobban az FA4-nél.

A fejlesztők több, különálló optimalizálást építettek a warp-specializált szerkezetre. A változó hosszúságú bemenetek miatt a csempéket a K és V adatok munkaterhelése szerint rendezik, majd váltakozó irányú kiosztással osztják el az SM-ek között. Ez az előreirányú kernelben körülbelül 20 százalékos javulást hozott a beszámoló szerint.

További technikáik között szerepel a broadcast-Q esetén szükséges dQ-redukció többlépcsős pufferelése, a tenzormemória korábbi felszabadítása, a ciklusok egy maszkolás nélküli főrészre és egy rövid, maszkolt végső részre bontása, valamint két CTA együttműködése egy mátrixszorzásban. Ez utóbbi megoldást az FA4 implementációjából vették át.

A kód elérhető a GitHubon

A PyTorch szerint a TLX használatával a nagy teljesítményű Blackwell-kernel fejlesztése közelebb kerül a Tritonban dolgozó modellezési mérnökökhöz. A kód magasabb szintű Triton-modellben marad, ezért a cég állítása szerint olvashatóbb, könnyebben módosítható és egyszerűbben bővíthető új figyelmi változatokkal, például csúszóablakos vagy blokkokra ritkított megoldásokkal.

A projekt forráskódját a Meta ads_model_kernel_library nevű GitHub-tárházának TLX JFA könyvtárában tették közzé. A bemutatott eredmények a GEM broadcast-Q munkaterhelésére és a bejegyzésben ismertetett bfloat16-os B200-mérésekre vonatkoznak.

Kövesd az AI Hírek oldalát a FacebookonA legfontosabb MI-hírek magyarul, rögtön a megjelenés után a hírfolyamodban.Követem

Kapcsolódó hírek

Kutatás
Kutatás2026. október 2. 00:26

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…

A Modal mindenki számára elérhetővé tette a többgépes GPU-klasztereket
Chipek és infrastruktúra2026. október 1.

A Modal mindenki számára elérhetővé tette a többgépes GPU-klasztereket

A Modal általánosan elérhetővé tette a Modal Clusters szolgáltatást, amely több számítási csomópontot kapcsol össze nagy léptékű modelltréninghez és következtetéshez. A…

Kutatás
Kutatás2026. szeptember 16. 20:55

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…