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 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 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.
