Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
50 changes: 26 additions & 24 deletions flash_attn/cute/flash_fwd_sm100.py
Original file line number Diff line number Diff line change
Expand Up @@ -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},
Expand Down Expand Up @@ -2027,27 +2043,20 @@ 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
else:
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
Comment thread
drisspg marked this conversation as resolved.
# 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,
Expand Down Expand Up @@ -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)
Expand Down
72 changes: 72 additions & 0 deletions tests/cute/test_flash_attn.py
Original file line number Diff line number Diff line change
Expand Up @@ -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():
Comment thread
drisspg marked this conversation as resolved.
"""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)
Expand Down