From 951fa05a09c7937bbb4cd4547c0674f4ecbece0f Mon Sep 17 00:00:00 2001 From: popsiclexu Date: Tue, 2 Jun 2026 11:40:38 +0800 Subject: [PATCH] [MoE] Support BF16 standard A2A with DeepGEMM runner (#26473) Co-authored-by: popsiclexu --- docker/Dockerfile | 2 +- python/pyproject.toml | 2 +- .../layers/deep_gemm_wrapper/compile_utils.py | 2 +- .../sglang/srt/layers/moe/ep_moe/kernels.py | 59 +++++++++++-------- .../srt/layers/moe/moe_runner/deep_gemm.py | 6 ++ 5 files changed, 45 insertions(+), 26 deletions(-) diff --git a/docker/Dockerfile b/docker/Dockerfile index fbe1b2b58..104d0f479 100644 --- a/docker/Dockerfile +++ b/docker/Dockerfile @@ -12,7 +12,7 @@ ARG DEEPEP_COMMIT=9af0e0d0e74f3577af1979c9b9e1ac2cad0104ee ARG BUILD_AND_DOWNLOAD_PARALLEL=8 ARG SGL_KERNEL_VERSION=0.4.3 ARG SGL_VERSION -ARG SGL_DEEP_GEMM_VERSION=0.1.0 +ARG SGL_DEEP_GEMM_VERSION=0.1.1 ARG USE_LATEST_SGLANG=0 ARG GDRCOPY_VERSION=2.5.1 ARG PIP_DEFAULT_INDEX diff --git a/python/pyproject.toml b/python/pyproject.toml index 278686fe9..be2f05b97 100755 --- a/python/pyproject.toml +++ b/python/pyproject.toml @@ -60,7 +60,7 @@ dependencies = [ "sentencepiece", "setproctitle", "flash-attn-4>=4.0.0b9", - "sgl-deep-gemm==0.1.0", + "sgl-deep-gemm==0.1.1", "sglang-kernel==0.4.3", "soundfile==0.13.1", "tiktoken", diff --git a/python/sglang/srt/layers/deep_gemm_wrapper/compile_utils.py b/python/sglang/srt/layers/deep_gemm_wrapper/compile_utils.py index 5a6a49bfb..e5c39f850 100644 --- a/python/sglang/srt/layers/deep_gemm_wrapper/compile_utils.py +++ b/python/sglang/srt/layers/deep_gemm_wrapper/compile_utils.py @@ -349,7 +349,7 @@ class _BF16GroupedContWarmupExecutor(_BaseWarmupExecutor): self.a[:m], self.b, self.out[:m], - m_indices=self.m_indices[:m], + self.m_indices[:m], ) diff --git a/python/sglang/srt/layers/moe/ep_moe/kernels.py b/python/sglang/srt/layers/moe/ep_moe/kernels.py index 7def543f8..39f46ccef 100644 --- a/python/sglang/srt/layers/moe/ep_moe/kernels.py +++ b/python/sglang/srt/layers/moe/ep_moe/kernels.py @@ -475,10 +475,12 @@ def _silu_and_mul_kernel( input_ptr_offs + token_index * stride_input_1 + size_n, mask=offs_in_d < size_n, other=0.0, - ) + ).to(tl.float32) gate = gate / (1 + tl.exp(-gate)) - gate = gate.to(input_ptr.dtype.element_ty) gate_up = up * gate + # Compute SiLU in fp32 for better precision, then cast back to the + # input dtype. + gate_up = gate_up.to(input_ptr.dtype.element_ty) tl.store( output_ptr_offs + token_index * stride_output_1, gate_up, @@ -702,17 +704,19 @@ def post_reorder_triton_kernel( offset = start_offset + vec mask = offset < hidden_size - sum_vec = tl.zeros([BLOCK_SIZE], dtype=InDtype) + sum_vec = tl.zeros([BLOCK_SIZE], dtype=tl.float32) for idx in range(topk): expert_id = tl.load(topk_ids_ptr + idx) - if expert_id > 0: + if expert_id >= 0: dst_idx_int32 = tl.load(src2dst_ptr + idx) dst_idx = dst_idx_int32.to(tl.int64) - weigh_scale = tl.load(topk_weights_ptr + idx).to(InDtype) + weigh_scale = tl.load(topk_weights_ptr + idx).to(tl.float32) load_ptr = down_output_ptr + dst_idx * hidden_size - in_data = tl.load(load_ptr + offset, mask=mask) + # accumulate expert outputs in fp32 for better precision + # before casting to the final output dtype. + in_data = tl.load(load_ptr + offset, mask=mask).to(tl.float32) sum_vec += in_data * weigh_scale - tl.store(store_ptr + offset, sum_vec, mask=mask) + tl.store(store_ptr + offset, sum_vec.to(InDtype), mask=mask) @triton.jit @@ -1136,6 +1140,7 @@ def fill_gateup_input_triton_kernel( hidden_size, scale_size, BLOCK_SIZE: tl.constexpr, + IS_FP8: tl.constexpr, ): src_idx_int32 = tl.program_id(0) @@ -1143,7 +1148,8 @@ def fill_gateup_input_triton_kernel( src2dst_ptr = src2dst_ptr + src_idx * topk topk_ids_ptr = topk_ids_ptr + src_idx * topk src_ptr = input_ptr + src_idx * hidden_size - scale_src_ptr = scale_ptr + src_idx * scale_size + if IS_FP8: + scale_src_ptr = scale_ptr + src_idx * scale_size vec = tl.arange(0, BLOCK_SIZE) for idx in range(topk): @@ -1157,12 +1163,14 @@ def fill_gateup_input_triton_kernel( mask = offset < hidden_size in_data = tl.load(src_ptr + offset, mask=mask) tl.store(dst_ptr + offset, in_data, mask=mask) - scale_dst_ptr = gateup_input_scale_ptr + dst_idx * scale_size - for start_offset in tl.range(0, scale_size, BLOCK_SIZE): - offset = start_offset + vec - mask = offset < scale_size - in_scale = tl.load(scale_src_ptr + offset, mask=mask) - tl.store(scale_dst_ptr + offset, in_scale, mask=mask) + + if IS_FP8: + scale_dst_ptr = gateup_input_scale_ptr + dst_idx * scale_size + for start_offset in tl.range(0, scale_size, BLOCK_SIZE): + offset = start_offset + vec + mask = offset < scale_size + in_scale = tl.load(scale_src_ptr + offset, mask=mask) + tl.store(scale_dst_ptr + offset, in_scale, mask=mask) def moe_ep_deepgemm_preprocess( @@ -1210,15 +1218,19 @@ def moe_ep_deepgemm_preprocess( block_shape = [128, 128] assert len(block_shape) == 2 block_n, block_k = block_shape[0], block_shape[1] + is_fp8 = output_dtype == torch.float8_e4m3fn + if is_fp8: + # TODO: fuse this with the preprocess + hidden_states, scale = per_token_group_quant_fp8(hidden_states, block_k) - # TODO: fuse this with the preprocess - hidden_states, scale = per_token_group_quant_fp8(hidden_states, block_k) - - gateup_input_scale = torch.empty( - (gateup_input.size(0), gateup_input.size(1), scale.size(1)), - device=hidden_states.device, - dtype=scale.dtype, - ) + gateup_input_scale = torch.empty( + (gateup_input.size(0), gateup_input.size(1), scale.size(1)), + device=hidden_states.device, + dtype=scale.dtype, + ) + else: + scale = None + gateup_input_scale = None fill_gateup_input_triton_kernel[(hidden_states.shape[0],)]( hidden_states, @@ -1229,8 +1241,9 @@ def moe_ep_deepgemm_preprocess( topk_ids, top_k, hidden_states.size(1), - scale.size(1), + scale.size(1) if is_fp8 else 0, BLOCK_SIZE=1024, + IS_FP8=is_fp8, ) return ( diff --git a/python/sglang/srt/layers/moe/moe_runner/deep_gemm.py b/python/sglang/srt/layers/moe/moe_runner/deep_gemm.py index bad52b959..c3ef55650 100644 --- a/python/sglang/srt/layers/moe/moe_runner/deep_gemm.py +++ b/python/sglang/srt/layers/moe/moe_runner/deep_gemm.py @@ -585,6 +585,11 @@ def pre_permute_standard_to_deep_gemm( topk_weights, topk_ids = topk_weights, topk_ids # PreReorder + output_dtype = ( + torch.bfloat16 + if quant_info.w13_weight.dtype == torch.bfloat16 + else torch.float8_e4m3fn + ) masked_m, expected_m, src2dst, hidden_states, hidden_states_scale = ( moe_ep_deepgemm_preprocess( topk_ids, @@ -592,6 +597,7 @@ def pre_permute_standard_to_deep_gemm( hidden_states, runner_config.top_k, quant_info.block_shape, + output_dtype=output_dtype, ) )