diff --git a/python/sglang/jit_kernel/csrc/deepseek_v4/fused_norm_rope_v2.cuh b/python/sglang/jit_kernel/csrc/deepseek_v4/fused_norm_rope_v2.cuh index 0e789263c..9374b4c47 100644 --- a/python/sglang/jit_kernel/csrc/deepseek_v4/fused_norm_rope_v2.cuh +++ b/python/sglang/jit_kernel/csrc/deepseek_v4/fused_norm_rope_v2.cuh @@ -202,7 +202,7 @@ INDEXER_KERNEL void fused_norm_rope_indexer(const __grid_constant__ FusedNormRop local_max = math::max(local_max, math::abs(data[i])); } const auto abs_max = warp::reduce_max(local_max); - const auto scale = fmaxf(1e-4f, abs_max) / math::FP8_E4M3_MAX; + const auto scale = fmaxf(1e-4f, abs_max) / kFP8E4M3Max; const auto inv_scale = 1.0f / scale; const int64_t page = out_loc >> kPageBits; const int64_t offset = out_loc & ((1 << kPageBits) - 1); @@ -488,7 +488,7 @@ FLASHMLA_KERNEL void fused_norm_rope_flashmla(const __grid_constant__ FusedNormR const auto x = cast(cast(data[0])); const auto y = cast(cast(data[1])); const auto abs_max = warp::reduce_max(fmaxf(fabs(x), fabs(y))); - const auto scale_raw = fmaxf(1e-4f, abs_max) / math::FP8_E4M3_MAX; + const auto scale_raw = fmaxf(1e-4f, abs_max) / kFP8E4M3Max; const auto scale_ue8m0 = cast_to_ue8m0(scale_raw); const auto inv_scale = inv_scale_ue8m0(scale_ue8m0); const auto result = pack_fp8(x * inv_scale, y * inv_scale); diff --git a/python/sglang/jit_kernel/csrc/deepseek_v4/store.cuh b/python/sglang/jit_kernel/csrc/deepseek_v4/store.cuh index 49f6f5596..96c548854 100644 --- a/python/sglang/jit_kernel/csrc/deepseek_v4/store.cuh +++ b/python/sglang/jit_kernel/csrc/deepseek_v4/store.cuh @@ -52,7 +52,7 @@ __global__ void fused_store_flashmla_cache(const __grid_constant__ FusedStoreCac if (wid != 7) { const auto [x, y] = cast(elems); const auto abs_max = warp::reduce_max(fmaxf(fabs(x), fabs(y))); - const auto scale_raw = fmaxf(1e-4f, abs_max) / math::FP8_E4M3_MAX; + const auto scale_raw = fmaxf(1e-4f, abs_max) / kFP8E4M3Max; const auto scale_ue8m0 = cast_to_ue8m0(scale_raw); const auto inv_scale = inv_scale_ue8m0(scale_ue8m0); const auto result = pack_fp8(x * inv_scale, y * inv_scale); @@ -104,7 +104,7 @@ __global__ void fused_store_indexer_cache(const __grid_constant__ FusedStoreCach const auto local_max = fmaxf(fmaxf(fabs(x0), fabs(x1)), fmaxf(fabs(y0), fabs(y1))); const auto abs_max = warp::reduce_max(local_max); // use normal fp32 scale - const auto scale = fmaxf(1e-4f, abs_max) / math::FP8_E4M3_MAX; + const auto scale = fmaxf(1e-4f, abs_max) / kFP8E4M3Max; const auto inv_scale = 1.0f / scale; const int32_t page = index >> kPageBits; const int32_t offset = index & ((1 << kPageBits) - 1); diff --git a/python/sglang/jit_kernel/csrc/dsa/fused_store_index_cache.cuh b/python/sglang/jit_kernel/csrc/dsa/fused_store_index_cache.cuh index e649fda57..5ba756894 100644 --- a/python/sglang/jit_kernel/csrc/dsa/fused_store_index_cache.cuh +++ b/python/sglang/jit_kernel/csrc/dsa/fused_store_index_cache.cuh @@ -26,7 +26,7 @@ struct FusedStoreCacheParam { [[maybe_unused]] SGL_DEVICE float fp8_e4m3_clip(float val) { namespace math = device::math; - return math::max(math::min(val, math::FP8_E4M3_MAX), -math::FP8_E4M3_MAX); + return math::max(math::min(val, kFP8E4M3Max), -kFP8E4M3Max); } [[maybe_unused]] @@ -63,7 +63,7 @@ __global__ void fused_store_indexer_cache(const __grid_constant__ FusedStoreCach const auto local_max = fmaxf(fmaxf(fabs(x0), fabs(x1)), fmaxf(fabs(y0), fabs(y1))); const auto abs_max = warp::reduce_max(local_max); // use normal fp32 scale - const auto scale = fmaxf(1e-4f, abs_max) / math::FP8_E4M3_MAX; + const auto scale = fmaxf(1e-4f, abs_max) / kFP8E4M3Max; const auto inv_scale = 1.0f / scale; const int32_t page = index >> kPageBits; const int32_t offset = index & ((1 << kPageBits) - 1); diff --git a/python/sglang/jit_kernel/include/sgl_kernel/deepseek_v4/fp8_utils.cuh b/python/sglang/jit_kernel/include/sgl_kernel/deepseek_v4/fp8_utils.cuh index 53a62755b..214c6f266 100644 --- a/python/sglang/jit_kernel/include/sgl_kernel/deepseek_v4/fp8_utils.cuh +++ b/python/sglang/jit_kernel/include/sgl_kernel/deepseek_v4/fp8_utils.cuh @@ -74,7 +74,15 @@ SGL_DEVICE uint8_t cvt_float_to_fp8_e4m3(float val) { uint8_t mant3; if (exp32 < kMinSubnormExp) { +#if HIP_FP8_TYPE_FNUZ + // E4M3FNUZ (gfx942) has no negative zero: byte 0x80 is NaN, not -0.0. + // Returning `sign` (0x80) for an underflowing negative injects NaN into the + // fp8 KV cache -> NaN attention/logits. Flush underflow to +0 instead. + return 0; +#else + // E4M3FN (gfx950): 0x80 == -0.0, harmless. return sign; +#endif } else if (exp32 < kMinNormExp) { // Subnormal range int32_t shift = -(kBias - 1) - exp32; // 1..3 diff --git a/python/sglang/jit_kernel/include/sgl_kernel/math.cuh b/python/sglang/jit_kernel/include/sgl_kernel/math.cuh index 4f9ac4814..2b203e4b1 100644 --- a/python/sglang/jit_kernel/include/sgl_kernel/math.cuh +++ b/python/sglang/jit_kernel/include/sgl_kernel/math.cuh @@ -17,7 +17,11 @@ inline constexpr float log2e = 1.44269504088896340736f; /// \brief Constant: ln(2) inline constexpr float loge2 = 0.693147180559945309417f; /// \brief Maximum representable value for FP8 E4M3 format. -inline constexpr float FP8_E4M3_MAX = 448.0f; +/// Arch-aware: 448 on CUDA / AMD OCP e4m3fn (gfx950), 224 on AMD e4m3fnuz +/// (gfx942). Mirrors kFP8E4M3Max so fp8 quant scale divisors and clamps in +/// the dsv4 compute path (indexer Q-quant, MoE silu+mul / dispatch quant, +/// GEMM per-tensor quant) do not over-saturate fnuz hardware. +inline constexpr float FP8_E4M3_MAX = ::kFP8E4M3Max; static_assert(log2e * loge2 == 1.0f, "log2e * loge2 must be 1"); /// \brief Returns the larger of `a` and `b`. diff --git a/python/sglang/srt/layers/attention/nsa/triton_decode/triton_mla_kernels_decode_fused.py b/python/sglang/srt/layers/attention/nsa/triton_decode/triton_mla_kernels_decode_fused.py index 1cc82a989..c0a4f3bc8 100644 --- a/python/sglang/srt/layers/attention/nsa/triton_decode/triton_mla_kernels_decode_fused.py +++ b/python/sglang/srt/layers/attention/nsa/triton_decode/triton_mla_kernels_decode_fused.py @@ -17,6 +17,12 @@ import torch import triton import triton.language as tl +from sglang.srt.utils.common import is_gfx942_supported + +# gfx942/MI300/MI325 stores e4m3fnuz (bias 8); +# gfx950/MI350 and CUDA store OCP e4m3fn (bias 7). +_KV_FP8_TY = tl.float8e4b8 if is_gfx942_supported() else tl.float8e4nv + def _bucket_total_tokens(total_tokens: int) -> int: """Round total_tokens up to the nearest power of 2 for autotune key stability.""" @@ -213,37 +219,37 @@ def _process_kv_block_aggressive( qk = tl.zeros([BLOCK_H, BLOCK_N], dtype=tl.float32) - nope_fp8_0 = nope_uint8_0.to(tl.float8e4nv, bitcast=True) + nope_fp8_0 = nope_uint8_0.to(_KV_FP8_TY, bitcast=True) kv_0 = (nope_fp8_0.to(tl.bfloat16) * scale_bf16_0[:, None]).to(tl.bfloat16) kv_0 = tl.where(valid_2d, kv_0, 0.0) qk += tl.dot(q_0, tl.trans(kv_0)).to(tl.float32) - nope_fp8_1 = nope_uint8_1.to(tl.float8e4nv, bitcast=True) + nope_fp8_1 = nope_uint8_1.to(_KV_FP8_TY, bitcast=True) kv_1 = (nope_fp8_1.to(tl.bfloat16) * scale_bf16_1[:, None]).to(tl.bfloat16) kv_1 = tl.where(valid_2d, kv_1, 0.0) qk += tl.dot(q_1, tl.trans(kv_1)).to(tl.float32) - nope_fp8_2 = nope_uint8_2.to(tl.float8e4nv, bitcast=True) + nope_fp8_2 = nope_uint8_2.to(_KV_FP8_TY, bitcast=True) kv_2 = (nope_fp8_2.to(tl.bfloat16) * scale_bf16_2[:, None]).to(tl.bfloat16) kv_2 = tl.where(valid_2d, kv_2, 0.0) qk += tl.dot(q_2, tl.trans(kv_2)).to(tl.float32) - nope_fp8_3 = nope_uint8_3.to(tl.float8e4nv, bitcast=True) + nope_fp8_3 = nope_uint8_3.to(_KV_FP8_TY, bitcast=True) kv_3 = (nope_fp8_3.to(tl.bfloat16) * scale_bf16_3[:, None]).to(tl.bfloat16) kv_3 = tl.where(valid_2d, kv_3, 0.0) qk += tl.dot(q_3, tl.trans(kv_3)).to(tl.float32) - nope_fp8_4 = nope_uint8_4.to(tl.float8e4nv, bitcast=True) + nope_fp8_4 = nope_uint8_4.to(_KV_FP8_TY, bitcast=True) kv_4 = (nope_fp8_4.to(tl.bfloat16) * scale_bf16_4[:, None]).to(tl.bfloat16) kv_4 = tl.where(valid_2d, kv_4, 0.0) qk += tl.dot(q_4, tl.trans(kv_4)).to(tl.float32) - nope_fp8_5 = nope_uint8_5.to(tl.float8e4nv, bitcast=True) + nope_fp8_5 = nope_uint8_5.to(_KV_FP8_TY, bitcast=True) kv_5 = (nope_fp8_5.to(tl.bfloat16) * scale_bf16_5[:, None]).to(tl.bfloat16) kv_5 = tl.where(valid_2d, kv_5, 0.0) qk += tl.dot(q_5, tl.trans(kv_5)).to(tl.float32) - nope_fp8_6 = nope_uint8_6.to(tl.float8e4nv, bitcast=True) + nope_fp8_6 = nope_uint8_6.to(_KV_FP8_TY, bitcast=True) kv_6 = (nope_fp8_6.to(tl.bfloat16) * scale_bf16_6[:, None]).to(tl.bfloat16) kv_6 = tl.where(valid_2d, kv_6, 0.0) qk += tl.dot(q_6, tl.trans(kv_6)).to(tl.float32) diff --git a/python/sglang/srt/layers/fused_qk_norm_rope_store.py b/python/sglang/srt/layers/fused_qk_norm_rope_store.py index 01a6ed8b7..1aa19ede5 100644 --- a/python/sglang/srt/layers/fused_qk_norm_rope_store.py +++ b/python/sglang/srt/layers/fused_qk_norm_rope_store.py @@ -96,6 +96,7 @@ def _fused_qk_norm_rope_store_kernel( BYTES_PER_TOKEN: tl.constexpr, SWA_PAGE_SIZE: tl.constexpr, BF16_STORE: tl.constexpr, + IS_FNUZ: tl.constexpr, ): pid_m = tl.program_id(0).to(tl.int64) pid_h = tl.program_id(1).to(tl.int64) @@ -259,7 +260,14 @@ def _fused_qk_norm_rope_store_kernel( x_scaled = tile_data * inv_scale[:, None] x_fp8 = tl.clamp(x_scaled, FP8_MIN, FP8_MAX) - x_fp8_cast = x_fp8.to(tl.float8e4nv) + # Encode with the SAME fp8 type the decode bitcasts (read + # _KV_FP8_TY = float8e4b8 for e4m3fnuz / float8e4nv otherwise). + # An implicit/fn cast mis-encodes under fnuz (exponent bias 7 vs + # fnuz bias 8), so the fnuz read decodes every element 2x too small. + if IS_FNUZ: + x_fp8_cast = x_fp8.to(tl.float8e4b8) + else: + x_fp8_cast = x_fp8.to(tl.float8e4nv) x_fp8_bytes = x_fp8_cast.to(tl.uint8, bitcast=True) fp8_byte_offs = value_base[:, None] + tile_start + nope_tile_offs[None, :] tl.store( @@ -396,6 +404,7 @@ def fused_qk_norm_rope_swa_store( BYTES_PER_TOKEN=bytes_per_token, SWA_PAGE_SIZE=swa_page_size, BF16_STORE=bf16_store, + IS_FNUZ=_fp8_fnuz, num_warps=num_warps, ) return q_out diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index e73895f8e..7d3265bd9 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -129,6 +129,7 @@ from sglang.srt.utils import ( add_prefix, get_bool_env_var, is_gfx95_supported, + is_gfx942_supported, log_info_on_rank0, make_layers, ) @@ -162,6 +163,7 @@ _use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip # SGLANG_SHARED_EXPERT_TP1=1 (replicated shared expert). Default OFF. _SHARED_EXPERT_LOCAL = get_bool_env_var("SGLANG_DP_SHARED_EXPERT_LOCAL") _is_gfx95_supported = is_gfx95_supported() +_is_gfx942_supported = is_gfx942_supported() if _use_aiter: if _is_gfx95_supported: @@ -950,7 +952,14 @@ class MQALayer(nn.Module): # dispatch the cheaper decode::head64 variant; attn_sink is sliced to # this rank and padded to match. padded_num_heads = 64 if self.n_local_heads <= 64 else self.n_heads - q_padded = x.new_empty(x.shape[0], padded_num_heads, self.head_dim) + # Only [0:n_local_heads] is written below. Uninitialized padded TP + # heads inject NaN into attention on gfx942 (fnuz), so zero-init + # there; other archs tolerate new_empty and skip the per-forward + # memset. + if _is_gfx942_supported: + q_padded = x.new_zeros(x.shape[0], padded_num_heads, self.head_dim) + else: + q_padded = x.new_empty(x.shape[0], padded_num_heads, self.head_dim) tp_slice = slice(0, self.n_local_heads) q_out = q_padded[:, tp_slice, :] if self._attn_sink_local is None: