From c58907727bb3bf1213c4c3cba0814e3943ca6c17 Mon Sep 17 00:00:00 2001 From: drisspg Date: Sat, 25 Jul 2026 20:42:13 +0000 Subject: [PATCH] Numeric tweaks to fp8 stack-info: PR: https://github.com/Dao-AILab/flash-attention/pull/2731, branch: drisspg/stack/49 --- flash_attn/cute/flash_fwd_sm100.py | 50 +++++++++++---------- tests/cute/test_flash_attn.py | 72 ++++++++++++++++++++++++++++++ 2 files changed, 98 insertions(+), 24 deletions(-) diff --git a/flash_attn/cute/flash_fwd_sm100.py b/flash_attn/cute/flash_fwd_sm100.py index fe5b9008269..281b1d9615c 100644 --- a/flash_attn/cute/flash_fwd_sm100.py +++ b/flash_attn/cute/flash_fwd_sm100.py @@ -80,6 +80,22 @@ # num_regs_correction: int — register count for correction warps (multiple of 8) # num_regs_other is derived: 512 - num_regs_softmax * 2 - num_regs_correction # (hd256 exception: num_regs_other is fixed at 32, not derived) + +# Note [Low Precision Scaling] +# P is in (0, 1] and is cast to the input dtype before P @ V, so scaling it by 2^max_offset +# spends the dtype's unused upper code points on the probability tail. A positive +# rescale_threshold lets the row max lag by that many log2 units, so P can reach +# 2^(max_offset + rescale_threshold); above the dtype max the top probabilities saturate +# while the FP32 denominator still counts them in full, shrinking the output (#2716). + +# log2 of the largest finite value representable in each supported input dtype. +_LOG2_DTYPE_MAX = { + cutlass.Float8E4M3FN: math.log2(448.0), + cutlass.Float8E5M2: math.log2(57344.0), + cutlass.Float16: math.log2(65504.0), + cutlass.BFloat16: math.log2(3.3895313892515355e38), +} + _TUNING_CONFIG = { (True, False, 128, False): {"ex2_emu_freq": 10, "ex2_emu_start_frg": 1, "num_regs_softmax": 176, "num_regs_correction": 88}, (False, True, 128, False): {"ex2_emu_freq": 16, "ex2_emu_start_frg": 1, "num_regs_softmax": 192, "num_regs_correction": 72}, @@ -2027,16 +2043,8 @@ def softmax_loop( qk_descale, _ = self._load_effective_descales(descale_tensors, batch_idx, kv_head_idx) - # P is scaled by 2^max_offset before the FP8 conversion. With rescale_threshold > 0 - # the row max can be stale by up to rescale_threshold (in log2 units), so P can reach - # 2^(max_offset + rescale_threshold). max_offset + rescale_threshold must stay within - # log2(fp8_max) (448 = 2^8.8 for e4m3fn, 57344 = 2^15.8 for e5m2), otherwise the - # largest probabilities saturate and accuracy degrades (#2716). - max_offset = ( - 4 if cutlass.const_expr(self.q_dtype is cutlass.Float8E4M3FN) else - 8 if cutlass.const_expr(self.q_dtype.width == 8) else - 0 - ) + # See Note [Low Precision Scaling] + max_offset = 8 if cutlass.const_expr(self.q_dtype.width == 8) else 0 if const_expr(self.score_mod is None): softmax_scale_log2_eff = softmax_scale_log2 * qk_descale softmax_scale_eff = None @@ -2044,10 +2052,11 @@ def softmax_loop( softmax_scale_log2_eff = softmax_scale_log2 softmax_scale_eff = softmax_scale * qk_descale - rescale_threshold = ( - 8.0 if const_expr(self.q_dtype.width == 16) else - 4.0 if const_expr(self.q_dtype.width == 8) else - 0.0 + rescale_threshold = 8.0 if const_expr(self.q_dtype.width == 16) else 0.0 + # See Note [Low Precision Scaling] + assert max_offset + rescale_threshold < _LOG2_DTYPE_MAX[self.q_dtype], ( + f"max_offset ({max_offset}) + rescale_threshold ({rescale_threshold}) must stay " + f"below log2(max {self.q_dtype} value) to avoid saturating P" ) softmax = SoftmaxSm100.create( softmax_scale_log2_eff, @@ -2468,17 +2477,10 @@ def correction_loop( else: softmax_scale_log2_eff = softmax_scale_log2 - # Must match the softmax warp's max_offset (see comment there; #2716); - # max_offset_scale = 2^max_offset. - max_offset = ( - Float32(4.0) if cutlass.const_expr(self.q_dtype is cutlass.Float8E4M3FN) else - Float32(8.0) if cutlass.const_expr(self.q_dtype.width == 8) else - Float32(0.0) - ) + # Must match the softmax warp's max_offset; max_offset_scale = 2^max_offset. + max_offset = Float32(8.0) if cutlass.const_expr(self.q_dtype.width == 8) else Float32(0.0) max_offset_scale = ( - Float32(16.0) if cutlass.const_expr(self.q_dtype is cutlass.Float8E4M3FN) else - Float32(256.0) if cutlass.const_expr(self.q_dtype.width == 8) else - Float32(1.0) + Float32(256.0) if cutlass.const_expr(self.q_dtype.width == 8) else Float32(1.0) ) seqlen = SeqlenInfoCls(batch_idx) n_block_min, n_block_max = block_info.get_n_block_min_max(seqlen, m_block, split_idx, num_splits) diff --git a/tests/cute/test_flash_attn.py b/tests/cute/test_flash_attn.py index ce48dd11909..cd615673b39 100644 --- a/tests/cute/test_flash_attn.py +++ b/tests/cute/test_flash_attn.py @@ -1835,6 +1835,78 @@ def _generate_block_kvcache( return k_cache, v_cache, page_table, k_cache_paged, v_cache_paged, num_blocks +def _run_fp8_paged_decode(q, k, v, page_size=128): + """Run a single-sequence FP8 paged decode with unit descales.""" + seqlen_k, nheads_kv, d = k.shape + num_pages = math.ceil(seqlen_k / page_size) + k_cache = torch.zeros(num_pages, page_size, nheads_kv, d, device=k.device, dtype=k.dtype) + v_cache = torch.zeros_like(k_cache) + k_cache.view(-1, nheads_kv, d)[:seqlen_k].copy_(k) + v_cache.view(-1, nheads_kv, d)[:seqlen_k].copy_(v) + page_table = torch.arange(num_pages, dtype=torch.int32, device=k.device).unsqueeze(0) + descale = torch.ones(1, nheads_kv, dtype=torch.float32, device=k.device) + return _flash_attn_fwd( + q, + k_cache, + v_cache, + cu_seqlens_q=torch.tensor([0, 1], dtype=torch.int32, device=q.device), + seqused_k=torch.tensor([seqlen_k], dtype=torch.int32, device=q.device), + page_table=page_table, + softmax_scale=d**-0.5, + causal=True, + q_descale=descale, + k_descale=descale, + v_descale=descale, + )[0] + + +def _fp8_decode_reference(q, k, v): + """Compute FP32 attention over dequantized FP8 decode inputs.""" + nheads = q.shape[1] + k = k.float().repeat_interleave(nheads // k.shape[1], dim=1) + v = v.float().repeat_interleave(nheads // v.shape[1], dim=1) + scores = torch.einsum("qhd,khd->hqk", q.float(), k) * q.shape[-1] ** -0.5 + return torch.einsum("hqk,khd->qhd", torch.softmax(scores, dim=-1), v) + + +@pytest.mark.skipif(not IS_SM100, reason="FP8 paged decode is SM100-only") +@maybe_fake_tensor_mode(USE_FAKE_TENSOR) +def test_flash_attn_fp8_paged_decode_tile_boundary(): + """A second KV tile must not saturate e4m3 softmax probabilities.""" + torch.manual_seed(0) + q = torch.randn(1, 6, 128, device="cuda", dtype=torch.bfloat16).to(torch.float8_e4m3fn) + k = torch.randn(129, 1, 128, device="cuda", dtype=torch.bfloat16).to(torch.float8_e4m3fn) + v = torch.randn(129, 1, 128, device="cuda", dtype=torch.bfloat16).to(torch.float8_e4m3fn) + + out = _run_fp8_paged_decode(q, k, v) + if is_fake_mode(): + return + + ref = _fp8_decode_reference(q, k, v) + cosine = torch.nn.functional.cosine_similarity(out.float().flatten(), ref.flatten(), dim=0) + assert cosine > 0.99, f"FP8 paged decode lost accuracy at the tile boundary: {cosine=}" + + +@pytest.mark.skipif(not IS_SM100, reason="FP8 paged decode is SM100-only") +@maybe_fake_tensor_mode(USE_FAKE_TENSOR) +def test_flash_attn_fp8_paged_decode_preserves_tail_mass(): + """Collectively significant e4m3 softmax tails must not flush to zero.""" + q = torch.zeros(1, 6, 128, device="cuda", dtype=torch.float8_e4m3fn) + q[..., 0] = 16.0 + k = torch.zeros(1024, 1, 128, device="cuda", dtype=torch.float8_e4m3fn) + # Decode visits KV blocks right-to-left, so k[-1] establishes the max before the tails. + k[:-1, ..., 0] = -7.0 + v = torch.ones_like(k) + v[-1] = 0.0 + + out = _run_fp8_paged_decode(q, k, v) + if is_fake_mode(): + return + + ref = _fp8_decode_reference(q, k, v) + torch.testing.assert_close(out.float(), ref, atol=0.01, rtol=0.1) + + @pytest.mark.parametrize("page_size", [16, 64, 256]) @pytest.mark.parametrize("seqlen_q", [64, 128, 256]) @maybe_fake_tensor_mode(USE_FAKE_TENSOR)