Change
79e7591
79e759114cee2dcdfc5e63df71647a2aed991838 · commit on GitHub
pytorch-blog-feed: changed (217739 bytes, HTTP 200)
raw/pytorch-blog-feed/response.xml modified
- Source
- pytorch-blog-feed
- Lines added
- +146
- Lines removed
- -390
- Stored bytes at this commit
- 217,739
- Timestamp
- origin
- Raw artifact at this commit
- raw/pytorch-blog-feed/response.xml
Recorded headers
| observed_at | 2026-10-02T05:37:33.738Z |
|---|---|
| origin_date | 2026-10-02T04:04:21.000Z |
| status | 200 |
| final URL | https://pytorch.org/blog/feed/ |
| etag | "9f592f61d5a6f99053cd4e8b3c24cdc4" |
| last-modified | Thu, 01 Oct 2026 22:32:14 GMT |
| date | Fri, 02 Oct 2026 05:37:33 GMT |
| age | 5592 |
| cache-control | public, max-age=60, s-maxage=43200, stale-while-revalidate=86400, stale-if-error=604800 |
| cf-cache-status | null |
| content-encoding | null |
| content-length | 217739 |
@
@@ -12,7 +12,7 @@ <atom:link href="https://pytorch.org/blog/feed/" rel="self" type="application/rss+xml" /> <link>https://pytorch.org</link> <description></description>-
<lastBuildDate>Wed, 30 Sep 2026 14:39:49 +0000</lastBuildDate>+
<lastBuildDate>Thu, 01 Oct 2026 22:32:14 +0000</lastBuildDate> <language>en-US</language> <sy:updatePeriod> hourly </sy:updatePeriod>@
@@ -28,6 +28,149 @@ <height>32</height></image> <item>+
<title>Optimizing Jagged Flash Attention with TLX: The Road Toward SOTA FA4 on Blackwell</title>+
<link>https://pytorch.org/blog/optimizing-jagged-flash-attention-with-tlx-the-road-toward-sota-fa4-on-blackwell/</link>+
+
<dc:creator><![CDATA[Han Xu, Jacky Zhou, Jackie (Jiaqi) Xu, Hongtao Yu, Peng Chen (Dev Infra), Darren Liu, Dev (Devashish) Shankar, Max Leung, Nick Riasanovsky, Hao Yan, Manman Ren, Yuanwei (Kevin) Fang]]></dc:creator>+
<pubDate>Thu, 01 Oct 2026 22:26:57 +0000</pubDate>+
<category><![CDATA[Blog]]></category>+
<guid isPermaLink="false">https://pytorch.org/?p=171511</guid>+
+
<description><![CDATA[TL;DR In this blog post, we present our work on Jagged Flash Attention (JFA) — the attention kernel behind Meta’s Generative Ads Model (GEM) — on NVIDIA Blackwell (B200), built...]]></description>+
<content:encoded><![CDATA[<h2><span style="font-weight: 400;">TL;DR</span></h2>+
<p><span style="font-weight: 400;">In this blog post, we present our work on Jagged Flash Attention (JFA) — the attention kernel behind Meta’s Generative Ads Model (GEM) — on NVIDIA Blackwell (B200), built with TLX (Triton Low-level Extensions), which add explicit, hardware-aware control on to…+
<p><span style="font-weight: 400;">Attention is the single slowest kernel in GEM, and reaching peak Blackwell performance has traditionally required hand-written CuteDSL or CUDA — slow to develop and hard to extend to new variants. We show that TLX closes this gap on both fronts. On development effi…+
<p><span style="font-weight: 400;">We organize the work in two parts — the structural changes the TLX rewrite makes possible, and the optimizations we layer on top — all benchmarked in bfloat16 on B200 against FA4, the current state of the art on Blackwell.</span></p>+
<p><span style="font-weight: 400;">Code available at: </span><a href="https://github.com/facebookresearch/ads_model_kernel_library/tree/main/tlx_jfa"><span style="font-weight: 400;">https://github.com/facebookresearch/ads_model_kernel_library/tree/main/tlx_jfa</span></a></p>+
<h2><span style="font-weight: 400;">Introduction</span></h2>+
<p><span style="font-weight: 400;">Meta’s ads models, including the Generative Ads Model (GEM) [2] and the Kunlun architecture [5], run attention over jagged (variable-length or ragged) user sequences. As described in our </span><a href="https://engineering.fb.com/2026/08/03/ml-applications/tr…+
<p><img fetchpriority="high" decoding="async" class="aligncenter wp-image-171512 size-full" src="https://pytorch.org/wp-content/uploads/2026/09/1.png" alt="" width="1376" height="768" srcset="https://pytorch.org/wp-content/uploads/2026/09/1.png 1376w, https://pytorch.org/wp-content/uploads/2026/09/1…+
<p><span style="font-weight: 400;">Hitting peak throughput on Blackwell means keeping the tensor cores continuously fed. Plain Triton leaves most of these decisions to the compiler; TLX [1] exposes them as first-class primitives (explicit SMEM/TMEM allocation, async_task warp specialization, barrier…+
<h2><span style="font-weight: 400;">1. The Challenge: Performant and Extensible Attention on Blackwell</span></h2>+
<p><span style="font-weight: 400;">Optimizing attention on Blackwell means solving two problems at once. The first is performance: attention is the single slowest kernel in GEM, and it reaches peak throughput only when the global-memory loads, softmax, and matmuls are tightly overlapped. The second …+
<p><b>Why a compiler-scheduled baseline falls short.</b><span style="font-weight: 400;"> Our starting point is a mature Triton-based JFA kernel — algorithmically correct, but it leaves all on-chip data movement and scheduling to the compiler: no control over shared-memory allocation or pipeline dept…+
<p><b>The production case, broadcast-Q, further shapes the design:</b><span style="font-weight: 400;"> a single dense Q is broadcast across every sequence in the batch, so its gradient </span><b>dQ must be summed across the entire batch</b><span style="font-weight: 400;"> — turning the dQ epilogue i…+
<h2><span style="font-weight: 400;">2. Optimizing Jagged Flash Attention with TLX</span></h2>+
<p><span style="font-weight: 400;">Our work falls into two categories. </span><b>Structural changes</b><span style="font-weight: 400;"> (§2.1) are the ground-up reorganization of the kernel that TLX makes possible — they change how work and memory are laid out across the CTA’s warps. </span><b…+
<h3><span style="font-weight: 400;">2.1 Structural changes</span></h3>+
<p><span style="font-weight: 400;">Warp specialization, explicit shared-/tensor-memory management, and barrier pipelining are standard ingredients of any high-performance Blackwell attention kernel, so we keep this brief — what matters is that TLX lets us express them in high-level Triton code rathe…+
<p><img decoding="async" class="aligncenter wp-image-171513 size-full" src="https://pytorch.org/wp-content/uploads/2026/09/2.png" alt="" width="1620" height="900" srcset="https://pytorch.org/wp-content/uploads/2026/09/2.png 1620w, https://pytorch.org/wp-content/uploads/2026/09/2-300x167.png 300w, ht…+
<p><span style="font-weight: 400;">The same explicit control extends to memory and scheduling: we hand-allocate every on-chip buffer and choose its pipeline depth (e.g. triple-buffered K/V so the load warp runs ahead), alias TMEM buffers with non-overlapping lifetimes (QK scores, P, and the softmax …+
<h3><span style="font-weight: 400;">2.2 Optimizations</span></h3>+
<p><span style="font-weight: 400;">With the warp-specialized, persistent structure in place, the remaining work is to eliminate the stalls and wasted work that keep the tensor cores from staying busy. The following optimizations are largely independent and each targets a specific bottleneck.</span><…+
<p><span style="font-weight: 400;">We found them with a consistent methodology: NVIDIA Nsight Compute (NCU) to read hardware counters such as tensor-core/TMEM pipeline utilization, SM-busy, and register spills (which show up as local-memory traffic); the ptxas dump logs to confirm spills; and Triton…+
<p><span style="font-weight: 400;">At a glance, the optimizations below are:</span></p>+
<ul>+
<li style="font-weight: 400;" aria-level="1"><b>Scheduling jagged tiles across SMs</b><span style="font-weight: 400;"> — software load balancing plus Cluster Launch Control keep every SM busy despite the order-of-magnitude length skew of jagged inputs (forward and backward).</span></li>+
<li style="font-weight: 400;" aria-level="1"><b>Multi-stage dQ staging</b><span style="font-weight: 400;"> — a double-buffered SMEM staging pipeline hides the heavily contended broadcast-Q dQ reduce-add, the #1 backward bottleneck.</span></li>+
<li style="font-weight: 400;" aria-level="1"><b>Early tensor-memory release</b><span style="font-weight: 400;"> — freeing the dQ tensor-memory buffer before its final stores lets the MMA warp start the next tile’s dQ matmul sooner.</span></li>+
<li style="font-weight: 400;" aria-level="1"><b>Loop peeling</b><span style="font-weight: 400;"> — splitting the KV loop into a branch-free bulk pass and a tiny masked tail removes per-iteration mask overhead and the register spills it causes.</span></li>+
<li style="font-weight: 400;" aria-level="1"><b>2-CTA collaborative MMA</b><span style="font-weight: 400;"> — two CTAs cooperate on a single matmul to raise tensor-core utilization in the matmul-heavy backward (adopted from FA4).</span></li>+
</ul>+
<h4>Scheduling jagged tiles across SMs</h4>+
<p><span style="font-weight: 400;">Jagged inputs create severe load imbalance across SMs: with sequence lengths varying by orders of magnitude, a naive tile-to-SM mapping leaves some SMs idle while others grind through the long sequences. An SM-occupancy heatmap (Triton-MPP) makes this visual — a fe…+
<p><span style="font-weight: 400;">In the forward pass the outer loop is over the dense Q and the inner loop over the jagged K/V, so a tile’s cost is proportional to its batch’s key length — the tiles are highly non-uniform. Balancing this is our own idea for the jagged case: on the host…+
<p><img decoding="async" class="aligncenter wp-image-171516 size-full" src="https://pytorch.org/wp-content/uploads/2026/09/3.png" alt="" width="1396" height="1038" srcset="https://pytorch.org/wp-content/uploads/2026/09/3.png 1396w, https://pytorch.org/wp-content/uploads/2026/09/3-300x223.png 300w, h…+
<p><span style="font-weight: 400;">The backward needs a different strategy. Its dK/dV pass loops over K/V blocks on the outside and the dense Q on the inside, so now the K/V blocks are the tiles: each one is roughly uniform in cost (a K/V block against the shared dense Q), but the number of tiles pe…+
<p><span style="font-weight: 400;">Static balancing handles the skew we can predict ahead of time; the residual variance that only shows up at runtime is handled dynamically by Cluster Launch Control (CLC) [3]. A Blackwell feature, CLC hands out the next tile index on demand, so whichever SM finishe…+
<p><img decoding="async" class="aligncenter wp-image-171519 size-full" src="https://pytorch.org/wp-content/uploads/2026/09/4.png" alt="" width="1400" height="700" srcset="https://pytorch.org/wp-content/uploads/2026/09/4.png 1400w, https://pytorch.org/wp-content/uploads/2026/09/4-300x150.png 300w, ht…+
<p style="text-align: center;"><i><span style="font-weight: 400;">CTAs workload heapmat without CLC</span></i></p>+
<p><img decoding="async" class="aligncenter wp-image-171522 size-full" src="https://pytorch.org/wp-content/uploads/2026/09/5.png" alt="" width="1400" height="700" srcset="https://pytorch.org/wp-content/uploads/2026/09/5.png 1400w, https://pytorch.org/wp-content/uploads/2026/09/5-300x150.png 300w, ht…+
<p style="text-align: center;"><i><span style="font-weight: 400;">CTAs workload heatmap with CLC</span></i></p>+
<p><span style="font-weight: 400;">We apply CLC to both the forward and backward kernels; in TLX it is a lightweight producer/consumer protocol over a shared scheduling context.</span></p>+
<p><span style="font-weight: 400;">Although software load balancing and CLC may look redundant — both spread work across SMs — they are complementary: CLC schedules dynamically yet cannot tell which tiles are empty, so on jagged (and especially sparse) inputs it still spends cycles handing out the s…+
<h4>Multi-stage dQ staging</h4>+
<p><span style="font-weight: 400;">Triton-MPP barrier analysis of the backward pinned the dQ epilogue as the #1 bottleneck: its reduce-add to HBM accounted for ~9–11% of lost tensor-core utilization, the single largest actionable gap. Two things make it slow. First, because dQ is summed across the w…+
<p><span style="font-weight: 400;">We make the staging explicit and double-buffered so that one slice’s reduce-add to HBM overlaps the next slice’s copy out of TMEM, keeping a store in flight at all times:</span></p>+
<pre><code class="language-python">NCOL = BLOCK_D // (EPILOGUE_SUBTILE * 2) # narrow column slices, sized to fit SMEM
+
STAGES = 2 # double-buffered staging
+
dq_smem = tlx.local_alloc((BLOCK_M, NCOL), dq.dtype, STAGES)
+
+
for s in range(BLOCK_D // NCOL):
+
# TMEM -> registers
+
dq = tlx.local_load(dq_tmem[:, s * NCOL : (s + 1) * NCOL]) * LN2
+
+
# registers -> SMEM (ping-pong)
+
tlx.local_store(dq_smem[s % STAGES], dq.to(dq.dtype))
+
tlx.fence_async_shared()
+
+
tlx.async_descriptor_store(desc_dq, dq_smem[s % STAGES],
+
offsets, store_reduce="add") # SMEM -> HBM, accumulating
+
+
# keep at most 1 store in flight
+
tlx.async_descriptor_store_wait(STAGES - 1)</code></pre>+
<p>The technique generalizes to any shared-memory-bound reduce-store epilogue, not just attention.</p>+
<h4>Early tensor-memory release</h4>+
<p>The same analysis showed why that epilogue stalls the math from the TMEM side: the reduction warp holds the dQ tensor-memory buffer until all of its slices have been drained to HBM, so the MMA warp blocks on the dq-empty barrier before it can start the next tile’s dQ matmul. A waterfall tes…+
<p>The improvement is to pre-load the last one or two slices into registers and release the TMEM buffer before issuing their stores, so the MMA warp can begin reusing that memory while the reduction warp finishes writing to HBM:</p>+
<pre><code class="language-python"><span style="color: #993366;"># Drain most slices the normal way
+
# (TMEM -> SMEM -> async reduce-add to HBM).
+
<span style="color: #0000ff;">for</span></span> s <span style="color: #0000ff;">in range</span>(N_SLICES - EARLY_RELEASE_SUBTILES):
+
reduce_add_slice(s)
+
<span style="color: #993366;">
+
# Pre-load the final 1-2 slices into registers FIRST, then release
+
# the dQ TMEM buffer immediately
+
# so the MMA warp can start the next tile's dQ matmul.</span>
+
dq_tail = tlx.local_load(dq_tmem[:, last_slice]) <span style="color: #993366;"># TMEM -> registers</span>
+
tlx.barrier_arrive(dq_empties[buf]) <span style="color: #993366;"># <-- TMEM freed early</span>
+
reduce_add_from_registers(dq_tail)</code></pre>+
<p>How many slices to release early is autotuned (1 or 2); the optimal choice is small because releasing more raises register pressure enough to hurt. Together with multi-stage staging, this lets the MMA warp begin the next dQ matmul while the previous tile’s dQ is still being written out.</p>+
<h4>Loop peeling</h4>+
<p><span style="font-weight: 400;">NCU profiling showed the forward was MMA-issue starved — not memory- or compute-bound — with the tensor-core (TMEM) pipeline at ~57% utilization versus FA4’s ~82% on the same shape, and the lost cycles coming from softmax-warp overhead bubbling into the MMA-i…+
<p><span style="font-weight: 400;">Two signals point at the culprit. First, the ptxas dumps and NCU show register spills (local-memory traffic) in the register-heavy softmax warp — in an earlier round, simply granting the warps more registers cut that traffic by ~40% for a ~6% forward gain, confirmi…+
<p><span style="font-weight: 400;">The reason this branch is costly is subtle and specific to how the compiler allocates registers: Triton/TLX assign registers to variables </span><i><span style="font-weight: 400;">statically</span></i><span style="font-weight: 400;">, so a mask branch living inside…+
<p><span style="font-weight: 400;">The fix is to peel the loop into a branch-free bulk pass plus a tiny masked tail. Because the mask flag is a compile-time constant, the bulk body contains no mask variables at all, so the compiler frees those registers and can schedule a tighter MMA-issue cadence:<…+
<pre><code class="language-python"><span style="color: #993366;"># APPLY_MASK is constexpr: with APPLY_MASK=False the compiler statically removes the
+
# comparison, the offs_n / mask tensors, and the select -> a branch-free body that
+
# does not reserve registers for the (rare) masked path.</span>
+
aligned = (klen // BLOCK_N) * BLOCK_N
+
+
<span style="color: #0000ff;">for</span> start_n <span style="color: #0000ff;">in</span> tl.<span style="color: #0000ff;">range</span>(lo, aligned, BLOCK_N): <span style="color: #993366;"># bulk: straight-line, no mask code</span>
+
softmax_iter(start_n, ..., APPLY_MASK=False)
+
+
<span style="color: #0000ff;">for</span> start_n <span style="color: #0000ff;">in</span> tl.<span style="color: #0000ff;">range</span>(aligned, klen, BLOCK_N): <span style="color: #993366;"># tail: 0 or 1 iteration</span>
+
softmax_iter(start_n, ..., APPLY_MASK=True) <span style="color: #993366;"># mask only the partial last tile</span></code></pre>+
<p>We apply the same peeling to the register-heavy paths in both the forward softmax loop and the backward, where the gain comes not from skipping a cheap runtime branch but from the reduced register pressure it unlocks; in the backward it specifically reclaims the ~9% latency that the correctness m…+
<h4>2-CTA collaborative MMA (backward)</h4>+
<p><span style="font-weight: 400;">The backward is matmul-heavy — computing dQ, dK, and dV takes five GEMMs per K/V block — and a single CTA does not fully utilize the Blackwell tensor cores on its own. We adopt FA4’s 2-CTA (paired-CTA tcgen05) scheme: two CTAs in a cluster cooperate on one wi…+
<p><span style="font-weight: 400;">What is ours here is the TLX implementation: porting the scheme to our jagged, broadcast-Q layout and running it under our persistent and CLC multi-tile schedulers.</span></p>+
<p><img decoding="async" class="aligncenter wp-image-171525 size-full" src="https://pytorch.org/wp-content/uploads/2026/09/6.png" alt="" width="896" height="1200" srcset="https://pytorch.org/wp-content/uploads/2026/09/6.png 896w, https://pytorch.org/wp-content/uploads/2026/09/6-224x300.png 224w, htt…+
<p><span style="font-weight: 400;">Scoped to the production broadcast-Q, HEAD_DIM=128 case, it adds about +12% throughput (−11% latency) in the backward over the single-CTA path.</span></p>+
<h3><span style="font-weight: 400;">Performance</span></h3>+
<p><span style="font-weight: 400;">We benchmark on B200 (bf16) in two regimes: the production jagged case (Hierarchical Seed Pooling (HSP) — a broadcast dense Q against jagged K/V) and an LLM-style dense case (equal-length Q/K/V), both against FA4, the state-of-the-art open-source CuteDSL FlashAtten…+
<p><b>Jagged (broadcast-Q) — forward TFLOPS (bf16, B200)</b></p>+
<p><img decoding="async" class="aligncenter wp-image-171530 size-full" src="https://pytorch.org/wp-content/uploads/2026/09/7.png" alt="" width="1692" height="730" srcset="https://pytorch.org/wp-content/uploads/2026/09/7.png 1692w, https://pytorch.org/wp-content/uploads/2026/09/7-300x129.png 300w, ht…+
<p><b>Jagged (broadcast-Q) — backward TFLOPS (bf16, B200)</b></p>+
<p><img decoding="async" class="aligncenter wp-image-171531 size-full" src="https://pytorch.org/wp-content/uploads/2026/09/8.png" alt="" width="1926" height="792" srcset="https://pytorch.org/wp-content/uploads/2026/09/8.png 1926w, https://pytorch.org/wp-content/uploads/2026/09/8-300x123.png 300w, ht…+
<p><b>LLM dense — forward TFLOPS (bf16, B200, B=768, H=4, head_dim=128)</b></p>+
<p><img decoding="async" class="aligncenter wp-image-171536 size-full" src="https://pytorch.org/wp-content/uploads/2026/09/9.png" alt="" width="1358" height="680" srcset="https://pytorch.org/wp-content/uploads/2026/09/9.png 1358w, https://pytorch.org/wp-content/uploads/2026/09/9-300x150.png 300w, ht…+
<p><b>LLM dense — backward TFLOPS (bf16, B200, B=768, H=4, head_dim=128)</b></p>+
<p><img decoding="async" class="aligncenter wp-image-171537 size-full" src="https://pytorch.org/wp-content/uploads/2026/09/10.png" alt="" width="1358" height="680" srcset="https://pytorch.org/wp-content/uploads/2026/09/10.png 1358w, https://pytorch.org/wp-content/uploads/2026/09/10-300x150.png 300w,…+
<h2><span style="font-weight: 400;">3. Attention Variants: the flexibility of TLX</span></h2>+
<p><span style="font-weight: 400;">A practical benefit of building on TLX is that the kernel is structured enough to </span><i><span style="font-weight: 400;">fork</span></i><span style="font-weight: 400;"> for new requirements without a ground-up rewrite. Because the warp specialization, memory all…+
<h4>Low-precision (MXFP8) attention</h4>+
<p><span style="font-weight: 400;">For FP8-tolerant training we built a microscaling-FP8 variant of both the forward and backward by swapping the BF16 matmuls for TLX’s block-scaled MMA (FP8 E4M3 data with E8M0 per-block scale factors), reusing the same warp-specialized skeleton, barriers, and…+
<p><span style="font-weight: 400;">After tuning, the MXFP8 forward lands above FA4’s own FP8 kernel and the backward reaches parity with FA4 at dense — both comfortably beating the BF16 baseline — a large throughput win obtained largely by changing the GEMM calls rather than the kernel structu…+
<h4>Block-sparse attention</h4>+
<p><span style="font-weight: 400;">For long jagged histories where full O(N²) attention dominates, we built a two-stage sparse variant [4]. A cheap scoring kernel average-pools each Q and K block and selects the top-k most relevant KV blocks per Q block (a tunable selection ratio); a TLX attention k…+
<p><span style="font-weight: 400;">It reuses the same warp-specialized structure, CLC dispatch, and software load balancing (here used to skip unselected and empty tiles), stays PT2-friendly, and still supports broadcast-Q, GQA, and windowing. At a 0.5 selection ratio the forward is roughly 1.3–1.5×…+
<p><span style="font-weight: 400;">In both cases the bulk of the kernel carried over unchanged; only the math — the GEMM precision, or the set of KV blocks each tile visits — differed. That reuse is the practical payoff of TLX: the same kernel can be re-targeted to new precisions and sparsity patter…+
<h2><span style="font-weight: 400;">4. Summary and Future Work</span></h2>+
<p><span style="font-weight: 400;">The structural changes — warp specialization, persistent execution, explicit SMEM/TMEM allocation with reuse aliasing, and barrier-mediated pipelining — set the stage, and the optimizations layered on top — software load balancing, Cluster Launch Control, multi-sta…+
<p><span style="font-weight: 400;">Beyond raw performance, the biggest win is development efficiency. Much of the low-level plumbing that CuteDSL writes by hand — async pipelines, barrier management, tcgen05 MMA setup — is instead handled by the TLX compiler, so we express intent rather than machine…+
<h2><span style="font-weight: 400;">References</span></h2>+
<p><span style="font-weight: 400;">[1] TLX: Enabling Cluster Launch Control with Triton — </span><a href="https://pytorch.org/blog/enabling-cluster-launch-control-with-tlx/"><span style="font-weight: 400;">https://pytorch.org/blog/enabling-cluster-launch-control-with-tlx/</span></a></p>+
<p><span style="font-weight: 400;">[2] GEM: Meta’s Generative Ads Model — </span><a href="https://engineering.fb.com/2025/11/10/ml-applications/metas-generative-ads-model-gem-the-central-brain-accelerating-ads-recommendation-ai-innovation/"><span style="font-weight: 400;">https://engineering.f…+
<p><span style="font-weight: 400;">[3] NVIDIA Blackwell Tuning Guide (thread-block clusters) — </span><a href="https://docs.nvidia.com/cuda/blackwell-tuning-guide/index.html#thread-block-clusters"><span style="font-weight: 400;">https://docs.nvidia.com/cuda/blackwell-tuning-guide/index.html#thread-b…+
<p><span style="font-weight: 400;">[4] TLX Block Attention: A Warp-Specialized Blackwell Kernel for Fixed-Block Sparse Self-Attention — </span><a href="https://pytorch.org/blog/tlx-block-attention-a-warp-specialized-blackwell-kernel-for-fixed-block-sparse-self-attention/"><span style="font-weight: 4…+
<p><span style="font-weight: 400;">[5] Kunlun: Establishing Scaling Laws for Massive-Scale Recommendation Systems through Unified Architecture Design — </span><a href="https://arxiv.org/abs/2602.10016"><span style="font-weight: 400;">https://arxiv.org/abs/2602.10016</span></a></p>+
<p><span style="font-weight: 400;">[6] FlashAttention-4: Algorithm and Kernel Pipelining Co-Design for Asymmetric Hardware Scaling — </span><a href="https://arxiv.org/abs/2603.05451"><span style="font-weight: 400;">https://arxiv.org/abs/2603.05451</span></a></p>+
]]></content:encoded>+
+
+
+
</item>+
<item> <title>A Ray-Focused Guide to PyTorch Conference North America</title> <link>https://pytorch.org/blog/a-ray-focused-guide-to-pytorch-conference-north-america/</link> @
@@ -109,7 +252,7 @@<p>The interesting engineering starts one step later, in the decisions the relay leaves to each backend: what are the test candidates, which dispatches deserve a build, which of PyTorch’s tens of thousands of tests are meaningful on your hardware, how to adapt tests written for CUDA without fo…<h2>Three Evolving Candidates under Test: OOT Accelerator PyTorch Backend, PyTorch Core, Test Suite</h2><p><!-- IMAGE PLACEHOLDER: Figure 1 - Testing Surface in PyTorch and Torch Spyre (re-upload via WP media library) --></p>-
<p><em><img fetchpriority="high" decoding="async" class="alignnone wp-image-171228 size-full" src="https://pytorch.org/wp-content/uploads/2026/09/image-1-1.png" alt="" width="1886" height="828" srcset="https://pytorch.org/wp-content/uploads/2026/09/image-1-1.png 1886w, https://pytorch.org/wp-content…+
<p><em><img decoding="async" class="alignnone wp-image-171228 size-full" src="https://pytorch.org/wp-content/uploads/2026/09/image-1-1.png" alt="" width="1886" height="828" srcset="https://pytorch.org/wp-content/uploads/2026/09/image-1-1.png 1886w, https://pytorch.org/wp-content/uploads/2026/09/imag…<p><em>Figure 1: Testing Surface in PyTorch and Torch Spyre</em></p><h3>Challenge 1: Deciding what to test: three moving targets, and which combination to run</h3><p>Often, OOT accelerators have to deal with three evolving candidates to test after receiving a dispatch from <a href="https://pytorch.org/blog/introducing-cross-repository-ci-relay-scalable-ci-for-pytorchs-out-of-tree-backends/">PyTorch CRCR</a>. First, an OOT accelerator PyTorch backend codebase,…@
@@ -1122,395 +1265,8 @@ Tue 20 Oct, 16:20–16:30, LL20AB</p> -
</item>-
<item>-
<title>Low Precision Flash Attention 4: End-to-End Block-Scaled Attention for Blackwell</title>-
<link>https://pytorch.org/blog/low-precision-flash-attention-4-end-to-end-block-scaled-attention-for-blackwell/</link>-
-
<dc:creator><![CDATA[Dev (Devashish) Shankar, Darren Liu, Chunzhi Yang, Jackie (Jiaqi) Xu,Markus Hoehnerbach, Jason Xie, Santosh Mohan, Han Xu, Rich Zhu, Josh Fromm, Hongtao Yu, Max Leung, John Bocharov, Alan Purvis]]></dc:creator>-
<pubDate>Wed, 16 Sep 2026 18:55:21 +0000</pubDate>-
<category><![CDATA[Blog]]></category>-
<guid isPermaLink="false">https://pytorch.org/?p=159902</guid>-
-
<description><![CDATA[TL;DR We extend FlashAttention-4 [1] with MXFP8 forward and backward, reaching 2.85 PF/s forward and 2 PF/s backward on LLM shapes. On our internal shapes, FA4 MX8 reaches 2.54 PF/s...]]></description>-
<content:encoded><![CDATA[<h2><span style="font-weight: 400;">TL;DR</span></h2>-
<p><span style="font-weight: 400;">We extend FlashAttention-4 [1] with MXFP8 forward and backward, reaching 2.85 PF/s forward and 2 PF/s backward on LLM shapes. On our internal shapes, FA4 MX8 reaches 2.54 PF/s forward and 1.58 PF/s backward, delivering up to 1.6× and 1.52× gains over BF16. We fuse …-
<h2><span style="font-weight: 400;">1. Introduction</span></h2>-
<p><span style="font-weight: 400;">Blackwell’s tensor cores introduce block-scaled MMA instructions (tcgen05.mma.block_scale) that operate natively on microscaling formats — MXFP8, MXFP6, MXFP4 & NVFP4 — delivering 2-4x the throughput of BF16 MMA [4,5]. However, exploiting this…-
<p><span style="font-weight: 400;">In this work, we extend the FA4 attention kernel with end-to-end MXFP8 support for both forward and backward passes, and integrate it into a cross-attention module for Ads training with fused producer and output epilogues. The key contributions are: (1) TMEM alloca…-
<h2><span style="font-weight: 400;">2. Implementation Details</span></h2>-
<h3><span style="font-weight: 400;">2.1 Attention Forward</span></h3>-
<p><span style="font-weight: 400;">Attention forward consists of the following primary operations:</span></p>-
<p><span style="font-weight: 400;">S = Q @ K.T</span></p>-
<p><span style="font-weight: 400;">P = Softmax(S)</span></p>-
<p><span style="font-weight: 400;">O = P @ V</span></p>-
<p><span style="font-weight: 400;">To enable blockscaled MMA, we follow existing CuTe DSL examples from Quack GEMM kernels and CUTLASS C++ examples [5,6]. We use TMA loads to fetch scale factors from GMEM to SMEM, and copy SFs from SMEM to TMEM before triggering the UMMA. The primary challenge here …-
<p><span style="font-weight: 400;">Currently, in the softmax warp, softmax computation happens in FP32, and then the results are converted to BF16 before the PV multiplication. We convert P to MXFP8 while also computing the scales. We deep dive into the PTX optimizations done to achieve this efficie…-
<p><span style="font-weight: 400;">One subtle thing to note is that for P.V blockscaled MMA to work, the scales need to be computed along the MMA K-dim. For Q and K, this is the embedding dimension (D) of attention, but for V, the scales and quantization need to be computed along the sequence dimens…-
<h4><span style="font-weight: 400;">2.1.1 TMEM Allocation and barrier synchronization</span></h4>-
<p><span style="font-weight: 400;">Blackwell architecture has a fixed TMEM size of 512 column, which is completely utilized for MMA operands and accumulators in existing Blackwell FA kernels. This makes it challenging to add block scaled MMA, since scales also need to be in TMEM.</span></p>-
<p><span style="font-weight: 400;">FA4 forward uses a ping-pong computation between two Q tiles. We load two Q tiles, Q0 & Q1 of size [128, 128], and loop over K/V tiles (N dimension). The order of GEMMs is:</span></p>-
<table style="border-collapse: collapse; width: 100%; max-width: 520px; font-family: Arial, Helvetica, sans-serif; font-size: 24px; line-height: 1.25; background: #ffffff;">-
<tbody>-
<tr>-
<td style="border: 2px solid #111111; padding: 14px 12px;"><strong>GEMM</strong></td>-
</tr>-
<tr>-
<td style="border: 2px solid #111111; padding: 14px 12px;"><strong>Prologue</strong></td>-
</tr>-
<tr>-
<td style="border: 2px solid #111111; padding: 14px 12px; color: #366fbc;"><strong>S0</strong> = Q0 @ K0</td>-
</tr>-
<tr>-
<td style="border: 2px solid #111111; padding: 14px 12px; color: #366fbc;"><strong>S1</strong> = Q1 @ K0</td>-
</tr>-
<tr>-
<td style="border: 2px solid #111111; padding: 14px 12px;"><strong>Mainloop</strong> (for n in 0 .. N-1)</td>-
</tr>-
<tr>-
<td style="border: 2px solid #111111; padding: 14px 12px; color: #6a4adf;"><strong>O0</strong> = P0_n * V_n</td>-
</tr>-
<tr>-
<td style="border: 2px solid #111111; padding: 14px 12px; color: #366fbc;"><strong>S0</strong> = Q0 * K_{n+1}</td>-
</tr>-
<tr>-
<td style="border: 2px solid #111111; padding: 14px 12px; color: #d7622b;"><strong>O1</strong> = P1_n * V_n</td>-
</tr>-
<tr>-
<td style="border: 2px solid #111111; padding: 14px 12px; color: #366fbc;"><strong>S1</strong> = Q1 * K_{n+1}</td>-
</tr>-
<tr>-
<td style="border: 2px solid #111111; padding: 14px 12px;"><strong>Epilogue</strong></td>-
</tr>-
<tr>-
<td style="border: 2px solid #111111; padding: 14px 12px; color: #6a4adf;"><strong>O0</strong> = P0_N * V_N</td>-
</tr>-
<tr>-
<td style="border: 2px solid #111111; padding: 14px 12px; color: #d7622b;"><strong>O1</strong> = P1_N * V_N</td>-
</tr>-
</tbody>-
</table>-
<p><span style="font-weight: 400;">This is the TMEM Alloc:</span><img decoding="async" class="aligncenter wp-image-159911 size-full" src="https://pytorch.org/wp-content/uploads/2026/08/1.png" alt="" width="1788" height="426" srcset="https://pytorch.org/wp-content/uploads/2026/08/1.png 1788w, https:/…-
<ul>-
<li style="font-weight: 400;" aria-level="1"><span style="font-weight: 400;">For the prologue S(i) GEMMs, we can use O(i) for S(i) SFs, as O(i) hasn’t started yet</span></li>-
<li style="font-weight: 400;" aria-level="1"><span style="font-weight: 400;">SFs for O(i) can live overlapped with S(i) – this is the same strategy used by regular FA which overlaps P(i) with S(i). Thus, there already exists a barrier which ensures that S(i) TMEM is consumed, before we copy O(…-
<li style="font-weight: 400;" aria-level="1"><span style="font-weight: 400;">SFs for S(i) live overlap with S(1-i). This requires an additional barrier between the MMA and Softmax warps, since MMA is executed asynchronously, it is possible that there is a write-write conflict between S(1-i) accumula…-
</ul>-
<p><span style="font-weight: 400;">Concretely, the GEMM execution order with SF placement is shown below:</span></p>-
<table style="border-collapse: collapse; width: 100%; max-width: 1040px; font-family: Arial, Helvetica, sans-serif; font-size: 24px; line-height: 1.25; background: #ffffff;">-
<colgroup>-
<col style="width: 50%;" />-
<col style="width: 50%;" /> </colgroup>-
<tbody>-
<tr>-
<td style="border: 2px solid #111111; padding: 14px 12px;"><strong>GEMM</strong></td>-
<td style="border: 2px solid #111111; padding: 14px 12px;"><strong>SF TMEM Region</strong></td>-
</tr>-
<tr>-
<td style="border: 2px solid #111111; padding: 14px 12px;" colspan="2"><strong>Prologue</strong></td>-
</tr>-
<tr>-
<td style="border: 2px solid #111111; padding: 14px 12px; color: #366fbc;"><strong>S0</strong> = Q0 @ K0</td>-
<td style="border: 2px solid #111111; padding: 14px 12px; color: #6a4adf;">O0 (free, not started)</td>-
</tr>-
<tr>-
<td style="border: 2px solid #111111; padding: 14px 12px; color: #366fbc;"><strong>S1</strong> = Q1 @ K0</td>-
<td style="border: 2px solid #111111; padding: 14px 12px; color: #d7622b;">O1 (free, not started)</td>-
</tr>-
<tr>-
<td style="border: 2px solid #111111; height: 62px; padding: 0;" colspan="2"></td>-
</tr>-
<tr>-
<td style="border: 2px solid #111111; padding: 14px 12px;" colspan="2"><strong>Mainloop</strong> (for n in 0 .. N-1)</td>-
</tr>-
<tr>-
<td style="border: 2px solid #111111; padding: 14px 12px; color: #6a4adf;"><strong>O0</strong> = P0_n * V_n</td>-
<td style="border: 2px solid #111111; padding: 14px 12px; color: #366fbc;">S0 (S consumed, P distinct)</td>-
</tr>-
<tr>-
<td style="border: 2px solid #111111; padding: 14px 12px; color: #366fbc;"><strong>S0</strong> = Q0 * K_{n+1}</td>-
<td style="border: 2px solid #111111; padding: 14px 12px;"><span style="color: #366fbc;">S1</span><br />-
<strong style="color: #c92a2a;"> (new barrier!)</strong></td>-
</tr>-
<tr>-
<td style="border: 2px solid #111111; padding: 14px 12px; color: #d7622b;"><strong>O1</strong> = P1_n * V_n</td>-
<td style="border: 2px solid #111111; padding: 14px 12px; color: #6a4adf;">O0 (existing barrier)</td>-
</tr>-
<tr>-
<td style="border: 2px solid #111111; padding: 14px 12px; color: #366fbc;"><strong>S1</strong> = Q1 * K_{n+1}</td>-
<td style="border: 2px solid #111111; padding: 14px 12px;"><span style="color: #366fbc;">S0</span><br />-
<strong style="color: #c92a2a;"> (new barrier!)</strong></td>-
</tr>-
<tr>-
<td style="border: 2px solid #111111; height: 62px; padding: 0;" colspan="2"></td>-
</tr>-
<tr>-
<td style="border: 2px solid #111111; padding: 14px 12px;" colspan="2"><strong>Epilogue</strong></td>-
</tr>-
<tr>-
<td style="border: 2px solid #111111; padding: 14px 12px; color: #6a4adf;"><strong>O0</strong> = P0_N * V_N</td>-
<td style="border: 2px solid #111111; padding: 14px 12px;"></td>-
</tr>-
<tr>-
<td style="border: 2px solid #111111; padding: 14px 12px; color: #d7622b;"><strong>O1</strong> = P1_N * V_N</td>-
<td style="border: 2px solid #111111; padding: 14px 12px;"></td>-
</tr>-
</tbody>-
</table>-
<h4> <span style="font-weight: 400;">2.1.2 Improved unroll-KV </span></h4>-
<p><span style="font-weight: 400;">As we move to lower precision, MMA throughput increases (2x for MXFP8, 4x for MXFP4 vs BF16), the softmax warp’s SFU-bound computation stays the same. This shifts the bottleneck: softmax bubbles that were hidden behind slow BF16 MMA now become exposed.</span>…-
<p><span style="font-weight: 400;">In the persistent kernel, the tile boundary is particularly problematic. Without unroll-KV, the last two GEMMs of a tile are both PV, followed by both QK of the next tile. Softmax for the next tile’s Q0 can’t start until QK0 completes — but QK0 can̵…-
<p><img decoding="async" class="aligncenter wp-image-159927 size-full" src="https://pytorch.org/wp-content/uploads/2026/08/2.png" alt="" width="1696" height="778" srcset="https://pytorch.org/wp-content/uploads/2026/08/2.png 1696w, https://pytorch.org/wp-content/uploads/2026/08/2-300x138.png 300w, ht…-
<p><span style="font-weight: 400;">Now softmax for Q0 triggers one GEMM earlier, hiding the tile boundary latency behind the PV pipeline.</span></p>-
<p><span style="font-weight: 400;">However, enabling this for BF16 caused a regression: the correction warp (which processes both stages sequentially) was delayed because PV1 was pushed later, which cascaded through the softmax_corr_empty barrier into the next softmax. We fixed this by moving the ba…-
<h4><span style="font-weight: 400;">2.1.3 Optimized Online MXFP8 conversion</span></h4>-
<p><span style="font-weight: 400;">Since we want to do both QK and PV GEMMs using block-scaled attention, this necessitates an online conversion of the post-softmax output (P) from FP32 to MXFP8, instead of BF16. Applying blockscaling is a 3-step process. For a block x consisting of 32 elements:</sp…-
<p><img decoding="async" class="aligncenter wp-image-159928 size-full" src="https://pytorch.org/wp-content/uploads/2026/08/3.png" alt="" width="606" height="150" srcset="https://pytorch.org/wp-content/uploads/2026/08/3.png 606w, https://pytorch.org/wp-content/uploads/2026/08/3-300x74.png 300w" sizes…-
<p><span style="font-weight: 400;">Here, a is the amax for the block, and sigma is the scaling factor. In order to do this optimally in Blackwell, we make use of 3-instruction max, and fmul2 instructions (exposed in CuteDSL nvvm) for step (1) and (3) – which happen for every element. For compu…-
<p><span style="font-weight: 400;">Additionally, we notice that computation of softmax already involves calculating row maxes for the 128 elements before the exponentiation. Since exp is a monotonically increasing function, we can re-use the max computed for softmax, thus preventing additional max o…-
<p><span style="font-weight: 400;">Our recent experiments indicate that performance can be further enhanced via constant scaling for P; this remains numerically robust given that the softmax operator naturally constrains outputs within the [0, 1] interval.</span></p>-
<h4><span style="font-weight: 400;">2.1.4 Handling variable sequence length tensors with TMA</span></h4>-
<p><span style="font-weight: 400;">Handling jagged data is important in Ads models, but it is especially challenging with MXFP8. Blackwell block-scaled MMA uses a </span><a href="https://docs.nvidia.com/cutlass/latest/media/docs/cpp/blackwell_functionality.html#scale-factor-layouts"><span style="fon…-
<p><img decoding="async" class="aligncenter wp-image-159929 size-full" src="https://pytorch.org/wp-content/uploads/2026/08/4.png" alt="" width="2048" height="1090" srcset="https://pytorch.org/wp-content/uploads/2026/08/4.png 2048w, https://pytorch.org/wp-content/uploads/2026/08/4-300x160.png 300w, h…-
<h3><span style="font-weight: 400;">2.2 Attention Backward</span></h3>-
<p><span style="font-weight: 400;">Attention backward consists of the following key operations:</span></p>-
<p><span style="font-weight: 400;">GEMMs:</span></p>-
<p><span style="font-weight: 400;"> S = K @ Q.T </span></p>-
<p><span style="font-weight: 400;"> dP = V @ dO.T </span></p>-
<p><span style="font-weight: 400;"> dV = P.T @ dO </span></p>-
<p><span style="font-weight: 400;"> dK = dS.T @ Q </span></p>-
<p><span style="font-weight: 400;"> dQ = dS @ K </span></p>-
<p><span style="font-weight: 400;">Compute:</span></p>-
<p><span style="font-weight: 400;">P = softmax(S)</span></p>-
<p><span style="font-weight: 400;">dS = dsoftmax(dP, P)</span></p>-
<p><span style="font-weight: 400;">Notice that in the backward pass, several GEMMs are transposed. For instance, consider:</span></p>-
<p><span style="font-weight: 400;"> dP = V @ dO.T </span></p>-
<p><span style="font-weight: 400;"> dV = P.T @ dO</span></p>-
<p><span style="font-weight: 400;">For the dP MMA, the dO quantization needs to be along the D (Embedding dimension), however, for the dV MMA, the quantization needs to be along the M (Sequence dimension). This is similar to what we saw with the P.V MMA in forward, where V needs to be quantized alon…-
<p><span style="font-weight: 400;">To solve this problem, we quantize Q, K, and dO using square [32,32] blocks, making the E4M3 payload transpose-invariant. Both GEMMs can therefore reuse one quantized representation instead of storing separate versions. The much smaller E8M0 scales are still laid o…-
<h4><span style="font-weight: 400;">2.2.1 TMEM Allocation</span></h4>-
<p><span style="font-weight: 400;">[Note: This scheme describes the 1-CTA path. We currently are using the 1-CTA path for MX8, as it is performing better currently]</span></p>-
<p><span style="font-weight: 400;">While the forward has 2 GEMMs (S and O) with 2-stage Q pipelining across 4 accumulators, the backward has 5 GEMMs that all need TMEM space for accumulators and scale factors. The TMEM layout packs all 512 columns:</span></p>-
<p><img decoding="async" class="aligncenter wp-image-159930 size-full" src="https://pytorch.org/wp-content/uploads/2026/08/5.png" alt="" width="1782" height="296" srcset="https://pytorch.org/wp-content/uploads/2026/08/5.png 1782w, https://pytorch.org/wp-content/uploads/2026/08/5-300x50.png 300w, htt…-
<p><span style="font-weight: 400;">The core constraint: dK and dV are persistent accumulators — they accumulate across the entire M-loop, so their TMEM regions are occupied for the kernel’s entire lifetime. This means SF placement can only use the S and dP regions. We use the dP region f…-
<table>-
<tbody>-
<tr>-
<td><b>Name</b></td>-
<td><b>GEMM SFs</b></td>-
<td><b>SF TMEM Region</b></td>-
<td><b>Why safe?</b></td>-
</tr>-
<tr>-
<td><span style="font-weight: 400;">Prologue</span></td>-
<td><span style="font-weight: 400;">SFK, SFQ, SFV, SFDO</span></td>-
<td><span style="font-weight: 400;" data-rich-links="{"dde_di":"kix.slzhx8cwgju5","dde-fdv":"dK","dde-sii":"dropdownItem.whf8o9a3vz9c","ddefe-ddi":{"cv":{"op":"set","opValue":[{"di-id&q…-
<td><span style="font-weight: 400;">dK free during prologue</span></td>-
</tr>-
<tr>-
<td><span style="font-weight: 400;">S (K@Q.T)</span></td>-
<td><span style="font-weight: 400;">SFK, SFQ</span></td>-
<td><span style="font-weight: 400;" data-rich-links="{"dde_di":"kix.slzhx8cwgju5","dde-fdv":"dP","dde-sii":"dropdownItem.9q1yfrhxqyq6","ddefe-ddi":{"cv":{"op":"set","opValue":[{"di-id&q…-
<td><span style="font-weight: 400;">pipeline_dP_drain.empty.wait ensures the previous dP value has been drained. In cluster-one this aliases the existing pipeline_dP barrier.</span></td>-
</tr>-
<tr>-
<td><span style="font-weight: 400;">dK (dS.T@Q)</span></td>-
<td><span style="font-weight: 400;">SFDS, SFQ_dK</span></td>-
<td><span style="font-weight: 400;" data-rich-links="{"dde_di":"kix.slzhx8cwgju5","dde-fdv":"dP","dde-sii":"dropdownItem.9q1yfrhxqyq6","ddefe-ddi":{"cv":{"op":"set","opValue":[{"di-id&q…-
<td><span style="font-weight: 400;">Since dS is the input to the dK GEMM, existing barriers ensure the region is free to use. Note that dS only takes 32 columns as it is FP8, so we have 96 columns free.</span></td>-
</tr>-
<tr>-
<td><span style="font-weight: 400;">dQ (dS@K)</span></td>-
<td><span style="font-weight: 400;">SFDS_dQ, SFK_dQ</span></td>-
<td><span style="font-weight: 400;" data-rich-links="{"dde_di":"kix.slzhx8cwgju5","dde-fdv":"S","dde-sii":"dropdownItem.egcjvhlb4vmx","ddefe-ddi":{"cv":{"op":"set","opValue":[{"di-id&qu…-
<td><span style="font-weight: 400;">pipeline_S_drain.empty.wait ensures S has been read before its TMEM region is reused.</span></td>-
</tr>-
<tr>-
<td><span style="font-weight: 400;">dP (V@dO.T)</span></td>-
<td><span style="font-weight: 400;">SFV, SFDO</span></td>-
<td><span style="font-weight: 400;" data-rich-links="{"dde_di":"kix.slzhx8cwgju5","dde-fdv":"S","dde-sii":"dropdownItem.egcjvhlb4vmx","ddefe-ddi":{"cv":{"op":"set","opValue":[{"di-id&qu…-
<td><span style="font-weight: 400;">Implicit ordering</span></td>-
</tr>-
<tr>-
<td><span style="font-weight: 400;">dV (P.T@dO)</span></td>-
<td><span style="font-weight: 400;">SFP, SFDO_dV</span></td>-
<td><span style="font-weight: 400;" data-rich-links="{"dde_di":"kix.slzhx8cwgju5","dde-fdv":"S","dde-sii":"dropdownItem.egcjvhlb4vmx","ddefe-ddi":{"cv":{"op":"set","opValue":[{"di-id&qu…-
<td><span style="font-weight: 400;">pipeline_S_P.empty.wait (already needed for P readiness)</span></td>-
</tr>-
</tbody>-
</table>-
<h4><span style="font-weight: 400;">2.2.2 Online square dS quantization</span></h4>-
<p><span style="font-weight: 400;">In backward, dS is computed in FP32 and consumed by two GEMMs:</span></p>-
<p><span style="font-weight: 400;">dK = dSᵀ @ Q, dQ = dS @ K.</span></p>-
<p><span style="font-weight: 400;">We quantize each 32×32 block of dS once. Each warp lane owns 32 values and computes a thread-local absolute maximum. Blackwell’s redux.sync.max.abs.f32 instruction then reduces these 32 partial maxima across the warp to obtain the AMAX for the full block. Because t…-
<p><span style="font-weight: 400;">Quantizing dS once avoids an additional E8M0 conversion, inverse-scale multiplication, E4M3 conversion, and payload fragment. The dK and dQ MMAs still copy the shared scale into their respective hardware layouts. We find the [32, 32] square quantization to be numer…-
<p><img decoding="async" class="aligncenter wp-image-159931 size-full" src="https://pytorch.org/wp-content/uploads/2026/08/6.png" alt="" width="1567" height="1034" srcset="https://pytorch.org/wp-content/uploads/2026/08/6.png 1567w, https://pytorch.org/wp-content/uploads/2026/08/6-300x198.png 300w, h…-
<h4><span style="font-weight: 400;">2.2.3 FP16 dQ Reduction</span></h4>-
<p><span style="font-weight: 400;">Through ablation studies, we identified dQ reduction as a critical throughput bottleneck. Because dQ necessitates writing a 128×128 tile to GMEM during every inner-loop iteration, it substantially increases global memory bandwidth consumption. To mitigate this…-
<h3><span style="font-weight: 400;">2.3 Hiding the Q/K/V Quantization overhead</span></h3>-
<p><span style="font-weight: 400;">While we obtain good speedups for the core attention kernel, one of the key challenges is quantization overhead, especially for smaller sequence lengths. To eliminate the quantization overhead, we fuse the MXFP8 quantization into the epilogue of the preceding kerne…-
<p><span style="font-weight: 400;">Another challenge for the backward pass is the transpose quantization, the fact that we need to quantize Q, dO and K tensors along both axes for the backward pass. Traditional block-scaling is not transpose invariant. Initially, we tried fusing dual quantization [3…Diff display stops at 400 lines. The line counts above are from the whole diff. 86 lines shown here cut at 300 characters. The raw artifact at this commit is linked above.