From e6d746e99ca1d67535433d9c94561d3883c1a64f Mon Sep 17 00:00:00 2001 From: Jin Tao Date: Tue, 6 Oct 2026 13:00:22 +0000 Subject: [PATCH 1/4] [Triton/Gluon] [gfx942] Read the fp8_scalar cache as e4m3fnuz on gfx942 The kernel decoded every fp8 byte as OCP e4m3, so gfx942, whose native fp8 (and vLLM's fp8 KV cache) is e4m3fnuz, took bf16 caches only. Add an FP8_FNUZ constexpr that picks the fp8 element type the dequant helpers read, and set it from the arch in the wrapper. gfx942 now takes the per-tensor fp8_scalar cache under bf16 dots; fp8 q, fp8 dots and fp8_dsv32_mla stay gfx950-only. A typed cache in the other arch's encoding is refused instead of misread; a uint8 view is taken to be the arch's own fp8. gfx942 (MI325X), GLM-5.3-Flash shape (16 heads, rope-free 512, top-k 2048), fp8 vs bf16 cache: decode 31.0 vs 33.2 us at 2 tokens, 33.1 vs 35.5 us at 16; 16K-token prefill 9.34 vs 10.54 ms. test_sparse_mla.py: 61 passed, 13 skipped (fp8 dots, fp8_dsv32_mla). Co-authored-by: Cursor --- .../gfx950/attention/sparse_mla.py | 27 ++++++++--- aiter/ops/triton/attention/sparse_mla.py | 47 ++++++++++++------- .../triton_tests/attention/test_sparse_mla.py | 47 +++++++++++++------ 3 files changed, 82 insertions(+), 39 deletions(-) diff --git a/aiter/ops/triton/_gluon_kernels/gfx950/attention/sparse_mla.py b/aiter/ops/triton/_gluon_kernels/gfx950/attention/sparse_mla.py index 269b72687be..dca8d483d76 100644 --- a/aiter/ops/triton/_gluon_kernels/gfx950/attention/sparse_mla.py +++ b/aiter/ops/triton/_gluon_kernels/gfx950/attention/sparse_mla.py @@ -153,14 +153,14 @@ def _cache_load( @gluon.jit -def _fp8_to_f32(x_u8): - return x_u8.to(gl.float8e4nv, bitcast=True).to(gl.float32) +def _fp8_to_f32(x_u8, FP8_TY: gl.constexpr): + return x_u8.to(FP8_TY, bitcast=True).to(gl.float32) @gluon.jit -def _fp8_to_bf16(x_u8): +def _fp8_to_bf16(x_u8, FP8_TY: gl.constexpr): # Exact: fp8's 3 mantissa bits fit bf16's 8. - return x_u8.to(gl.float8e4nv, bitcast=True).to(gl.bfloat16) + return x_u8.to(FP8_TY, bitcast=True).to(gl.bfloat16) @gluon.jit @@ -328,6 +328,11 @@ class Cfg: IDX_BUFFER_LOAD: gl.constexpr FP8_MFMA: gl.constexpr # "fp8_scalar" only: feed the matrix core the cache's # own fp8 instead of dequantizing to bf16 + # fp8 encoding of the cache, q and the dot operands: e4m3fnuz (gfx942's native, + # which its fp8 MFMA reads) or OCP e4m3 (gfx950's) + FP8_FNUZ: gl.constexpr + FP8_TY: gl.constexpr + FP8_MAX: gl.constexpr # Cache policy per load site GATHER_CACHE: gl.constexpr IDX_CACHE: gl.constexpr @@ -386,6 +391,7 @@ def __init__( DSV4_WALK=False, SCL_DWORD=False, STAGED_K32=False, + FP8_FNUZ=False, ): self.BLOCK_M = gl.constexpr(BLOCK_M) self.BLOCK_K = gl.constexpr(BLOCK_K) @@ -407,6 +413,9 @@ def __init__( self.HEAD_ALIGNED = gl.constexpr(HEAD_ALIGNED) self.IDX_BUFFER_LOAD = gl.constexpr(IDX_BUFFER_LOAD) self.FP8_MFMA = gl.constexpr(FP8_MFMA) + self.FP8_FNUZ = gl.constexpr(FP8_FNUZ) + self.FP8_TY = gl.constexpr(gl.float8e4b8 if FP8_FNUZ else gl.float8e4nv) + self.FP8_MAX = gl.constexpr(240.0 if FP8_FNUZ else 448.0) self.GATHER_CACHE = gl.constexpr(GATHER_CACHE) self.IDX_CACHE = gl.constexpr(IDX_CACHE) self.ASYNC_LDS = gl.constexpr(ASYNC_LDS) @@ -735,13 +744,13 @@ def _deq_store(x_u8, sc, kv_smem, off, cfg, fmt, AXIS: gl.constexpr): x_u8.to(gl.float8e4nv, bitcast=True), sc, gl.bfloat16 ) elif fmt.KIND == "fp8_scalar": - val = _fp8_to_bf16(x_u8) + val = _fp8_to_bf16(x_u8, cfg.FP8_TY) else: gl.static_assert( fmt.KIND == "fp8_g64" or fmt.KIND == "fp8_dsv32_mla", "fp8_dsv4_mla dequantizes with DEQ upcast or asm", ) - val = (_fp8_to_f32(x_u8) * sc).to(gl.bfloat16) + val = (_fp8_to_f32(x_u8, cfg.FP8_TY) * sc).to(gl.bfloat16) if AXIS == 1: kv_smem.slice(off, x_u8.shape[1], dim=1).store(val) else: @@ -1207,7 +1216,7 @@ def _stage(cfg, seg, x_u8, sc, k_rope, kv_smem, rope_smem): elif fmt.KIND == "fp8_dsv32_mla": rope_smem.store(k_rope) elif fmt.KIND == "fp8_scalar" and cfg.ROPE_SEPARATE: - rope_smem.store(_fp8_to_bf16(k_rope)) + rope_smem.store(_fp8_to_bf16(k_rope, cfg.FP8_TY)) # "fp8_g64": the whole head is one fp8 tile; nothing else to store. @@ -2239,6 +2248,9 @@ def _sparse_mla( IDX_BUFFER_LOAD: gl.constexpr, HAS_INVALID: gl.constexpr, FP8_MFMA: gl.constexpr = False, + # fp8 is e4m3fnuz (gfx942) rather than OCP e4m3: the cache, the q this kernel + # quantizes, and the fp8 dot operands. + FP8_FNUZ: gl.constexpr = False, # q already quantized to e4m3 by the caller, plus the scalar f32 scale it # was quantized with. This is the calling convention aiter's asm # mla_decode_fwd uses, where vLLM passes layer._q_scale. @@ -2397,6 +2409,7 @@ def _sparse_mla( DSV4_WALK, CS0_ALIGN >= 4, STAGED_K32, + FP8_FNUZ, ) main_fmt = Fmt( cfg, diff --git a/aiter/ops/triton/attention/sparse_mla.py b/aiter/ops/triton/attention/sparse_mla.py index 350bfae68d9..7c1d3b766f0 100644 --- a/aiter/ops/triton/attention/sparse_mla.py +++ b/aiter/ops/triton/attention/sparse_mla.py @@ -73,10 +73,14 @@ def _cache_pointers(fmt, kv, d_qk, kv_scale): SUPPORTED_ARCHS = ("gfx942", "gfx950") -# The kernel reads every fp8 byte, in q, the cache and the dot operands, as OCP -# e4m3, which is gfx950's native fp8. gfx942's is fnuz, and the kernel does not -# decode it yet, so gfx942 takes bf16 q and a bf16 cache only. +# The kernel reads fp8 in the arch's native encoding: e4m3fnuz on gfx942 (what its +# fp8 matrix core and vLLM's fp8 KV cache use), OCP e4m3 on gfx950. +FNUZ_ARCHS = ("gfx942",) + +# fp8 q, fp8 dots and the per-128 fp8_dsv32_mla cache are gfx950-only. The +# per-tensor fp8_scalar cache also runs on gfx942, under bf16 dots. FP8_ARCHS = ("gfx950",) +FP8_SCALAR_ARCHS = ("gfx942", "gfx950") # The packed caches (fp8_dsv4_mla, fp8_g64) and the SWA+top-k two-loop do not # reach the kernel below: they route to pa_decode_sparse, whose packed driver is @@ -103,7 +107,10 @@ def _check_packed_arch(arch: str) -> None: which says nothing about the arch being the reason. """ if arch not in PACKED_ARCHS: - flat = "bf16, fp8_scalar or fp8_dsv32_mla" if arch in FP8_ARCHS else "bf16" + if arch in FP8_ARCHS: + flat = "bf16, fp8_scalar or fp8_dsv32_mla" + else: + flat = "bf16 or fp8_scalar" if arch in FP8_SCALAR_ARCHS else "bf16" raise ValueError( f"the fp8_dsv4_mla and fp8_g64 caches and the SWA+top-k two-loop are " f"{'/'.join(PACKED_ARCHS)}-only and have no implementation on {arch}. " @@ -112,27 +119,26 @@ def _check_packed_arch(arch: str) -> None: def _check_fp8_arch(arch: str, fmt: str, q_dtype: torch.dtype) -> None: - """fp8 q and caches only where OCP e4m3 is the native fp8. + """fp8 q and caches only where the kernel decodes them. - A cache arrives as bytes, usually behind a uint8 view, so a gfx942 cache in - its native fnuz cannot be told apart from OCP here. Read as OCP it comes out - 2x too large, and NaN wherever the quantizer saturated at 240. + A cache arrives as bytes, usually behind a uint8 view, so its encoding is + taken to be the arch's native one (see FNUZ_ARCHS). """ if arch in FP8_ARCHS: return + caches = ("bf16", "fp8_scalar") if arch in FP8_SCALAR_ARCHS else ("bf16",) got = [ what for what, is_fp8 in ( (f"q is {q_dtype}", q_dtype.itemsize == 1), - (f"the cache is {fmt}", fmt != "bf16"), + (f"the cache is {fmt}", fmt not in caches), ) if is_fp8 ] if got: raise ValueError( - f"sparse_mla_fwd takes bf16 q and a bf16 cache on {arch}, but " - f"{' and '.join(got)}. The kernel reads fp8 as OCP e4m3, and {arch}'s " - "native fp8 is fnuz, which it does not decode yet." + f"sparse_mla_fwd takes bf16 q and a {' or '.join(caches)} cache on " + f"{arch}, but {' and '.join(got)}." ) @@ -260,9 +266,11 @@ def _classify_flat(kv, width, slots, kv_scale, what): """Flat pool rows are one QK row per slot; dtype and kv_scale pick the tag.""" if kv.dtype == torch.bfloat16: return "bf16" # kv_scale, if any, is ignored - if kv.dtype == torch.float8_e4m3fnuz: + arch = arch_info.get_arch() + native = torch.float8_e4m3fnuz if arch in FNUZ_ARCHS else torch.float8_e4m3fn + if kv.dtype in (torch.float8_e4m3fn, torch.float8_e4m3fnuz) and kv.dtype != native: raise ValueError( - f"{what}: float8_e4m3fnuz is the gfx942 encoding; gfx950 reads OCP e4m3" + f"{what}: {kv.dtype} is not {arch}'s fp8; the kernel reads {native}" ) if kv.element_size() != 1: raise ValueError(f"{what}: unsupported cache dtype {kv.dtype}") @@ -487,9 +495,10 @@ def sparse_mla_fwd( Supported KV cache formats. The format is inferred from kv_buffer's shape, dtype and kv_scale; each row gives what the caller has to pass. R is the - QK width, kv_lora_rank + qk_rope_head_dim. Every fp8 format, like fp8 q, is - gfx950-only: the kernel reads fp8 as OCP e4m3, and gfx942's native fp8 is - fnuz. + QK width, kv_lora_rank + qk_rope_head_dim. fp8 is read in the arch's native + encoding: e4m3fnuz on gfx942, OCP e4m3 on gfx950. gfx942 takes the + fp8_scalar cache under bf16 dots; the other fp8 formats, like fp8 q, are + gfx950-only. format kv_buffer kv_scale geometry args bf16 [slots, R], [nb, block, R], None as the model @@ -547,7 +556,8 @@ def sparse_mla_fwd( "bf16" (default): the KV tile is dequantized to bf16 on its way into LDS and both dots are bf16. Works with every cache format - the arch takes: all of them on gfx950, bf16 alone on gfx942. + the arch takes: all of them on gfx950, bf16 and fp8_scalar on + gfx942. "fp8": the cache's own code points go to the fp8 matrix core with no dequant, and the per-tensor scale folds outside the tile loop. tensor scale fp8 kv cache only; gfx950 only. @@ -900,6 +910,7 @@ def _rows(c, bs): IDX_BUFFER_LOAD=idx_use_buffer_load, HAS_INVALID=has_invalid, FP8_MFMA=fp8_dots, + FP8_FNUZ=arch in FNUZ_ARCHS, ASYNC_LDS=async_lds_on, GATHER_CACHE="", KV_LDS_PAD=kv_lds_pad, diff --git a/op_tests/triton_tests/attention/test_sparse_mla.py b/op_tests/triton_tests/attention/test_sparse_mla.py index fe5ce8e111a..9d5904840a4 100644 --- a/op_tests/triton_tests/attention/test_sparse_mla.py +++ b/op_tests/triton_tests/attention/test_sparse_mla.py @@ -13,14 +13,15 @@ import aiter.ops.triton.attention.sparse_mla as smd from aiter.ops.triton.attention.sparse_mla import ( FP8_ARCHS, + FP8_SCALAR_ARCHS, SUPPORTED_ARCHS, sparse_mla_fwd, ) from aiter.ops.triton.utils._triton import arch_info from aiter.ops.triton.utils.types import get_fp8_e4m3_dtype -# The arch-native fp8, as a producer on this machine writes it. That is OCP e4m3, -# what the kernel reads, only on FP8_ARCHS; the fp8 cases skip everywhere else. +# The arch-native fp8, as a producer on this machine writes it, which is the +# encoding the kernel reads. FP8_DTYPE = get_fp8_e4m3_dtype() FP8_MAX = torch.finfo(FP8_DTYPE).max KV_LORA, ROPE = 512, 64 @@ -31,8 +32,12 @@ def _skip_unless_supported(dots="bf16", fmt="bf16"): arch = arch_info.get_arch() if arch not in SUPPORTED_ARCHS: pytest.skip(f"sparse_mla_fwd does not support {arch}") - if (dots == "fp8" or fmt != "bf16") and arch not in FP8_ARCHS: - pytest.skip(f"fp8 is read as OCP e4m3, and {arch}'s native fp8 is fnuz") + if dots == "fp8" and arch not in FP8_ARCHS: + pytest.skip(f"fp8 dots are {'/'.join(FP8_ARCHS)}-only") + if fmt == "tensor" and arch not in FP8_SCALAR_ARCHS: + pytest.skip(f"the fp8_scalar cache is {'/'.join(FP8_SCALAR_ARCHS)}-only") + if fmt == "dsmla" and arch not in FP8_ARCHS: + pytest.skip(f"the fp8_dsv32_mla cache is {'/'.join(FP8_ARCHS)}-only") def quantize_flat_fp8(kv): @@ -186,34 +191,48 @@ def test_packed_cache_arch_gate(arch): def test_fp8_arch_gate(arch): """fp8 q and fp8 caches, for every arch, from any machine.""" smd._check_fp8_arch(arch, "bf16", torch.bfloat16) + if arch in FP8_SCALAR_ARCHS: + smd._check_fp8_arch(arch, "fp8_scalar", torch.bfloat16) for fmt, q_dtype in ( - ("fp8_scalar", torch.bfloat16), ("fp8_dsv32_mla", torch.bfloat16), ("bf16", torch.float8_e4m3fn), ): if arch in FP8_ARCHS: smd._check_fp8_arch(arch, fmt, q_dtype) else: - with pytest.raises(ValueError, match="fnuz"): + with pytest.raises(ValueError, match="takes bf16 q"): smd._check_fp8_arch(arch, fmt, q_dtype) -@pytest.mark.parametrize("fmt", ["tensor", "dsmla"]) -def test_native_fp8_cache_rejected(fmt): - """This arch's own fp8 behind a uint8 view, through the public wrapper. - - Only the arch tells those bytes apart from OCP, and going through +def test_native_dsmla_cache_rejected(): + """This arch's own fp8_dsv32_mla records behind a uint8 view, through the + public wrapper, where only fp8_scalar is supported. Going through sparse_mla_fwd also catches the gate losing its one call site. """ arch = arch_info.get_arch() if arch not in SUPPORTED_ARCHS or arch in FP8_ARCHS: - pytest.skip(f"fp8 caches run on {arch}") - q, cache, ks, idx, ptr, _ = _build(fmt, 1, 16, 64, 1024, ragged=False) + pytest.skip(f"fp8_dsv32_mla runs on {arch}") + q, cache, ks, idx, ptr, _ = _build("dsmla", 1, 16, 64, 1024, ragged=False) assert cache.dtype == torch.uint8 - with pytest.raises(ValueError, match="fnuz"): + with pytest.raises(ValueError, match="takes bf16 q"): sparse_mla_fwd(q, cache, ptr, idx, D_QK**-0.5, kv_scale=ks) +def test_foreign_fp8_cache_dtype_rejected(): + """A typed fp8 cache in the other arch's encoding is refused, not misread.""" + arch = arch_info.get_arch() + if arch not in FP8_SCALAR_ARCHS: + pytest.skip(f"no fp8 cache runs on {arch}") + q, cache, ks, idx, ptr, _ = _build("tensor", 1, 16, 64, 1024, ragged=False) + foreign = ( + torch.float8_e4m3fn + if FP8_DTYPE == torch.float8_e4m3fnuz + else torch.float8_e4m3fnuz + ) + with pytest.raises(ValueError, match=f"is not {arch}'s fp8"): + sparse_mla_fwd(q, cache.view(foreign), ptr, idx, D_QK**-0.5, kv_scale=ks) + + @pytest.mark.parametrize("arch", SUPPORTED_ARCHS) def test_launch_config_published(arch): """Every supported arch ships its launch config, checked from any machine.""" From a83baccfbd24630d6a522de5d447c05e41ece116 Mon Sep 17 00:00:00 2001 From: Jin Tao Date: Tue, 6 Oct 2026 13:43:40 +0000 Subject: [PATCH 2/4] [Triton/Gluon] [gfx942] Run fp8 dots on gfx942 dot_precision="fp8" was gfx950-only because the kernel fed the matrix core OCP e4m3. The fp8 staging, the in-kernel Q quantization (now to the arch's fp8 range, 240 on fnuz) and P now use the arch's fp8 type, and the MFMA layouts take version 3 for fnuz operands, which only have CDNA3 intrinsics. P is quantized as p * 128 (exact; p <= 1 stays under 240) and the epilogue divides it out with the V-side scale. Unscaled, softmax tails below fnuz's smallest subnormal flush to zero: with one key scoring 9 above 2047 others that carry a fifth of the output, rel-L2 error is 14.5% unscaled and 1.7% scaled (split-K off). gfx950 keeps its numerics (P scale 1, OCP, version 4). gfx942 publishes a _sparse_mla_fp8 config at BLOCK_K 64: one-byte tiles fit twice the bf16 tile's rows. The LDS model takes the element size, and matches the compiled kernels (bf16 34304/34816/38912/39424 B, fp8 34816/39936 B). MI325X, GLM-5.3-Flash shape (16 heads, rope-free 512, top-k 2048), fp8 cache under fp8 vs bf16 dots, and the bf16 cache: decode at 2 tokens 23.7 / 30.7 / 33.2 us, at 16 tokens 25.4 / 32.7 / 37.4 us; 16K-token prefill 5.78 / 9.35 / 10.55 ms. BLOCK_K 32 and 8 warps were slower. Rel-L2 vs f32 attention over the same fp8 cache, fp8 dots: 3.4% on N(0,1) inputs, 4.5% on peaky ones (bf16 dots: 0.2-0.3%). test_sparse_mla.py on gfx942: 77 passed, 1 skipped (fp8_dsv32_mla). Co-authored-by: Cursor --- .../gfx950/attention/sparse_mla.py | 46 +++++++++------- aiter/ops/triton/attention/sparse_mla.py | 54 +++++++++++-------- .../gluon/attention/sparse_mla/DEFAULT.json | 4 ++ .../triton_tests/attention/test_sparse_mla.py | 34 +++++++++--- 4 files changed, 89 insertions(+), 49 deletions(-) diff --git a/aiter/ops/triton/_gluon_kernels/gfx950/attention/sparse_mla.py b/aiter/ops/triton/_gluon_kernels/gfx950/attention/sparse_mla.py index dca8d483d76..26023362883 100644 --- a/aiter/ops/triton/_gluon_kernels/gfx950/attention/sparse_mla.py +++ b/aiter/ops/triton/_gluon_kernels/gfx950/attention/sparse_mla.py @@ -333,6 +333,9 @@ class Cfg: FP8_FNUZ: gl.constexpr FP8_TY: gl.constexpr FP8_MAX: gl.constexpr + # fp8 PV dots quantize p * P_SCALE, so small probabilities do not flush to + # zero; the epilogue divides it back out with the V-side scale. + P_SCALE: gl.constexpr # Cache policy per load site GATHER_CACHE: gl.constexpr IDX_CACHE: gl.constexpr @@ -416,6 +419,8 @@ def __init__( self.FP8_FNUZ = gl.constexpr(FP8_FNUZ) self.FP8_TY = gl.constexpr(gl.float8e4b8 if FP8_FNUZ else gl.float8e4nv) self.FP8_MAX = gl.constexpr(240.0 if FP8_FNUZ else 448.0) + # p <= 1, so 128 keeps p * P_SCALE under fnuz's 240 and is exact. + self.P_SCALE = gl.constexpr(128.0 if FP8_FNUZ else 1.0) self.GATHER_CACHE = gl.constexpr(GATHER_CACHE) self.IDX_CACHE = gl.constexpr(IDX_CACHE) self.ASYNC_LDS = gl.constexpr(ASYNC_LDS) @@ -427,13 +432,15 @@ def __init__( # bf16 dots on 16x16x32 (gfx950's full rate) on the dsv4 and STAGED_K32 walks. MFMA_K = 32 if FP8_MFMA or DSV4_WALK or STAGED_K32 else 16 self.MFMA_K = gl.constexpr(MFMA_K) + # fnuz fp8 operands only have CDNA3 (version 3) matrix-core intrinsics. + MFMA_VERSION = 3 if FP8_MFMA and FP8_FNUZ else 4 # Warps tile N; past 16 heads they also tile M, 16 heads per warp. M_WARPS = max(1, min(BLOCK_M // 16, NUM_WARPS)) self.N_WARPS = gl.constexpr(NUM_WARPS // M_WARPS) self.qk_layout = gl.constexpr( gl.amd.AMDMFMALayout( - version=4, + version=MFMA_VERSION, instr_shape=[16, 16, MFMA_K], transposed=True, warps_per_cta=[M_WARPS, NUM_WARPS // M_WARPS], @@ -441,7 +448,7 @@ def __init__( ) self.pv_layout = gl.constexpr( gl.amd.AMDMFMALayout( - version=4, + version=MFMA_VERSION, instr_shape=[16, 16, MFMA_K], transposed=True, warps_per_cta=[M_WARPS, NUM_WARPS // M_WARPS], @@ -932,7 +939,7 @@ def _qk_scores(cfg, q_dot, q_rope_dot, kv_smem, rope_smem): else: k = kv_smem.permute([1, 0]).load(cfg.k_layout) # [KV_DIM, BLOCK_K] if cfg.ASYNC_LDS: - k = k.to(gl.float8e4nv, bitcast=True) # raw cache bytes; layout-preserving + k = k.to(cfg.FP8_TY, bitcast=True) # raw cache bytes; layout-preserving S = gl.amd.cdna4.mfma(q_dot, k, S) if cfg.ROPE_SEPARATE: if cfg.ASYNC_LDS and cfg.RELAXED_LOAD: @@ -942,7 +949,7 @@ def _qk_scores(cfg, q_dot, q_rope_dot, kv_smem, rope_smem): else: k_rope = rope_smem.permute([1, 0]).load(cfg.k_layout) if cfg.ASYNC_LDS: - k_rope = k_rope.to(gl.float8e4nv, bitcast=True) + k_rope = k_rope.to(cfg.FP8_TY, bitcast=True) S = gl.amd.cdna4.mfma(q_rope_dot, k_rope, S) return S @@ -1206,9 +1213,9 @@ def _stage(cfg, seg, x_u8, sc, k_rope, kv_smem, rope_smem): # No dequant: the scale is folded outside the loop (qk_scale on the K # side, the accumulator on the V side), so what lands in LDS is exactly # what the gather returned. - kv_smem.store(x_u8.to(gl.float8e4nv, bitcast=True)) + kv_smem.store(x_u8.to(cfg.FP8_TY, bitcast=True)) if cfg.ROPE_SEPARATE: - rope_smem.store(k_rope.to(gl.float8e4nv, bitcast=True)) + rope_smem.store(k_rope.to(cfg.FP8_TY, bitcast=True)) else: _deq_store_tile(x_u8, sc, kv_smem, cfg, fmt) if fmt.KIND == "fp8_dsv4_mla": @@ -1341,13 +1348,13 @@ def _qkpv_lds( else: v = kv_smem.load(cfg.v_layout) if cfg.ASYNC_LDS: - v = v.to(gl.float8e4nv, bitcast=True) + v = v.to(cfg.FP8_TY, bitcast=True) # "fp8_scalar": V was staged as raw fp8 code points; apply the per-tensor scale # on the small side (p) and leave l scale-free: out = sum(p*s*V)/l exactly. if seg.fmt.KIND == "fp8_scalar" and not cfg.FP8_MFMA: p = p * v_scale if cfg.FP8_MFMA: - p_dot = gl.convert_layout(p.to(gl.float8e4nv), cfg.p_layout) + p_dot = gl.convert_layout((p * cfg.P_SCALE).to(cfg.FP8_TY), cfg.p_layout) else: p_dot = gl.convert_layout(p.to(gl.bfloat16), cfg.p_layout) alpha_pv = gl.convert_layout(alpha, gl.SliceLayout(1, cfg.pv_layout)) @@ -2341,8 +2348,8 @@ def _sparse_mla( "inconsistent", ) # The fp8 path needs one positive scalar scale per cache, since that is what - # folds outside the loop, and OCP e4m3 code points, which is what the matrix - # core reads. + # folds outside the loop, and the arch's own fp8 code points (FP8_TY), which + # is what the matrix core reads. gl.static_assert( (not FP8_MFMA) or (MAIN_FMT == "fp8_scalar" and (not HAS_EXTRA or EXTRA_FMT == "fp8_scalar")), @@ -2479,7 +2486,6 @@ def _sparse_mla( if FP8_MFMA and not Q_FP8: # bf16 q: quantize here, one e4m3 scale for this program's whole Q tile # (nope and rope), so the fold below is one extra factor on qk_scale. - E4M3_MAX: gl.constexpr = 448.0 q_amax = gl.max(gl.max(gl.abs(q).to(gl.float32), axis=1), axis=0) if cfg.Q_LDS: q_dot = gl.allocate_shared_memory( @@ -2522,17 +2528,17 @@ def _sparse_mla( q_amax, gl.max(gl.max(gl.abs(q_rope).to(gl.float32), axis=1), axis=0) ) q_amax = gl.maximum(q_amax, 1e-30) - q_rcp = E4M3_MAX / q_amax + q_rcp = cfg.FP8_MAX / q_amax q_dot = gl.convert_layout( - (q.to(gl.float32) * q_rcp).to(gl.float8e4nv), cfg.q_layout + (q.to(gl.float32) * q_rcp).to(cfg.FP8_TY), cfg.q_layout ) if ROPE_SEPARATE: q_rope_dot = gl.convert_layout( - (q_rope.to(gl.float32) * q_rcp).to(gl.float8e4nv), cfg.q_layout + (q_rope.to(gl.float32) * q_rcp).to(cfg.FP8_TY), cfg.q_layout ) else: q_rope_dot = q_dot - q_scale = q_amax / E4M3_MAX + q_scale = q_amax / cfg.FP8_MAX main_qk_scale = main_qk_scale * q_scale extra_qk_scale = extra_qk_scale * q_scale @@ -2549,7 +2555,7 @@ def _sparse_mla( acc = gl.zeros([BLOCK_M, HEAD_SIZE], gl.float32, layout=cfg.pv_layout) # An fp8 buffer is half the bytes of the bf16 staging it replaces. - SMEM_DT: gl.constexpr = gl.float8e4nv if FP8_MFMA else gl.bfloat16 + SMEM_DT: gl.constexpr = cfg.FP8_TY if FP8_MFMA else gl.bfloat16 # The LDS-DMA converts nothing, so an async buffer's element type has to be the # cache's own u8; the dot operands bitcast on read (free, layout-preserving). BUF_DT: gl.constexpr = gl.uint8 if ASYNC_LDS else SMEM_DT @@ -2685,10 +2691,10 @@ def _sparse_mla( ) if FP8_MFMA: - # The fp8 PV dot ran on raw code points, so the V-side scale comes off - # here, once per program instead of once per tile. l is untouched, so - # out = acc*s/l is what the bf16 path computes. - acc = acc * (extra_v_scale if HAS_EXTRA else main_v_scale) + # The fp8 PV dot ran on raw code points and p * P_SCALE, so the V-side + # scale and P_SCALE come off here, once per program instead of once per + # tile. l is untouched, so out = acc*s/l is what the bf16 path computes. + acc = acc * ((extra_v_scale if HAS_EXTRA else main_v_scale) / cfg.P_SCALE) # Move the row reductions into pv-slice space for output/partials. m_pv = gl.convert_layout(m_i, gl.SliceLayout(1, cfg.pv_layout)) diff --git a/aiter/ops/triton/attention/sparse_mla.py b/aiter/ops/triton/attention/sparse_mla.py index 7c1d3b766f0..ee0346389fe 100644 --- a/aiter/ops/triton/attention/sparse_mla.py +++ b/aiter/ops/triton/attention/sparse_mla.py @@ -77,8 +77,8 @@ def _cache_pointers(fmt, kv, d_qk, kv_scale): # fp8 matrix core and vLLM's fp8 KV cache use), OCP e4m3 on gfx950. FNUZ_ARCHS = ("gfx942",) -# fp8 q, fp8 dots and the per-128 fp8_dsv32_mla cache are gfx950-only. The -# per-tensor fp8_scalar cache also runs on gfx942, under bf16 dots. +# fp8 q and the per-128 fp8_dsv32_mla cache are gfx950-only. The per-tensor +# fp8_scalar cache, and the fp8 dots that run on it, also run on gfx942. FP8_ARCHS = ("gfx950",) FP8_SCALAR_ARCHS = ("gfx942", "gfx950") @@ -89,15 +89,20 @@ def _cache_pointers(fmt, kv, d_qk, kv_scale): PACKED_ARCHS = ("gfx950",) -def _get_config(arch: str | None = None) -> dict: +def _get_config(arch: str | None = None, fp8_dots: bool = False) -> dict: """The _sparse_mla launch config published for arch, the running one by default. BLOCK_K is per arch because gfx950's tile does not fit gfx942's 64 KB of LDS. num_warps is its own entry rather than BLOCK_K // 16, so the smaller - tile does not halve the warps too. + tile does not halve the warps too. An arch may publish a separate + _sparse_mla_fp8 entry for fp8 dots, whose one-byte tiles fit twice the + BLOCK_K in the same LDS. """ cfg_dir = resolve_config_dir("attention", "SPARSE_MLA", backend="gluon", arch=arch) - return dict(load_config_json(f"{cfg_dir}/DEFAULT.json")["_sparse_mla"]) + configs = load_config_json(f"{cfg_dir}/DEFAULT.json") + if fp8_dots and "_sparse_mla_fp8" in configs: + return dict(configs["_sparse_mla_fp8"]) + return dict(configs["_sparse_mla"]) def _check_packed_arch(arch: str) -> None: @@ -142,15 +147,18 @@ def _check_fp8_arch(arch: str, fmt: str, q_dtype: torch.dtype) -> None: ) -# Row pitch padding, and the scratch the kernel takes beyond the tiles. Both -# hold only for bf16 tiles with the async path off, which is every gfx942 -# launch: fp8 dots are rejected there and its launch config keeps ASYNC_LDS off. -# A nonzero KV_LDS_PAD replaces the pad on the KV tile alone. +# Row pitch padding, and the scratch the kernel takes beyond the tiles, in tile +# elements, with the async path off, which is every gfx942 launch. bf16 tiles +# pad by 8; fp8 tiles (fp8 dots) are one byte per element and pad by 16. A +# nonzero KV_LDS_PAD replaces the pad on the KV tile alone. _LDS_PAD = 8 -_LDS_SCRATCH_PER_BLOCK_K = 32 +_LDS_PAD_FP8 = 16 +_LDS_SCRATCH_PER_BLOCK_K = 16 -def _check_lds_budget(arch, block_k, kv_lora_rank, qk_rope_head_dim, kv_lds_pad): +def _check_lds_budget( + arch, block_k, kv_lora_rank, qk_rope_head_dim, kv_lds_pad, fp8_dots=False +): """Reject a geometry whose tiles cannot fit, naming what would. kv_lds_pad is the launch's KV_LDS_PAD, so prefill is checked at the wider @@ -162,9 +170,10 @@ def _check_lds_budget(arch, block_k, kv_lora_rank, qk_rope_head_dim, kv_lds_pad) if arch != "gfx942": return budget = arch_info._LDS_CAP_BYTES[arch] - rope = block_k * (qk_rope_head_dim + _LDS_PAD) * 2 if qk_rope_head_dim else 0 - need = ( - block_k * (kv_lora_rank + (kv_lds_pad or _LDS_PAD)) * 2 + elt, pad = (1, _LDS_PAD_FP8) if fp8_dots else (2, _LDS_PAD) + rope = block_k * (qk_rope_head_dim + pad) if qk_rope_head_dim else 0 + need = elt * ( + block_k * (kv_lora_rank + (kv_lds_pad or pad)) + rope + _LDS_SCRATCH_PER_BLOCK_K * block_k ) @@ -239,10 +248,9 @@ def _resolve_dot_precision(dot_precision: str, fmt: str, arch: str) -> bool: ) if dot_precision == "bf16": return False - if arch not in FP8_ARCHS: + if arch not in FP8_SCALAR_ARCHS: raise ValueError( - f"dot_precision='fp8' is not supported on {arch}: the kernel feeds the " - f"matrix core OCP e4m3, but {arch}'s native fp8 is fnuz. Use " + f"dot_precision='fp8' is not supported on {arch}. Use " "dot_precision='bf16'." ) if fmt == "fp8_dsv32_mla": @@ -497,8 +505,8 @@ def sparse_mla_fwd( dtype and kv_scale; each row gives what the caller has to pass. R is the QK width, kv_lora_rank + qk_rope_head_dim. fp8 is read in the arch's native encoding: e4m3fnuz on gfx942, OCP e4m3 on gfx950. gfx942 takes the - fp8_scalar cache under bf16 dots; the other fp8 formats, like fp8 q, are - gfx950-only. + fp8_scalar cache under either dot precision; the other fp8 formats, like + fp8 q, are gfx950-only. format kv_buffer kv_scale geometry args bf16 [slots, R], [nb, block, R], None as the model @@ -560,7 +568,7 @@ def sparse_mla_fwd( gfx942. "fp8": the cache's own code points go to the fp8 matrix core with no dequant, and the per-tensor scale folds outside the tile loop. - tensor scale fp8 kv cache only; gfx950 only. + tensor scale fp8 kv cache only; gfx942 and gfx950. q is adapted to the choice. bf16 q is quantized in the kernel prologue, one scale per (query, head-block) tile; fp8 q is passed @@ -723,7 +731,7 @@ def sparse_mla_fwd( # Tuned launch config (gfx950 / MI355). H < 16 runs natively at # BLOCK_M = next_pow2(H) instead of padding heads block_m = 16 if num_heads >= 16 else max(8, 1 << (num_heads - 1).bit_length()) - cfg = _get_config() + cfg = _get_config(fp8_dots=fp8_dots) block_k = cfg["BLOCK_K"] num_warps = cfg["num_warps"] @@ -825,7 +833,9 @@ def _rows(c, bs): kv_lds_pad = ( 16 if num_queries >= _PREFILL_MIN_ROWS and not fp8_dots and not own_pad else 0 ) - _check_lds_budget(arch, block_k, kv_lora_rank, qk_rope_head_dim, kv_lds_pad) + _check_lds_budget( + arch, block_k, kv_lora_rank, qk_rope_head_dim, kv_lds_pad, fp8_dots + ) # The 32/64-head programs and the decode XCD remap below are tuned on gfx950. staged = fmt == "fp8_scalar" or (fmt == "bf16" and qk_rope_head_dim == 0) diff --git a/aiter/ops/triton/configs/gfx942/gluon/attention/sparse_mla/DEFAULT.json b/aiter/ops/triton/configs/gfx942/gluon/attention/sparse_mla/DEFAULT.json index 41b7c0519c8..ad17ddc2433 100644 --- a/aiter/ops/triton/configs/gfx942/gluon/attention/sparse_mla/DEFAULT.json +++ b/aiter/ops/triton/configs/gfx942/gluon/attention/sparse_mla/DEFAULT.json @@ -2,5 +2,9 @@ "_sparse_mla": { "BLOCK_K": 32, "num_warps": 4 + }, + "_sparse_mla_fp8": { + "BLOCK_K": 64, + "num_warps": 4 } } diff --git a/op_tests/triton_tests/attention/test_sparse_mla.py b/op_tests/triton_tests/attention/test_sparse_mla.py index 9d5904840a4..e4d0d27dbfc 100644 --- a/op_tests/triton_tests/attention/test_sparse_mla.py +++ b/op_tests/triton_tests/attention/test_sparse_mla.py @@ -32,8 +32,8 @@ def _skip_unless_supported(dots="bf16", fmt="bf16"): arch = arch_info.get_arch() if arch not in SUPPORTED_ARCHS: pytest.skip(f"sparse_mla_fwd does not support {arch}") - if dots == "fp8" and arch not in FP8_ARCHS: - pytest.skip(f"fp8 dots are {'/'.join(FP8_ARCHS)}-only") + if dots == "fp8" and arch not in FP8_SCALAR_ARCHS: + pytest.skip(f"fp8 dots are {'/'.join(FP8_SCALAR_ARCHS)}-only") if fmt == "tensor" and arch not in FP8_SCALAR_ARCHS: pytest.skip(f"the fp8_scalar cache is {'/'.join(FP8_SCALAR_ARCHS)}-only") if fmt == "dsmla" and arch not in FP8_ARCHS: @@ -162,14 +162,15 @@ def test_sparse_mla(fmt, dots, tol, H, C, topk, ragged, pool): def test_dot_precision_arch_gate(arch): """The fp8-dot gate, for every arch, from any machine. - The matrix above skips its fp8 cases off gfx950, so the gate has no - coverage on any arch without this. + The matrix above only runs its fp8 cases on the arch it is on. """ assert smd._resolve_dot_precision("bf16", "fp8_scalar", arch) is False - if arch in FP8_ARCHS: + if arch in FP8_SCALAR_ARCHS: assert smd._resolve_dot_precision("fp8", "fp8_scalar", arch) is True + with pytest.raises(ValueError, match="needs an fp8 cache"): + smd._resolve_dot_precision("fp8", "bf16", arch) else: - with pytest.raises(ValueError, match="fnuz"): + with pytest.raises(ValueError, match="not supported"): smd._resolve_dot_precision("fp8", "fp8_scalar", arch) @@ -237,12 +238,13 @@ def test_foreign_fp8_cache_dtype_rejected(): def test_launch_config_published(arch): """Every supported arch ships its launch config, checked from any machine.""" assert {"BLOCK_K", "num_warps"} <= smd._get_config(arch).keys() + assert {"BLOCK_K", "num_warps"} <= smd._get_config(arch, fp8_dots=True).keys() def test_lds_budget_gfx950_is_unchecked(): """gfx950 is left to the launcher, though arch_info lists its LDS too. - The footprint model holds only for gfx942's bf16, non-async tiles. + The footprint model holds only for gfx942's non-async tiles. """ smd._check_lds_budget("gfx950", 64, 2048, 64, kv_lds_pad=16) @@ -288,6 +290,24 @@ def test_lds_budget_gfx942_boundary(kv_lds_pad, kv_lora_rank, rope, need): smd._check_lds_budget("gfx942", block_k, kv_lora_rank, rope, kv_lds_pad) +@pytest.mark.parametrize( + "kv_lora_rank, rope, need", + [(512, 0, None), (512, 64, None), (1024, 0, 67584), (512, 512, 68608)], +) +def test_lds_budget_gfx942_fp8_tile(kv_lora_rank, rope, need): + """CPU-only: fp8 dots' one-byte tiles, at their own BLOCK_K on gfx942. + + 512/0 and 512/64 compile to the 34816 and 39936 B the model gives them. + """ + block_k = smd._get_config("gfx942", fp8_dots=True)["BLOCK_K"] + args = ("gfx942", block_k, kv_lora_rank, rope, 0) + if need is None: + smd._check_lds_budget(*args, fp8_dots=True) + return + with pytest.raises(ValueError, match=rf"needs {need} B of LDS"): + smd._check_lds_budget(*args, fp8_dots=True) + + def test_ds_mla_format(): _skip_unless_supported(fmt="dsmla") _run_and_check("dsmla", C=8, H=16, topk=2048, ragged=True) From 2782fad36b9a1a9a906872ec40750f9508da35c7 Mon Sep 17 00:00:00 2001 From: Jin Tao Date: Thu, 8 Oct 2026 08:25:39 +0000 Subject: [PATCH 3/4] [Triton/Gluon] [gfx942] Name the fp8 switches in sparse MLA traces, bench fp8 dots on gfx942 _sparse_mla_repr left out FP8_MFMA and FP8_FNUZ, so bf16 and fp8 dots over the same fp8_scalar cache got the same trace name whenever they ran at the same BLOCK_K, which gfx950 does for prefill-sized grids and with the async path off. Add both keys. bench_sparse_mla.py gated its fp8-dot series on FP8_ARCHS and pre-quantized Q, so gfx942 skipped the path with a stale note that the kernel reads OCP e4m3. Gate on FP8_SCALAR_ARCHS, pre-quantize Q only on FP8_ARCHS (gfx942 quantizes bf16 Q inside the kernel), and count Q at its own element size in the bandwidth metric. MI325X decode at 1/8/64 sequences (16 heads, context 8192, top-k 2048): bf16 dots 39.5/43.6/103.8 us, fp8 dots 25.0/27.0/72.3 us. test_sparse_mla.py: 77 passed, 1 skipped. Co-authored-by: Cursor --- .../gfx950/attention/sparse_mla.py | 11 ++++++- .../op_benchmarks/triton/bench_sparse_mla.py | 31 ++++++++++++------- 2 files changed, 29 insertions(+), 13 deletions(-) diff --git a/aiter/ops/triton/_gluon_kernels/gfx950/attention/sparse_mla.py b/aiter/ops/triton/_gluon_kernels/gfx950/attention/sparse_mla.py index 26023362883..5c790061364 100644 --- a/aiter/ops/triton/_gluon_kernels/gfx950/attention/sparse_mla.py +++ b/aiter/ops/triton/_gluon_kernels/gfx950/attention/sparse_mla.py @@ -2163,7 +2163,16 @@ def _xcd_work(GRID_ORDER: gl.constexpr, NUM_XCDS: gl.constexpr, SPLIT_K: gl.cons _sparse_mla_repr = make_kernel_repr( "_sparse_mla", - ["BLOCK_M", "BLOCK_K", "HEAD_SIZE", "SPLIT_K", "MAIN_FMT", "ROPE_SEPARATE"], + [ + "BLOCK_M", + "BLOCK_K", + "HEAD_SIZE", + "SPLIT_K", + "MAIN_FMT", + "ROPE_SEPARATE", + "FP8_MFMA", + "FP8_FNUZ", + ], ) diff --git a/op_tests/op_benchmarks/triton/bench_sparse_mla.py b/op_tests/op_benchmarks/triton/bench_sparse_mla.py index ecc019ba5d3..b8dc594cd26 100644 --- a/op_tests/op_benchmarks/triton/bench_sparse_mla.py +++ b/op_tests/op_benchmarks/triton/bench_sparse_mla.py @@ -21,6 +21,7 @@ from aiter.ops.triton.attention.sparse_mla import ( FP8_ARCHS, + FP8_SCALAR_ARCHS, SUPPORTED_ARCHS, sparse_mla_fwd, ) @@ -74,7 +75,7 @@ def device_time_ms(func, warmup=25, rep=100, flush=True): return total / rep / 1e3 -def bytes_moved(num_tokens, num_heads, nnz, kv_elem_bytes): +def bytes_moved(num_tokens, num_heads, nnz, kv_elem_bytes, q_elem_bytes): """Bytes the launch actually moves, counting rows as gathered. Split-K partials are left out: the split count is the kernel's own decision, @@ -82,7 +83,7 @@ def bytes_moved(num_tokens, num_heads, nnz, kv_elem_bytes): """ kv = nnz * D_QK * kv_elem_bytes # the gather, and the bulk of it idx = nnz * 4 # int32 index stream, read once - q = num_tokens * num_heads * D_QK * kv_elem_bytes + q = num_tokens * num_heads * D_QK * q_elem_bytes out = num_tokens * num_heads * KV_LORA_RANK * 2 return kv + idx + q + out @@ -156,19 +157,19 @@ def run_benchmark(args): for tokens in args.num_tokens ] - # sparse_mla_fwd raises on fp8 where the arch's native fp8 is not OCP e4m3, - # and triton's harness does not catch it, so drop the series instead. + # sparse_mla_fwd raises on fp8 dots outside FP8_SCALAR_ARCHS, and triton's + # harness does not catch it, so drop the series instead. dot_vals = ["bf16"] dot_names = ["bf16 dots"] dot_styles = [("green", "-")] - if arch_info.get_arch() in FP8_ARCHS: + if arch_info.get_arch() in FP8_SCALAR_ARCHS: dot_vals.append("fp8") dot_names.append("fp8 dots") dot_styles.append(("blue", "-")) else: print( - f"note: skipping the fp8-dot series, {arch_info.get_arch()}'s native " - "fp8 is fnuz and the kernel reads OCP e4m3" + f"note: skipping the fp8-dot series, {arch_info.get_arch()} has no " + "fp8 dots" ) benchmark = triton.testing.Benchmark( @@ -199,10 +200,14 @@ def bench_sparse_mla( device=q.device, ) if dots == "fp8": - # Hand q over already quantized, the way production does. - cache, scale = kv_fp8, kv_scale - q_scale = (q.float().abs().amax() / E4M3_MAX).clamp_min(1e-30).reshape(1) - q = (q.float() / q_scale).clamp(-E4M3_MAX, E4M3_MAX).to(E4M3_DTYPE) + cache, scale, q_scale = kv_fp8, kv_scale, None + if arch_info.get_arch() in FP8_ARCHS: + # Hand q over already quantized, the way production does. The + # other archs take bf16 q only and quantize it inside the kernel. + q_scale = ( + (q.float().abs().amax() / E4M3_MAX).clamp_min(1e-30).reshape(1) + ) + q = (q.float() / q_scale).clamp(-E4M3_MAX, E4M3_MAX).to(E4M3_DTYPE) else: cache, scale, q_scale = kv, None, None @@ -226,7 +231,9 @@ def func(): num_tokens = q.shape[0] # QK reads the whole row, PV only the latent half flops = 2.0 * num_heads * nnz * (D_QK + KV_LORA_RANK) - moved = bytes_moved(num_tokens, num_heads, nnz, 1 if dots == "fp8" else 2) + moved = bytes_moved( + num_tokens, num_heads, nnz, cache.element_size(), q.element_size() + ) if metric == "time": return time_ms From 013f408ebe8491e80ddab1fbab4c3dbcbc98bd74 Mon Sep 17 00:00:00 2001 From: Jin Tao Date: Thu, 8 Oct 2026 08:45:20 +0000 Subject: [PATCH 4/4] [Triton/Gluon] [gfx942] Refuse non-e4m3 byte caches, test the P_SCALE softmax tail _classify_flat took any one-byte dtype other than the other arch's e4m3, and the kernel then read the bytes as this arch's e4m3, so a typed e5m2 (or int8, bool) cache was decoded with the wrong encoding. Accept only the native e4m3 dtype or a uint8 view of it, as the sparse_mla_fwd docstring says. The fp8 accuracy cases use mild random inputs, where every p stays near 1, and they all pass with P_SCALE forced to 1. test_fp8_dots_keep_the_softmax_tail gives each query a first key scoring 9 above 2047 others that carry a fifth of the output at p ~ 1e-4, with split-K off. On MI325X the fp8-dot output is 1.8% off its f32 reference (max-rel); with P_SCALE = 1 the tail's share comes out zero and the error is 100%. It runs on gfx942, the arch with a P_SCALE. test_sparse_mla.py on gfx942: 82 passed, 1 skipped. Co-authored-by: Cursor --- aiter/ops/triton/attention/sparse_mla.py | 7 +- .../triton_tests/attention/test_sparse_mla.py | 64 +++++++++++++++++++ 2 files changed, 69 insertions(+), 2 deletions(-) diff --git a/aiter/ops/triton/attention/sparse_mla.py b/aiter/ops/triton/attention/sparse_mla.py index ee0346389fe..75e7a7e9d30 100644 --- a/aiter/ops/triton/attention/sparse_mla.py +++ b/aiter/ops/triton/attention/sparse_mla.py @@ -280,8 +280,11 @@ def _classify_flat(kv, width, slots, kv_scale, what): raise ValueError( f"{what}: {kv.dtype} is not {arch}'s fp8; the kernel reads {native}" ) - if kv.element_size() != 1: - raise ValueError(f"{what}: unsupported cache dtype {kv.dtype}") + if kv.dtype not in (native, torch.uint8): + raise ValueError( + f"{what}: unsupported cache dtype {kv.dtype}; a flat fp8 cache is " + f"{native} or a uint8 view of it" + ) if kv_scale is None: raise ValueError( f"{what}: a flat fp8 cache needs kv_scale, [1] f32 (fp8_scalar) or " diff --git a/op_tests/triton_tests/attention/test_sparse_mla.py b/op_tests/triton_tests/attention/test_sparse_mla.py index e4d0d27dbfc..632c94c761f 100644 --- a/op_tests/triton_tests/attention/test_sparse_mla.py +++ b/op_tests/triton_tests/attention/test_sparse_mla.py @@ -12,6 +12,7 @@ import aiter.ops.triton.attention.sparse_mla as smd from aiter.ops.triton.attention.sparse_mla import ( + FNUZ_ARCHS, FP8_ARCHS, FP8_SCALAR_ARCHS, SUPPORTED_ARCHS, @@ -234,6 +235,19 @@ def test_foreign_fp8_cache_dtype_rejected(): sparse_mla_fwd(q, cache.view(foreign), ptr, idx, D_QK**-0.5, kv_scale=ks) +@pytest.mark.parametrize( + "dtype", [torch.float8_e5m2, torch.float8_e5m2fnuz, torch.int8, torch.bool] +) +def test_non_e4m3_byte_cache_dtype_rejected(dtype): + """Other one-byte dtypes are refused rather than decoded as e4m3.""" + arch = arch_info.get_arch() + if arch not in FP8_SCALAR_ARCHS: + pytest.skip(f"no fp8 cache runs on {arch}") + q, cache, ks, idx, ptr, _ = _build("tensor", 1, 16, 64, 1024, ragged=False) + with pytest.raises(ValueError, match="unsupported cache dtype"): + sparse_mla_fwd(q, cache.view(dtype), ptr, idx, D_QK**-0.5, kv_scale=ks) + + @pytest.mark.parametrize("arch", SUPPORTED_ARCHS) def test_launch_config_published(arch): """Every supported arch ships its launch config, checked from any machine.""" @@ -389,3 +403,53 @@ def test_sparse_mla_rope_free(fmt, dots, tol, H, C, topk, ragged, pool): rope=0, dot_precision=dots, ) + + +def test_fp8_dots_keep_the_softmax_tail(): + """Each query's first key scores 9 above the other 2047, which carry a fifth + of the output at p ~ 1e-4 each. That is below fnuz's smallest subnormal + unless P_SCALE lifts p before the PV dot, and with split-K off a single + program quantizes every tail p after the max. + """ + arch = arch_info.get_arch() + if arch not in FNUZ_ARCHS: + pytest.skip(f"P_SCALE is 1 on {arch}") + C, H, topk, pool = 4, 16, 2048, 1 << 13 + sm = KV_LORA**-0.5 + g = torch.Generator().manual_seed(0) + # Head h reads axis h alone, where the dominant key is 1 (score 9) and the + # tail keys are 0 (score 0). Elsewhere the tail keys share one direction. + q = torch.zeros(C, H, KV_LORA) + q[:, range(H), range(H)] = 9.0 / sm + u = torch.randn(KV_LORA, generator=g) + kv = torch.nn.functional.normalize( + u + 0.3 * torch.randn(pool, KV_LORA, generator=g), dim=-1 + ) + kv = kv * KV_LORA**0.5 + kv[:, :H] = 0 + kv[-1] = 0 + kv[-1, :H] = 1 + idx = torch.stack( + [ + torch.cat([torch.tensor([pool - 1]), torch.randperm(pool - 1, generator=g)]) + for _ in range(C) + ] + )[:, :topk] + idx = idx.flatten().to(torch.int32).cuda() + ptr = torch.arange(0, (C + 1) * topk, topk, dtype=torch.int32, device="cuda") + q = q.to(torch.bfloat16).cuda() + cache, ks = quantize_flat_fp8(kv.cuda()) + ref = reference(q, dequant_flat_fp8(cache, ks).to(torch.bfloat16), idx, ptr, sm) + out, _ = sparse_mla_fwd( + q, + cache, + ptr, + idx, + sm, + kv_scale=ks, + qk_rope_head_dim=0, + dot_precision="fp8", + kv_splits=1, + ) + e = rel_err(out, ref) + assert e < 5e-2, f"softmax tail: rel-err {e:.3e}"