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

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

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

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

Kapcsolódó hírek

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

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…

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…

Általánosan elérhetővé váltak a Modal többgépes AI-fürtjei
Chipek és infrastruktúra2026. október 1.

Általánosan elérhetővé váltak a Modal többgépes AI-fürtjei

A Modal általánosan elérhetővé tette a Modal Clusters szolgáltatást, amely több csomópontból álló AI-fürtök indítását és kezelését egyszerűsíti. A vállalat szerint a…