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 FlashAttention-4 májusi változatánál. A TLX-re épülő megoldás a fejlesztők szerint jóval kevesebb kóddal valósítja meg a Blackwell hardverének alacsony szintű vezérlését.
- A TLX-es Jagged Flash Attention NVIDIA B200 GPU-kra készült.
- A PyTorch mérése szerint a kernel 13 százalékkal gyorsabb volt forward, 50 százalékkal backward számításban az FA4-nél.
- A kernel körülbelül 3,2 ezer sor kódból áll, az FA4 mintegy 10 ezer soros megoldásaival szemben.
- A JFA kitöltés nélkül dolgozik a változó hosszúságú felhasználói sorozatokon.
- A kód elérhető a Meta ads_model_kernel_library GitHub-tárházában.
A változó hosszúságú sorozatok kezelése a cél
A PyTorch október 1-jén közzétett blogbejegyzése a Jagged Flash Attention, röviden JFA optimalizálását mutatja be. Ez az a figyelmi kernel, amelyet a Meta Generative Ads Model, vagyis GEM használ. A bejegyzés szerint a figyelem számítása a GEM leglassabb kernelje.
A Meta hirdetési modelljei változó hosszúságú, úgynevezett jagged felhasználói sorozatokon futtatják a figyelmet. A sorozatok rögzített hosszúságra történő kitöltése a GEM korábbi, rendszer-szintű bemutatása szerint akár a számítási kapacitás 50 százalékát is elpazarolhatja. A GEM ezért a sorozatokat egymás mellé, folytonosan tárolja, a határaikat pedig egy offsets tenzor rögzíti.
A JFA közvetlenül erre a kitöltés nélküli, tömörített Q, K és V adatreprezentációra alkalmazza a FlashAttention algoritmust. Így a rendszernek nem kell kitöltött tokeneket létrehoznia.
A TLX hardverközelibb irányítást ad a Tritonhoz
A Triton magas szintű, csempékre épülő programozási modellje sok döntést a fordítóra bíz. A PyTorch szerint ez korlátozza a megosztott memória, a csővezetékek, a szinkronizáció és a warpok működésének közvetlen irányítását. A TLX, vagyis Triton Low-level Extensions ezeket a lehetőségeket első osztályú primitívekként teszi elérhetővé.
A megoldás explicit SMEM- és TMEM-kezelést, aszinkron feladatokat, warp-specializációt, akadályokat, aszinkron TMA- és MMA-műveleteket, valamint Cluster Launch Controlt használ. A CTA külön feladatokra specializált warpokat kap: egyesek a TMA-betöltéseket, mások a tenzormagos mátrixszorzásokat, a softmaxhoz szükséges számításokat vagy az eredmények kiírását végzik. A visszafelé számításban külön warp foglalkozik a dQ redukciójával.
A struktúra része az on-chip pufferek kézi kiosztása, a csővezeték mélységének megválasztása és az adatok explicit producer-consumer akadályokon keresztüli átadása. A kernel mindkét irányban perzisztens módon működik, vagyis egy CTA egy SM-en több csempét dolgoz fel.
Kevesebb kód, magasabb mért teljesítmény
A PyTorch szerint a TLX-alapú figyelmi kernel körülbelül 3,2 ezer sor tömör Triton-szintű kódból áll. Ez nagyjából háromszor kevesebb, mint a FlashAttention-4 állapot-of-the-art, körülbelül 10 ezer soros CuteDSL-kerneljei. A vállalat szerint ez a megközelítés a modellezési mérnökök számára is olvashatóbbá és könnyebben bővíthetővé teszi a kernelt, nem csak a kernel-specialisták számára.
A bfloat16 pontosságú, B200 GPU-kon végzett mérésekben a TLX-kernel a GEM szempontjából fontos jagged alakzatokon körülbelül 13 százalékkal múlta felül a FlashAttention-4 2026. májusi változatát az előre irányú számításban. A visszafelé számításnál az előny körülbelül 50 százalék volt. A PyTorch a mérésekhez többek között NVIDIA Nsight Compute-ot, ptxas naplókat és TritonBench-et használt.
Az optimalizálások között szerepel a jagged csempék SM-ek közötti terheléskiegyenlítése, a többfázisú dQ-staging, a tenzormemória korábbi felszabadítása és a ciklusok kettébontása. A forward kernel esetében a csempék KV-terhelés szerinti rendezése és cikcakkos kiosztása a beszámoló szerint körülbelül 20 százalékot nyert vissza.
A kód elérhető a Meta kutatási tárházában
A PyTorch bejegyzése szerint a TLX-es JFA-kód nyilvánosan elérhető a Meta ads_model_kernel_library GitHub-tárházának tlx_jfa könyvtárában. A fejlesztők célja olyan Blackwellre optimalizált figyelmi kernel létrehozása volt, amely a kézzel írt CuteDSL- vagy CUDA-megoldások teljesítményéhez közelít, miközben a Triton magasabb szintű modelljében marad.
Ez a GEM-hez kapcsolódó fejlesztéseknél lehet fontos, mivel a figyelmi kernelben új modellezési változatok, például csúszóablakos és blokkszórt megoldások jelenhetnek meg. A PyTorch szerint a TLX ezeket a módosításokat könnyebben olvasható és bővíthető kódbázisban teszi lehetővé.

