From 3136c7331454291b4d762d61c9a20854cd6ebaae Mon Sep 17 00:00:00 2001 From: niushengxiao Date: Fri, 17 Jul 2026 16:57:37 +0800 Subject: [PATCH 1/3] fix: reduce memory occupation and raise performance for multi-nodes --- .dockerignore | 30 ++ docker/Dockerfile | 6 +- lightllm/__init__.py | 4 + .../fused_moe/fused_moe_weight.py | 8 + .../fused_moe/impl/deepgemm_impl.py | 107 ++----- .../fused_moe/deepep_scatter_gather.py | 203 +++++++++++++ .../fused_moe/grouped_fused_moe_ep.py | 277 ++++++++++++------ lightllm/common/quantization/__init__.py | 26 +- lightllm/distributed/communication_op.py | 29 +- .../layer_infer/transformer_layer_infer.py | 18 +- .../layer_infer/transformer_layer_infer.py | 18 +- 11 files changed, 530 insertions(+), 196 deletions(-) create mode 100644 .dockerignore diff --git a/.dockerignore b/.dockerignore new file mode 100644 index 0000000000..1ac2bb0d48 --- /dev/null +++ b/.dockerignore @@ -0,0 +1,30 @@ +.git +.github +.conda +.venv +.idea +.vscode + +__pycache__ +*.py[cod] +.pytest_cache +.mypy_cache +.ruff_cache + +build +dist +*.egg-info +docs +test +unit_tests +benchmark +logs +tmp + +*.bin +*.ckpt +*.gguf +*.onnx +*.pt +*.pth +*.safetensors diff --git a/docker/Dockerfile b/docker/Dockerfile index 5acee29028..ff661fa1f3 100644 --- a/docker/Dockerfile +++ b/docker/Dockerfile @@ -7,9 +7,8 @@ ARG VLLM_VERSION=0.21.0 ARG NIXL_REF=v1.2.0 ARG FLASH_MLA_REF=47c35a7 ARG DEEPGEMM_REF=891d57b4db1071624b5c8fa0d1e51cb317fa709f -ARG DEEPEP_REF=099d5f2bad488b9c534ea785062b12f2e91d1d41 +ARG DEEPEP_REF=60d44037a702f651a6e18bd4aea65ed8409051c2 ARG DEEPEP_NCCL_VERSION=2.30.4 -ARG DEEPEP_NVSHMEM_VERSION=3.3.24 ARG TARGETPLATFORM ARG ENABLE_DEEPEP=1 ARG ENABLE_NIXL=1 @@ -92,8 +91,7 @@ RUN if [ "${ENABLE_DEEPEP}" = "1" ]; then \ set -e; \ ln -sf /usr/lib/x86_64-linux-gnu/libmlx5.so.1 /usr/lib/x86_64-linux-gnu/libmlx5.so; \ python -m pip install --upgrade --no-deps \ - "nvidia-nccl-cu13==${DEEPEP_NCCL_VERSION}" \ - "nvidia-nvshmem-cu13==${DEEPEP_NVSHMEM_VERSION}"; \ + "nvidia-nccl-cu13==${DEEPEP_NCCL_VERSION}"; \ cd /root && git clone https://github.com/deepseek-ai/DeepEP.git && cd DeepEP && git checkout ${DEEPEP_REF}; \ ln -sf /opt/conda/lib/python${PYTHON_VERSION}/site-packages/nvidia/nvshmem/lib/libnvshmem_host.so.3 /opt/conda/lib/python${PYTHON_VERSION}/site-packages/nvidia/nvshmem/lib/libnvshmem_host.so; \ ln -sf /opt/conda/lib/python${PYTHON_VERSION}/site-packages/nvidia/nccl/lib/libnccl.so.2 /opt/conda/lib/python${PYTHON_VERSION}/site-packages/nvidia/nccl/lib/libnccl.so; \ diff --git a/lightllm/__init__.py b/lightllm/__init__.py index e9ba6f3041..bc09ec5a17 100644 --- a/lightllm/__init__.py +++ b/lightllm/__init__.py @@ -2,3 +2,7 @@ if is_musa(): import torchada # noqa: F401 +else: + import torch + + torch._C._accelerator_setAllocatorSettings("expandable_segments:True") diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py index 26f2b338b7..63fdf32b68 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py @@ -222,20 +222,28 @@ def masked_group_gemm( def prefilled_group_gemm( self, num_recv_tokens_per_expert_list, + num_unaligned_recv_tokens_per_expert: torch.Tensor, + recv_src_metadata: torch.Tensor, recv_x: Tuple[torch.Tensor], recv_topk_idx: torch.Tensor, recv_topk_weights: torch.Tensor, hidden_dtype=torch.bfloat16, + workspace_index: int = 0, + workspace_count: int = 1, ): assert self.enable_ep_moe, "prefilled_group_gemm is only supported when enable_ep_moe is True" return self.fuse_moe_impl.prefilled_group_gemm( num_recv_tokens_per_expert_list=num_recv_tokens_per_expert_list, + num_unaligned_recv_tokens_per_expert=num_unaligned_recv_tokens_per_expert, + recv_src_metadata=recv_src_metadata, recv_x=recv_x, recv_topk_idx=recv_topk_idx, recv_topk_weights=recv_topk_weights, w13=self.w13, w2=self.w2, hidden_dtype=hidden_dtype, + workspace_index=workspace_index, + workspace_count=workspace_count, ) def low_latency_combine( diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py index 5da17c57e1..f419fdd8d2 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py @@ -2,7 +2,6 @@ from typing import Optional, Tuple, Any from .triton_impl import FuseMoeTriton from lightllm.distributed import dist_group_manager -from lightllm.common.triton_utils.autotuner import Autotuner from lightllm.common.quantization.quantize_method import WeightPack from lightllm.utils.envs_utils import ( get_deepep_num_max_dispatch_tokens_per_rank_prefill, @@ -12,15 +11,10 @@ fused_experts, get_ep_num_sms, masked_group_gemm, - deepgemm_grouped_fp8_nt_contiguous, + get_prefill_moe_workspace, + expanded_moe_chunked_reduce, quantize_fused_experts_input, ) -from lightllm.common.basemodel.triton_kernel.quantization.fp8act_quant_kernel import ( - per_token_group_quant_fp8, - tma_align_input_scale, -) -from lightllm.common.basemodel.triton_kernel.fused_moe.deepep_scatter_gather import ep_scatter, ep_gather -from lightllm.common.basemodel.triton_kernel.fused_moe.moe_silu_and_mul import silu_and_mul_fwd from lightllm.common.basemodel.triton_kernel.redundancy_topk_ids_repair import redundancy_topk_ids_repair @@ -179,6 +173,8 @@ def dispatch( allocate_on_comm_stream=True, do_cpu_sync=True, do_handle_copy=False, + do_expand=True, + use_tma_aligned_col_major_sf=True, ) def hook(): @@ -211,87 +207,35 @@ def masked_group_gemm( def prefilled_group_gemm( self, num_recv_tokens_per_expert_list, + num_unaligned_recv_tokens_per_expert: torch.Tensor, + recv_src_metadata: torch.Tensor, recv_x: Tuple[torch.Tensor], recv_topk_idx: torch.Tensor, recv_topk_weights: torch.Tensor, w13: WeightPack, w2: WeightPack, hidden_dtype=torch.bfloat16, + workspace_index: int = 0, + workspace_count: int = 1, ): - device = recv_x[0].device w13_weight, w13_scale = w13.weight, w13.weight_scale w2_weight, w2_scale = w2.weight, w2.weight_scale - _, K = recv_x[0].shape - _, N, _ = w13_weight.shape - block_size = self.quant_method.block_size - # scatter - all_tokens = sum(num_recv_tokens_per_expert_list) # calcu padding all nums. - # gather_out shape [recive_num_tokens, hidden] - gather_out = torch.empty_like(recv_x[0], device=device, dtype=hidden_dtype) - if all_tokens > 0: - input_tensor = [ - torch.empty((all_tokens, K), device=device, dtype=recv_x[0].dtype), - torch.empty((all_tokens, K // 128), device=device, dtype=torch.float32), - ] - # when m_indices is filled ok. - # m_indices show token use which expert, example, [0, 0, 0, 0, .... 1, 1, 1, 1,...., cur_expert_num - 1, ..] - # the count of 0 is num_recv_tokens_per_expert_list[0], the count of 1 is num_recv_tokens_per_expert_list[1] - # ... - m_indices = torch.empty(all_tokens, device=device, dtype=torch.int32) - # output_index shape [recive_num_tokens, topk_num] - # output_index use to show the token index in input_tensor - output_index = torch.empty_like(recv_topk_idx) - - num_recv_tokens_per_expert = torch.tensor( - num_recv_tokens_per_expert_list, dtype=torch.int32, pin_memory=True, device="cpu" - ).cuda(non_blocking=True) - - expert_start_loc = torch.empty_like(num_recv_tokens_per_expert) - - ep_scatter( - recv_x[0], - recv_x[1], - recv_topk_idx, - num_recv_tokens_per_expert, - expert_start_loc, - input_tensor[0], - input_tensor[1], - m_indices, - output_index, - ) - input_tensor[1] = tma_align_input_scale(input_tensor[1]) - # groupgemm (contiguous layout) - gemm_out_a = torch.empty((all_tokens, N), device=device, dtype=hidden_dtype) - - deepgemm_grouped_fp8_nt_contiguous(input_tensor, (w13_weight, w13_scale), gemm_out_a, m_indices) - - # silu_and_mul_fwd + qaunt - # TODO fused kernel - silu_out = torch.empty((all_tokens, N // 2), device=device, dtype=hidden_dtype) - - silu_and_mul_fwd(gemm_out_a.view(-1, N), silu_out) - qsilu_out, qsilu_out_scale = per_token_group_quant_fp8( - silu_out, block_size, dtype=w13_weight.dtype, column_major_scales=True, scale_tma_aligned=True - ) - - # groupgemm (contiguous layout) - gemm_out_b = torch.empty((all_tokens, K), device=device, dtype=hidden_dtype) - - deepgemm_grouped_fp8_nt_contiguous( - (qsilu_out, qsilu_out_scale), (w2_weight, w2_scale), gemm_out_b, m_indices - ) - # gather and local reduce - ep_gather(gemm_out_b, recv_topk_idx, recv_topk_weights, output_index, gather_out) - else: - ######################################## warning ################################################## - # here is used to match autotune feature, make moe model run same triton kernel in different rank. - # in some special case, one rank will recv 0 token, so add a token to make it run triton kernel. - if Autotuner.is_autotune_warmup(): - _gemm_out_a = torch.zeros((1, N), device=device, dtype=hidden_dtype) - _silu_out = torch.zeros((1, N // 2), device=device, dtype=hidden_dtype) - silu_and_mul_fwd(_gemm_out_a.view(-1, N), _silu_out) - _gemm_out_a, _silu_out = None, None - + assert recv_topk_idx is None + gather_out = expanded_moe_chunked_reduce( + num_recv_tokens_per_expert_list, + num_unaligned_recv_tokens_per_expert, + recv_x, + recv_topk_weights, + recv_src_metadata, + w13_weight, + w13_scale, + w2_weight, + w2_scale, + self.quant_method.block_size, + get_prefill_moe_workspace(workspace_index, workspace_count), + hidden_dtype, + ) + del recv_x return gather_out def low_latency_combine( @@ -312,7 +256,8 @@ def combine( handle: Any, overlap_event: Optional[Any] = None, ): - # normal combine + # The prefill kernel keeps expanded routing metadata while pointing its + # single valid slot at each pre-reduced dense row. combined_x, _, event = dist_group_manager.ep_buffer.combine( gemm_out_b, handle, diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/deepep_scatter_gather.py b/lightllm/common/basemodel/triton_kernel/fused_moe/deepep_scatter_gather.py index 101d316937..d37f3ee039 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/deepep_scatter_gather.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/deepep_scatter_gather.py @@ -152,6 +152,209 @@ def ep_scatter( return +@torch.no_grad() +def ep_fill_m_indices( + num_recv_tokens_per_expert: torch.Tensor, + m_indices: torch.Tensor, +): + """Build DeepGEMM's contiguous expert index vector without scattering data.""" + block_e = 128 + num_experts = num_recv_tokens_per_expert.shape[0] + assert m_indices.shape[0] % block_e == 0 + + expert_start_loc = torch.empty_like(num_recv_tokens_per_expert) + _fwd_kernel_ep_scatter_1[(num_experts,)]( + num_recv_tokens_per_expert, + expert_start_loc, + m_indices, + num_experts=num_experts, + num_warps=8, + BLOCK_E=block_e, + BLOCK_EXPERT_NUM=triton.next_power_of_2(num_experts), + ) + return expert_start_loc + + +@triton.jit +def _zero_expanded_padding_kernel( + recv_x, + recv_x_stride_m, + recv_x_stride_k, + recv_x_scale, + recv_x_scale_stride_m, + recv_x_scale_stride_k, + recv_topk_weights, + num_recv_tokens_per_expert, + num_unaligned_recv_tokens_per_expert, + expert_start_loc, + hidden_size: tl.constexpr, + scale_hidden_size: tl.constexpr, + BLOCK_M: tl.constexpr, + BLOCK_K: tl.constexpr, + BLOCK_SCALE_K: tl.constexpr, +): + expert_id = tl.program_id(0) + pad_block_id = tl.program_id(1) + hidden_block_id = tl.program_id(2) + expert_start = tl.load(expert_start_loc + expert_id) + aligned_count = tl.load(num_recv_tokens_per_expert + expert_id) + actual_count = tl.load(num_unaligned_recv_tokens_per_expert + expert_id) + pad_offsets = pad_block_id * BLOCK_M + tl.arange(0, BLOCK_M) + row_offsets = (expert_start + actual_count + pad_offsets).to(tl.int64) + row_mask = pad_offsets < aligned_count - actual_count + + hidden_offsets = hidden_block_id * BLOCK_K + tl.arange(0, BLOCK_K) + x_ptrs = recv_x + row_offsets[:, None] * recv_x_stride_m + hidden_offsets[None, :] * recv_x_stride_k + tl.store(x_ptrs, 0.0, mask=row_mask[:, None] & (hidden_offsets[None, :] < hidden_size)) + if hidden_block_id == 0: + scale_offsets = tl.arange(0, BLOCK_SCALE_K) + scale_ptrs = ( + recv_x_scale + row_offsets[:, None] * recv_x_scale_stride_m + scale_offsets[None, :] * recv_x_scale_stride_k + ) + tl.store(scale_ptrs, 0.0, mask=row_mask[:, None] & (scale_offsets[None, :] < scale_hidden_size)) + tl.store(recv_topk_weights + row_offsets, 0.0, mask=row_mask) + + +@torch.no_grad() +def ep_zero_expanded_padding( + recv_x: torch.Tensor, + recv_x_scale: torch.Tensor, + recv_topk_weights: torch.Tensor, + num_recv_tokens_per_expert: torch.Tensor, + num_unaligned_recv_tokens_per_expert: torch.Tensor, + expert_start_loc: torch.Tensor, +): + block_m = 8 + block_k = 256 + scale_hidden_size = recv_x_scale.shape[1] + grid = ( + num_recv_tokens_per_expert.shape[0], + triton.cdiv(127, block_m), + triton.cdiv(recv_x.shape[1], block_k), + ) + _zero_expanded_padding_kernel[grid]( + recv_x, + recv_x.stride(0), + recv_x.stride(1), + recv_x_scale, + recv_x_scale.stride(0), + recv_x_scale.stride(1), + recv_topk_weights, + num_recv_tokens_per_expert, + num_unaligned_recv_tokens_per_expert, + expert_start_loc, + hidden_size=recv_x.shape[1], + scale_hidden_size=scale_hidden_size, + BLOCK_M=block_m, + BLOCK_K=block_k, + BLOCK_SCALE_K=triton.next_power_of_2(scale_hidden_size), + num_warps=4, + ) + + +@triton.jit +def _accumulate_expanded_chunk_kernel( + total_recv_tokens, + chunk, + chunk_stride_m, + chunk_stride_k, + chunk_start, + chunk_end, + weights, + recv_src_metadata, + metadata_stride_m, + metadata_stride_k, + output, + output_stride_m, + output_stride_k, + TOPK: tl.constexpr, + BLOCK_D: tl.constexpr, +): + hidden_block_id = tl.program_id(0) + start_recv_token_id = tl.program_id(1) + recv_token_grid_size = tl.num_programs(1) + hidden_offsets = hidden_block_id * BLOCK_D + tl.arange(0, BLOCK_D) + + for recv_token_id in range(start_recv_token_id, total_recv_tokens, recv_token_grid_size): + output_ptrs = output + recv_token_id * output_stride_m + hidden_offsets * output_stride_k + accumulator = tl.load(output_ptrs).to(tl.float32) + for topk_id in range(TOPK): + slot = tl.load(recv_src_metadata + recv_token_id * metadata_stride_m + (topk_id + 2) * metadata_stride_k) + if slot >= chunk_start and slot < chunk_end: + local_row = (slot - chunk_start).to(tl.int64) + value = tl.load(chunk + local_row * chunk_stride_m + hidden_offsets * chunk_stride_k) + weight = tl.load(weights + slot) + accumulator += value.to(tl.float32) * weight + tl.store(output_ptrs, accumulator) + + +@torch.no_grad() +def ep_accumulate_expanded_chunk( + chunk: torch.Tensor, + chunk_start: int, + weights: torch.Tensor, + recv_src_metadata: torch.Tensor, + output: torch.Tensor, +): + """Accumulate one contiguous expanded W2 chunk into dense receive-token rows.""" + topk = recv_src_metadata.shape[1] - 2 + block_d = 1024 + assert chunk.shape[1] == output.shape[1] and output.shape[1] % block_d == 0 + grid = (triton.cdiv(output.shape[1], block_d), min(output.shape[0], 1024)) + _accumulate_expanded_chunk_kernel[grid]( + output.shape[0], + chunk, + chunk.stride(0), + chunk.stride(1), + chunk_start, + chunk_start + chunk.shape[0], + weights, + recv_src_metadata, + recv_src_metadata.stride(0), + recv_src_metadata.stride(1), + output, + output.stride(0), + output.stride(1), + TOPK=topk, + BLOCK_D=block_d, + num_warps=2, + ) + + +@triton.jit +def _compact_expanded_metadata_kernel( + recv_src_metadata, + metadata_stride_m, + metadata_stride_k, + TOPK: tl.constexpr, + BLOCK_TOPK: tl.constexpr, +): + recv_token_id = tl.program_id(0) + topk_id = tl.arange(0, BLOCK_TOPK) + slot = tl.where(topk_id == 0, recv_token_id, -1) + tl.store( + recv_src_metadata + recv_token_id * metadata_stride_m + (topk_id + 2) * metadata_stride_k, + slot, + mask=topk_id < TOPK, + ) + + +@torch.no_grad() +def ep_compact_expanded_metadata(recv_src_metadata: torch.Tensor): + """Point expanded combine metadata at pre-reduced dense token rows.""" + topk = recv_src_metadata.shape[1] - 2 + if recv_src_metadata.shape[0] == 0: + return + _compact_expanded_metadata_kernel[(recv_src_metadata.shape[0],)]( + recv_src_metadata, + recv_src_metadata.stride(0), + recv_src_metadata.stride(1), + TOPK=topk, + BLOCK_TOPK=triton.next_power_of_2(topk), + num_warps=1, + ) + + @triton.jit def _fwd_kernel_ep_gather( total_token_num, diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py index cb2e370cb9..48c0f00582 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py @@ -1,7 +1,8 @@ """Fused MoE kernel.""" import torch import triton -from typing import Any, Callable, Dict, Optional, Tuple +import triton.language as tl +from typing import Any, Callable, Dict, List, Optional, Tuple from lightllm.distributed import dist_group_manager from lightllm.utils.log_utils import init_logger from lightllm.common.basemodel.triton_kernel.fused_moe.moe_silu_and_mul import silu_and_mul_fwd @@ -10,9 +11,13 @@ ) from lightllm.common.basemodel.triton_kernel.quantization.fp8act_quant_kernel import ( per_token_group_quant_fp8, - tma_align_input_scale, ) -from lightllm.common.basemodel.triton_kernel.fused_moe.deepep_scatter_gather import ep_scatter, ep_gather +from lightllm.common.basemodel.triton_kernel.fused_moe.deepep_scatter_gather import ( + ep_accumulate_expanded_chunk, + ep_compact_expanded_metadata, + ep_fill_m_indices, + ep_zero_expanded_padding, +) from lightllm.utils.envs_utils import ( get_deepep_num_max_dispatch_tokens_per_rank_prefill, get_deepep_num_max_dispatch_tokens_per_rank_decode, @@ -75,12 +80,11 @@ def masked_group_gemm( expected_m = min(expected_m, padded_m) qsilu_out_scale = torch.empty((E, padded_m, N // 2 // block_size), device=recv_x[0].device, dtype=torch.float32) qsilu_out = torch.empty((E, padded_m, N // 2), dtype=w1.dtype, device=recv_x[0].device) - # groupgemm (masked layout) - gemm_out_b = torch.empty_like(recv_x[0], device=recv_x[0].device, dtype=dtype) - _deepgemm_grouped_fp8_nt_masked(recv_x, (w1, w1_scale), gemm_out_a, masked_m, expected_m) silu_and_mul_masked_post_quant_fwd(gemm_out_a, qsilu_out, qsilu_out_scale, block_size, masked_m) + del gemm_out_a + gemm_out_b = torch.empty_like(recv_x[0], device=recv_x[0].device, dtype=dtype) _deepgemm_grouped_fp8_nt_masked((qsilu_out, qsilu_out_scale), (w2, w2_scale), gemm_out_b, masked_m, expected_m) return gemm_out_b @@ -241,9 +245,6 @@ def fused_experts_impl( assert w2.is_contiguous(), "Expert weights2 must be contiguous" assert hidden_states.dtype in [torch.float32, torch.float16, torch.bfloat16] - M, K = hidden_states.shape - E, N, _ = w1.shape - # qaunt hidden_states assert use_fp8_w8a8 and use_fp8_all2all, "use_fp8_w8a8 and use_fp8_all2all must be True" @@ -258,11 +259,9 @@ def fused_experts_impl( if is_prefill: qinput_tensor, input_scale = per_token_group_quant_fp8(hidden_states, block_size_k, dtype=w1.dtype) allocate_on_comm_stream = previous_event is not None - # normal dispatch - # recv_x [recive_num_tokens, hidden] recv_x_scale [recive_num_tokens, hidden // block_size] - # recv_topk_idx [recive_num_tokens, topk_num] - # recv_topk_weights [recive_num_tokens, topk_num] - # num_recv_tokens_per_expert_list list [cur_node_expert_num] padding with expert_alignment=128 + # Expanded dispatch directly produces expert-contiguous FP8 input and + # TMA-aligned scales for DeepGEMM. DeepEP also keeps the metadata needed + # to reduce the expanded W2 output in combine. recv_x, recv_topk_idx, recv_topk_weights, handle, _ = buffer.dispatch( (qinput_tensor, input_scale), topk_idx=topk_idx, @@ -274,76 +273,32 @@ def fused_experts_impl( allocate_on_comm_stream=allocate_on_comm_stream, do_cpu_sync=True, do_handle_copy=False, + do_expand=True, + use_tma_aligned_col_major_sf=True, + ) + # Dispatch is synchronous in this path. Its FP8 source is no longer + # needed once the received tensors have been produced. + del qinput_tensor, input_scale + + assert recv_topk_idx is None + gather_out = expanded_moe_chunked_reduce( + handle.num_recv_tokens_per_expert_list, + handle.num_unaligned_recv_tokens_per_expert, + recv_x, + recv_topk_weights, + handle.recv_src_metadata, + w1, + w1_scale, + w2, + w2_scale, + block_size_k, + get_prefill_moe_workspace(), + hidden_states.dtype, ) + del recv_x - # scatter - all_tokens = sum(handle.num_recv_tokens_per_expert_list) # calcu padding all nums. - # gather_out shape [recive_num_tokens, hidden] - gather_out = torch.empty_like(recv_x[0], device=hidden_states.device, dtype=hidden_states.dtype) - if all_tokens > 0: - input_tensor = [ - torch.empty((all_tokens, K), device=hidden_states.device, dtype=qinput_tensor.dtype), - torch.empty((all_tokens, K // 128), device=hidden_states.device, dtype=torch.float32), - ] - # when m_indices is filled ok. - # m_indices show token use which expert, example, [0, 0, 0, 0, .... 1, 1, 1, 1,...., cur_expert_num - 1, ..] - # the count of 0 is num_recv_tokens_per_expert_list[0], the count of 1 is num_recv_tokens_per_expert_list[1] - # ... - m_indices = torch.empty(all_tokens, device=hidden_states.device, dtype=torch.int32) - # output_index shape [recive_num_tokens, topk_num] - # output_index use to show the token index in input_tensor - output_index = torch.empty_like(recv_topk_idx) - - num_recv_tokens_per_expert = torch.tensor( - handle.num_recv_tokens_per_expert_list, dtype=torch.int32, pin_memory=True, device="cpu" - ).cuda(non_blocking=True) - - expert_start_loc = torch.empty_like(num_recv_tokens_per_expert) - - ep_scatter( - recv_x[0], - recv_x[1], - recv_topk_idx, - num_recv_tokens_per_expert, - expert_start_loc, - input_tensor[0], - input_tensor[1], - m_indices, - output_index, - ) - - # groupgemm (contiguous layout) - gemm_out_a = torch.empty((all_tokens, N), device=hidden_states.device, dtype=hidden_states.dtype) - input_tensor[1] = tma_align_input_scale(input_tensor[1]) - deepgemm_grouped_fp8_nt_contiguous(input_tensor, (w1, w1_scale), gemm_out_a, m_indices) - - # silu_and_mul_fwd + qaunt - # TODO fused kernel - silu_out = torch.empty((all_tokens, N // 2), device=hidden_states.device, dtype=hidden_states.dtype) - - silu_and_mul_fwd(gemm_out_a.view(-1, N), silu_out) - qsilu_out, qsilu_out_scale = per_token_group_quant_fp8( - silu_out, block_size_k, dtype=w1.dtype, column_major_scales=True, scale_tma_aligned=True - ) - - # groupgemm (contiguous layout) - gemm_out_b = torch.empty((all_tokens, K), device=hidden_states.device, dtype=hidden_states.dtype) - - deepgemm_grouped_fp8_nt_contiguous((qsilu_out, qsilu_out_scale), (w2, w2_scale), gemm_out_b, m_indices) - - # gather and local reduce - ep_gather(gemm_out_b, recv_topk_idx, recv_topk_weights, output_index, gather_out) - else: - ######################################## warning ################################################## - # here is used to match autotune feature, make moe model run same triton kernel in different rank. - # in some special case, one rank will recv 0 token, so add a token to make it run triton kernel. - if Autotuner.is_autotune_warmup(): - _gemm_out_a = torch.zeros((1, N), device=hidden_states.device, dtype=hidden_states.dtype) - _silu_out = torch.zeros((1, N // 2), device=hidden_states.device, dtype=hidden_states.dtype) - silu_and_mul_fwd(_gemm_out_a.view(-1, N), _silu_out) - _gemm_out_a, _silu_out = None, None - - # normal combine + # W2 chunks were reduced to the deduplicated receive-token layout. Keep + # the expanded handle for routing, but point its slots at the dense rows. combined_x, _, event = buffer.combine( gather_out, handle, @@ -387,6 +342,164 @@ def deepgemm_grouped_fp8_nt_contiguous( raise RuntimeError("deep_gemm does not provide grouped_gemm_fp8 NT contiguous GEMM kernel in this version") +def get_prefill_moe_workspace( + workspace_index: int = 0, + workspace_count: int = 1, +): + """Map prefill MoE temporaries onto the idle low-latency RDMA buffer. + + Prefill uses the ElasticBuffer while decode uses the legacy low-latency + buffer, so their communication phases are mutually exclusive. The model + clears the low-latency buffer after every prefill before decode can use it. + """ + + workspace = dist_group_manager.ep_prefill_workspace + assert 0 <= workspace_index < workspace_count + workspace_size = workspace.numel() // workspace_count + workspace = workspace.narrow(0, workspace_index * workspace_size, workspace_size) + return workspace + + +def expanded_moe_chunked_reduce( + num_recv_tokens_per_expert_list: List[int], + num_unaligned_recv_tokens_per_expert: torch.Tensor, + recv_x: Tuple[torch.Tensor, torch.Tensor], + recv_topk_weights: torch.Tensor, + recv_src_metadata: torch.Tensor, + w1: torch.Tensor, + w1_scale: torch.Tensor, + w2: torch.Tensor, + w2_scale: torch.Tensor, + block_size_k: int, + workspace: torch.Tensor, + hidden_dtype: torch.dtype, +): + """Run expanded W1/W2 in bounded chunks and reduce to dense rows.""" + all_tokens = sum(num_recv_tokens_per_expert_list) + assert all_tokens == recv_x[0].shape[0] + intermediate_twice = w1.shape[1] + intermediate_size = intermediate_twice // 2 + hidden_size = w2.shape[1] + if all_tokens == 0: + if Autotuner.is_autotune_warmup(): + gemm_out = torch.zeros((1, intermediate_twice), device=recv_x[0].device, dtype=hidden_dtype) + silu_out = torch.zeros((1, intermediate_size), device=recv_x[0].device, dtype=hidden_dtype) + silu_and_mul_fwd(gemm_out, silu_out) + return torch.empty((0, hidden_size), device=recv_x[0].device, dtype=hidden_dtype) + + m_indices = torch.empty(all_tokens, device=recv_x[0].device, dtype=torch.int32) + num_recv_tokens_per_expert = torch.tensor( + num_recv_tokens_per_expert_list, dtype=torch.int32, pin_memory=True, device="cpu" + ).cuda(non_blocking=True) + expert_start_loc = ep_fill_m_indices(num_recv_tokens_per_expert, m_indices) + ep_zero_expanded_padding( + recv_x[0], + recv_x[1], + recv_topk_weights, + num_recv_tokens_per_expert, + num_unaligned_recv_tokens_per_expert, + expert_start_loc, + ) + del num_recv_tokens_per_expert, expert_start_loc + gather_rows = recv_src_metadata.shape[0] + gather_bytes = gather_rows * hidden_size * hidden_dtype.itemsize + silu_row_bytes = intermediate_size * hidden_dtype.itemsize + gemm_a_row_bytes = intermediate_twice * hidden_dtype.itemsize + gemm_b_row_bytes = hidden_size * hidden_dtype.itemsize + quant_row_bytes = intermediate_size * w2.dtype.itemsize + scale_row_bytes = (intermediate_size // block_size_k) * torch.float32.itemsize + quant_with_scale_row_bytes = quant_row_bytes + scale_row_bytes + # The same region is reused in three non-overlapping phases: + # W1: [SwiGLU output][W1 output] + # quant: [SwiGLU output]...[FP8 output + TMA scales] + # W2: [W2 output]......[FP8 output + TMA scales] + # Keeping the quantized activation at the end lets W2 overwrite the old + # SwiGLU/W1 storage without allocating another tensor from the CUDA heap. + temp_row_bytes = max( + silu_row_bytes + gemm_a_row_bytes, + silu_row_bytes + quant_with_scale_row_bytes, + gemm_b_row_bytes + quant_with_scale_row_bytes, + ) + max_chunk_rows = ((workspace.numel() - gather_bytes) // temp_row_bytes // 128) * 128 + if max_chunk_rows <= 0: + raise RuntimeError( + f"DeepEP workspace cannot hold dense output: need {gather_bytes} bytes, have {workspace.numel()} bytes" + ) + + gather_out = workspace.narrow(0, 0, gather_bytes).view(hidden_dtype).view(gather_rows, hidden_size) + gather_out.zero_() + temp_offset = gather_bytes + + for chunk_start in range(0, all_tokens, max_chunk_rows): + chunk_end = min(chunk_start + max_chunk_rows, all_tokens) + chunk_rows = chunk_end - chunk_start + silu_bytes = chunk_rows * silu_row_bytes + gemm_a_bytes = chunk_rows * gemm_a_row_bytes + gemm_b_bytes = chunk_rows * gemm_b_row_bytes + quant_bytes = chunk_rows * quant_row_bytes + aligned_chunk_rows = (chunk_rows + 3) // 4 * 4 + scale_storage_shape = (intermediate_size // block_size_k, aligned_chunk_rows) + scale_bytes = scale_storage_shape[0] * scale_storage_shape[1] * torch.float32.itemsize + temp_bytes = chunk_rows * temp_row_bytes + silu_out = workspace.narrow(0, temp_offset, silu_bytes).view(hidden_dtype).view(chunk_rows, intermediate_size) + gemm_region_offset = temp_offset + silu_bytes + gemm_out_a = ( + workspace.narrow(0, gemm_region_offset, gemm_a_bytes) + .view(hidden_dtype) + .view(chunk_rows, intermediate_twice) + ) + deepgemm_grouped_fp8_nt_contiguous( + (recv_x[0][chunk_start:chunk_end], recv_x[1][chunk_start:chunk_end]), + (w1, w1_scale), + gemm_out_a, + m_indices[chunk_start:chunk_end], + ) + silu_and_mul_fwd(gemm_out_a, silu_out) + del gemm_out_a + + quant_offset = temp_offset + temp_bytes - quant_bytes - scale_bytes + qsilu_workspace = ( + workspace.narrow(0, quant_offset, quant_bytes).view(w2.dtype).view(chunk_rows, intermediate_size) + ) + scale_workspace = ( + workspace.narrow(0, quant_offset + quant_bytes, scale_bytes).view(torch.float32).view(scale_storage_shape) + ) + + def workspace_quant_alloc(shape, dtype, device): + if tuple(shape) == tuple(qsilu_workspace.shape) and dtype == qsilu_workspace.dtype: + return qsilu_workspace + if tuple(shape) == scale_storage_shape and dtype == torch.float32: + return scale_workspace + raise RuntimeError(f"unexpected prefill quant allocation: shape={shape}, dtype={dtype}") + + qsilu_out, qsilu_out_scale = per_token_group_quant_fp8( + silu_out, + block_size_k, + dtype=w2.dtype, + column_major_scales=True, + scale_tma_aligned=True, + alloc_func=workspace_quant_alloc, + ) + gemm_out_b = workspace.narrow(0, temp_offset, gemm_b_bytes).view(hidden_dtype).view(chunk_rows, hidden_size) + deepgemm_grouped_fp8_nt_contiguous( + (qsilu_out, qsilu_out_scale), + (w2, w2_scale), + gemm_out_b, + m_indices[chunk_start:chunk_end], + ) + del qsilu_out, qsilu_out_scale, silu_out + ep_accumulate_expanded_chunk( + gemm_out_b, + chunk_start, + recv_topk_weights, + recv_src_metadata, + gather_out, + ) + + ep_compact_expanded_metadata(recv_src_metadata) + return gather_out + + def _deepgemm_grouped_fp8_nt_masked( input_tuple: Tuple[torch.Tensor, torch.Tensor], w_tuple: Tuple[torch.Tensor, torch.Tensor], diff --git a/lightllm/common/quantization/__init__.py b/lightllm/common/quantization/__init__.py index cd534d53ec..f676e35d64 100644 --- a/lightllm/common/quantization/__init__.py +++ b/lightllm/common/quantization/__init__.py @@ -43,6 +43,7 @@ def _parse_network_config(self, network_config): self.quantized_weight = False self.static_activation = False self.hf_quantization_config = None + self._mapping_expert_quant_method() return self.quantized_weight = True activation_scheme = network_config.get("activation_scheme", "dynamic") @@ -50,6 +51,19 @@ def _parse_network_config(self, network_config): self.hf_quantization_config = hf_quantization_config self.hf_quantization_method = hf_quantization_config["quant_method"] self._mapping_quant_method() + self._mapping_expert_quant_method() + + def _mapping_expert_quant_method(self): + expert_dtype = self.expert_dtype or self.network_config_.get("expert_dtype", None) + if expert_dtype is None: + return + target = self._get_expert_quant_type(expert_dtype) + for layer_num in range(self.layer_num): + if self.expert_dtype is not None: + self.quant_cfg[layer_num]["fused_moe"] = target + else: + self.quant_cfg[layer_num].setdefault("fused_moe", target) + logger.info(f"select fused_moe quant way from expert_dtype=`{expert_dtype}`: {target}") def _mapping_quant_method(self): if self.hf_quantization_method == "fp8": @@ -63,18 +77,6 @@ def _mapping_quant_method(self): self.quant_type = "vllm-fp8w8a8-b128" logger.info(f"select fp8w8a8-b128 quant way: {self.quant_type}") - # fp8 量化下,部分 MoE 模型(如 DeepSeek-V4),可以单独声明 expert 权重精度, - # 按其值给 fused_moe 选用对应的 deepgemm 量化方法。 - expert_dtype = self.expert_dtype or self.network_config_.get("expert_dtype", None) - if expert_dtype is None: - return - target = self._get_expert_quant_type(expert_dtype) - for layer_num in range(self.layer_num): - if self.expert_dtype is not None: - self.quant_cfg[layer_num]["fused_moe"] = target - else: - self.quant_cfg[layer_num].setdefault("fused_moe", target) - logger.info(f"select fused_moe quant way from expert_dtype=`{expert_dtype}`: {target}") elif self.hf_quantization_method == "awq": self.quant_type = "awq" if is_awq_marlin_compatible(self.hf_quantization_config): diff --git a/lightllm/distributed/communication_op.py b/lightllm/distributed/communication_op.py index f15badde25..83dd8932c5 100644 --- a/lightllm/distributed/communication_op.py +++ b/lightllm/distributed/communication_op.py @@ -109,6 +109,7 @@ def __init__(self): self.groups = [] self.ep_buffer = None self.ep_low_latency_buffer = None + self.ep_prefill_workspace = None self.ep_mega_moe_buffer = None self.ep_num_sms = None @@ -145,6 +146,7 @@ def new_deepep_group( if not enable_ep_moe: self.ep_buffer = None self.ep_low_latency_buffer = None + self.ep_prefill_workspace = None self.ep_mega_moe_buffer = None self.ep_num_sms = None return @@ -162,22 +164,23 @@ def new_deepep_group( hidden=self.ll_hidden, num_topk=num_experts_per_tok, use_fp8_dispatch=True, - allow_multiple_reduction=False, + allow_multiple_reduction=True, ) self.ep_mega_moe_buffer = None self.ep_low_latency_buffer = None - if not is_sm100_gpu(): - num_rdma_bytes = deep_ep.Buffer.get_low_latency_rdma_size_hint( - self.ll_decode_num_tokens, self.ll_hidden, global_world_size, self.ll_num_experts - ) - self.ep_low_latency_buffer = deep_ep.Buffer( - deepep_group, - int(1e9), - num_rdma_bytes, - low_latency_mode=True, - num_qps_per_rank=(self.ll_num_experts // global_world_size), - ) - else: + num_rdma_bytes = deep_ep.Buffer.get_low_latency_rdma_size_hint( + self.ll_decode_num_tokens, self.ll_hidden, global_world_size, self.ll_num_experts + ) + self.ep_low_latency_buffer = deep_ep.Buffer( + deepep_group, + num_rdma_bytes=num_rdma_bytes, + low_latency_mode=True, + num_qps_per_rank=(self.ll_num_experts // global_world_size), + ) + self.ep_prefill_workspace = self.ep_low_latency_buffer.get_local_buffer_tensor( + torch.uint8, use_rdma_buffer=True + ) + if is_sm100_gpu(): if moe_intermediate_size is None: raise ValueError("SM100 Mega MoE requires moe_intermediate_size or intermediate_size in model config") diff --git a/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py index be819c94a0..bd93cd39ab 100644 --- a/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py @@ -499,7 +499,14 @@ def overlap_tpsp_context_forward( # 0 moe calu _0_moe_out = layer_weight.experts.prefilled_group_gemm( - _0_num_recv_tokens_per_expert_list, _0_recv_x, _0_recv_topk_idx, _0_recv_topk_weight + _0_num_recv_tokens_per_expert_list, + _0_handle.num_unaligned_recv_tokens_per_expert, + _0_handle.recv_src_metadata, + _0_recv_x, + _0_recv_topk_idx, + _0_recv_topk_weight, + workspace_index=0, + workspace_count=2, ) # 1 dispatch execute @@ -525,7 +532,14 @@ def overlap_tpsp_context_forward( # 1 moe calc _1_moe_out = layer_weight.experts.prefilled_group_gemm( - _1_num_recv_tokens_per_expert_list, _1_recv_x, _1_recv_topk_idx, _1_recv_topk_weight + _1_num_recv_tokens_per_expert_list, + _1_handle.num_unaligned_recv_tokens_per_expert, + _1_handle.recv_src_metadata, + _1_recv_x, + _1_recv_topk_idx, + _1_recv_topk_weight, + workspace_index=1, + workspace_count=2, ) # wait 0 combine diff --git a/lightllm/models/qwen3_moe/layer_infer/transformer_layer_infer.py b/lightllm/models/qwen3_moe/layer_infer/transformer_layer_infer.py index 8879aa2d27..1698b850d1 100644 --- a/lightllm/models/qwen3_moe/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/qwen3_moe/layer_infer/transformer_layer_infer.py @@ -313,7 +313,14 @@ def overlap_tpsp_context_forward( # 0 moe calu _0_moe_out = layer_weight.experts.prefilled_group_gemm( - _0_num_recv_tokens_per_expert_list, _0_recv_x, _0_recv_topk_idx, _0_recv_topk_weight + _0_num_recv_tokens_per_expert_list, + _0_handle.num_unaligned_recv_tokens_per_expert, + _0_handle.recv_src_metadata, + _0_recv_x, + _0_recv_topk_idx, + _0_recv_topk_weight, + workspace_index=0, + workspace_count=2, ) # 1 dispatch execute @@ -339,7 +346,14 @@ def overlap_tpsp_context_forward( # 1 moe calc _1_moe_out = layer_weight.experts.prefilled_group_gemm( - _1_num_recv_tokens_per_expert_list, _1_recv_x, _1_recv_topk_idx, _1_recv_topk_weight + _1_num_recv_tokens_per_expert_list, + _1_handle.num_unaligned_recv_tokens_per_expert, + _1_handle.recv_src_metadata, + _1_recv_x, + _1_recv_topk_idx, + _1_recv_topk_weight, + workspace_index=1, + workspace_count=2, ) # wait 0 combine From af3c9390b248543244fbccbac6a88d17ac7548c8 Mon Sep 17 00:00:00 2001 From: niushengxiao Date: Mon, 20 Jul 2026 18:03:05 +0800 Subject: [PATCH 2/3] feat: refine code --- .../fused_moe/deepep_scatter_gather.py | 30 ++--- .../fused_moe/grouped_fused_moe_ep.py | 111 ++++++++---------- lightllm/distributed/communication_op.py | 10 +- 3 files changed, 73 insertions(+), 78 deletions(-) diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/deepep_scatter_gather.py b/lightllm/common/basemodel/triton_kernel/fused_moe/deepep_scatter_gather.py index d37f3ee039..9b292e43cf 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/deepep_scatter_gather.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/deepep_scatter_gather.py @@ -14,16 +14,19 @@ def _fwd_kernel_ep_scatter_1( num_experts: tl.constexpr, BLOCK_E: tl.constexpr, BLOCK_EXPERT_NUM: tl.constexpr, + ALIGN_COUNTS: tl.constexpr, ): cur_expert = tl.program_id(0) offset_cumsum = tl.arange(0, BLOCK_EXPERT_NUM) tokens_per_expert = tl.load(num_recv_tokens_per_expert + offset_cumsum, mask=offset_cumsum < num_experts, other=0) - cumsum = tl.cumsum(tokens_per_expert) - tokens_per_expert - tl.store(expert_start_loc + offset_cumsum, cumsum, mask=offset_cumsum < num_experts) - - cur_expert_start = tl.load(expert_start_loc + cur_expert) + if ALIGN_COUNTS: + tokens_per_expert = (tokens_per_expert + BLOCK_E - 1) // BLOCK_E * BLOCK_E + cur_expert_start = tl.sum(tl.where(offset_cumsum < cur_expert, tokens_per_expert, 0)) cur_expert_token_num = tl.load(num_recv_tokens_per_expert + cur_expert) + if ALIGN_COUNTS: + cur_expert_token_num = (cur_expert_token_num + BLOCK_E - 1) // BLOCK_E * BLOCK_E + tl.store(expert_start_loc + cur_expert, cur_expert_start) m_indices_start_ptr = m_indices + cur_expert_start off_expert = tl.arange(0, BLOCK_E) @@ -117,6 +120,7 @@ def ep_scatter( num_warps=num_warps, BLOCK_E=BLOCK_E, BLOCK_EXPERT_NUM=triton.next_power_of_2(num_experts), + ALIGN_COUNTS=False, ) grid = min(recv_topk.shape[0], 1024 * 8) @@ -154,23 +158,24 @@ def ep_scatter( @torch.no_grad() def ep_fill_m_indices( - num_recv_tokens_per_expert: torch.Tensor, + num_unaligned_recv_tokens_per_expert: torch.Tensor, m_indices: torch.Tensor, ): - """Build DeepGEMM's contiguous expert index vector without scattering data.""" + """Build aligned expert offsets and DeepGEMM's expert index vector.""" block_e = 128 - num_experts = num_recv_tokens_per_expert.shape[0] + num_experts = num_unaligned_recv_tokens_per_expert.shape[0] assert m_indices.shape[0] % block_e == 0 - expert_start_loc = torch.empty_like(num_recv_tokens_per_expert) + expert_start_loc = torch.empty_like(num_unaligned_recv_tokens_per_expert) _fwd_kernel_ep_scatter_1[(num_experts,)]( - num_recv_tokens_per_expert, + num_unaligned_recv_tokens_per_expert, expert_start_loc, m_indices, num_experts=num_experts, num_warps=8, BLOCK_E=block_e, BLOCK_EXPERT_NUM=triton.next_power_of_2(num_experts), + ALIGN_COUNTS=True, ) return expert_start_loc @@ -184,7 +189,6 @@ def _zero_expanded_padding_kernel( recv_x_scale_stride_m, recv_x_scale_stride_k, recv_topk_weights, - num_recv_tokens_per_expert, num_unaligned_recv_tokens_per_expert, expert_start_loc, hidden_size: tl.constexpr, @@ -197,8 +201,8 @@ def _zero_expanded_padding_kernel( pad_block_id = tl.program_id(1) hidden_block_id = tl.program_id(2) expert_start = tl.load(expert_start_loc + expert_id) - aligned_count = tl.load(num_recv_tokens_per_expert + expert_id) actual_count = tl.load(num_unaligned_recv_tokens_per_expert + expert_id) + aligned_count = (actual_count + 127) // 128 * 128 pad_offsets = pad_block_id * BLOCK_M + tl.arange(0, BLOCK_M) row_offsets = (expert_start + actual_count + pad_offsets).to(tl.int64) row_mask = pad_offsets < aligned_count - actual_count @@ -220,7 +224,6 @@ def ep_zero_expanded_padding( recv_x: torch.Tensor, recv_x_scale: torch.Tensor, recv_topk_weights: torch.Tensor, - num_recv_tokens_per_expert: torch.Tensor, num_unaligned_recv_tokens_per_expert: torch.Tensor, expert_start_loc: torch.Tensor, ): @@ -228,7 +231,7 @@ def ep_zero_expanded_padding( block_k = 256 scale_hidden_size = recv_x_scale.shape[1] grid = ( - num_recv_tokens_per_expert.shape[0], + num_unaligned_recv_tokens_per_expert.shape[0], triton.cdiv(127, block_m), triton.cdiv(recv_x.shape[1], block_k), ) @@ -240,7 +243,6 @@ def ep_zero_expanded_padding( recv_x_scale.stride(0), recv_x_scale.stride(1), recv_topk_weights, - num_recv_tokens_per_expert, num_unaligned_recv_tokens_per_expert, expert_start_loc, hidden_size=recv_x.shape[1], diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py index 48c0f00582..db32769064 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py @@ -259,10 +259,20 @@ def fused_experts_impl( if is_prefill: qinput_tensor, input_scale = per_token_group_quant_fp8(hidden_states, block_size_k, dtype=w1.dtype) allocate_on_comm_stream = previous_event is not None - # Expanded dispatch directly produces expert-contiguous FP8 input and - # TMA-aligned scales for DeepGEMM. DeepEP also keeps the metadata needed - # to reduce the expanded W2 output in combine. - recv_x, recv_topk_idx, recv_topk_weights, handle, _ = buffer.dispatch( + # Expanded dispatch directly produces expert-contiguous, alignment-padded inputs: + # recv_x[0]: [num_expanded_tokens, hidden] + # recv_x[1]: [num_expanded_tokens, hidden // block_size_k], with a + # TMA-aligned column-major physical layout + # recv_topk_weights: [num_expanded_tokens] + # Here, num_expanded_tokens is the sum of each local expert's token count padded to expert_alignment. + # handle.num_recv_tokens_per_expert_list: a Python list of length num_local_experts; + # each value is the expert's token count padded to expert_alignment, and + # their sum is num_expanded_tokens + # handle.num_unaligned_recv_tokens_per_expert: [num_local_experts], the actual + # token counts before alignment padding + # handle.recv_src_metadata: [num_recv_tokens, topk + 2]; the last topk columns + # map each deduplicated receive token to rows in the expanded tensors + recv_x, _, recv_topk_weights, handle, _ = buffer.dispatch( (qinput_tensor, input_scale), topk_idx=topk_idx, topk_weights=topk_weights, @@ -280,7 +290,6 @@ def fused_experts_impl( # needed once the received tensors have been produced. del qinput_tensor, input_scale - assert recv_topk_idx is None gather_out = expanded_moe_chunked_reduce( handle.num_recv_tokens_per_expert_list, handle.num_unaligned_recv_tokens_per_expert, @@ -346,14 +355,7 @@ def get_prefill_moe_workspace( workspace_index: int = 0, workspace_count: int = 1, ): - """Map prefill MoE temporaries onto the idle low-latency RDMA buffer. - - Prefill uses the ElasticBuffer while decode uses the legacy low-latency - buffer, so their communication phases are mutually exclusive. The model - clears the low-latency buffer after every prefill before decode can use it. - """ - - workspace = dist_group_manager.ep_prefill_workspace + workspace = dist_group_manager.prefill_moe_workspace assert 0 <= workspace_index < workspace_count workspace_size = workspace.numel() // workspace_count workspace = workspace.narrow(0, workspace_index * workspace_size, workspace_size) @@ -374,12 +376,12 @@ def expanded_moe_chunked_reduce( workspace: torch.Tensor, hidden_dtype: torch.dtype, ): - """Run expanded W1/W2 in bounded chunks and reduce to dense rows.""" - all_tokens = sum(num_recv_tokens_per_expert_list) - assert all_tokens == recv_x[0].shape[0] - intermediate_twice = w1.shape[1] - intermediate_size = intermediate_twice // 2 - hidden_size = w2.shape[1] + """Run bounded expanded MoE and rewrite metadata for dense DeepEP combine.""" + alignment = 128 + all_tokens, intermediate_twice = recv_x[0].shape[0], w1.shape[1] + intermediate_size, hidden_size = intermediate_twice // 2, w2.shape[1] + assert all_tokens == sum(num_recv_tokens_per_expert_list) and all_tokens % alignment == 0 + assert workspace.dtype == torch.uint8 and workspace.ndim == 1 and workspace.is_contiguous() if all_tokens == 0: if Autotuner.is_autotune_warmup(): gemm_out = torch.zeros((1, intermediate_twice), device=recv_x[0].device, dtype=hidden_dtype) @@ -388,45 +390,43 @@ def expanded_moe_chunked_reduce( return torch.empty((0, hidden_size), device=recv_x[0].device, dtype=hidden_dtype) m_indices = torch.empty(all_tokens, device=recv_x[0].device, dtype=torch.int32) - num_recv_tokens_per_expert = torch.tensor( - num_recv_tokens_per_expert_list, dtype=torch.int32, pin_memory=True, device="cpu" - ).cuda(non_blocking=True) - expert_start_loc = ep_fill_m_indices(num_recv_tokens_per_expert, m_indices) + expert_start_loc = ep_fill_m_indices(num_unaligned_recv_tokens_per_expert, m_indices) ep_zero_expanded_padding( recv_x[0], recv_x[1], recv_topk_weights, - num_recv_tokens_per_expert, num_unaligned_recv_tokens_per_expert, expert_start_loc, ) - del num_recv_tokens_per_expert, expert_start_loc + del expert_start_loc + gather_rows = recv_src_metadata.shape[0] gather_bytes = gather_rows * hidden_size * hidden_dtype.itemsize silu_row_bytes = intermediate_size * hidden_dtype.itemsize gemm_a_row_bytes = intermediate_twice * hidden_dtype.itemsize gemm_b_row_bytes = hidden_size * hidden_dtype.itemsize - quant_row_bytes = intermediate_size * w2.dtype.itemsize - scale_row_bytes = (intermediate_size // block_size_k) * torch.float32.itemsize - quant_with_scale_row_bytes = quant_row_bytes + scale_row_bytes + q_data_row_bytes = intermediate_size * w2.dtype.itemsize + scale_cols = intermediate_size // block_size_k + scale_row_bytes = scale_cols * torch.float32.itemsize # The same region is reused in three non-overlapping phases: # W1: [SwiGLU output][W1 output] # quant: [SwiGLU output]...[FP8 output + TMA scales] # W2: [W2 output]......[FP8 output + TMA scales] - # Keeping the quantized activation at the end lets W2 overwrite the old - # SwiGLU/W1 storage without allocating another tensor from the CUDA heap. + quant_row_bytes = q_data_row_bytes + scale_row_bytes temp_row_bytes = max( silu_row_bytes + gemm_a_row_bytes, - silu_row_bytes + quant_with_scale_row_bytes, - gemm_b_row_bytes + quant_with_scale_row_bytes, + silu_row_bytes + quant_row_bytes, + gemm_b_row_bytes + quant_row_bytes, ) - max_chunk_rows = ((workspace.numel() - gather_bytes) // temp_row_bytes // 128) * 128 + max_chunk_rows = (workspace.numel() - gather_bytes) // temp_row_bytes // alignment * alignment if max_chunk_rows <= 0: + minimum_bytes = gather_bytes + alignment * temp_row_bytes raise RuntimeError( - f"DeepEP workspace cannot hold dense output: need {gather_bytes} bytes, have {workspace.numel()} bytes" + f"DeepEP workspace needs at least {minimum_bytes} bytes " + f"({gather_bytes} dense + {alignment * temp_row_bytes} temporary), have {workspace.numel()} bytes" ) - gather_out = workspace.narrow(0, 0, gather_bytes).view(hidden_dtype).view(gather_rows, hidden_size) + gather_out = workspace[:gather_bytes].view(hidden_dtype).view(gather_rows, hidden_size) gather_out.zero_() temp_offset = gather_bytes @@ -436,18 +436,15 @@ def expanded_moe_chunked_reduce( silu_bytes = chunk_rows * silu_row_bytes gemm_a_bytes = chunk_rows * gemm_a_row_bytes gemm_b_bytes = chunk_rows * gemm_b_row_bytes - quant_bytes = chunk_rows * quant_row_bytes - aligned_chunk_rows = (chunk_rows + 3) // 4 * 4 - scale_storage_shape = (intermediate_size // block_size_k, aligned_chunk_rows) - scale_bytes = scale_storage_shape[0] * scale_storage_shape[1] * torch.float32.itemsize + q_data_bytes = chunk_rows * q_data_row_bytes + scale_storage_shape = (scale_cols, chunk_rows) + scale_bytes = chunk_rows * scale_row_bytes temp_bytes = chunk_rows * temp_row_bytes - silu_out = workspace.narrow(0, temp_offset, silu_bytes).view(hidden_dtype).view(chunk_rows, intermediate_size) - gemm_region_offset = temp_offset + silu_bytes - gemm_out_a = ( - workspace.narrow(0, gemm_region_offset, gemm_a_bytes) - .view(hidden_dtype) - .view(chunk_rows, intermediate_twice) + silu_out = ( + workspace[temp_offset : temp_offset + silu_bytes].view(hidden_dtype).view(chunk_rows, intermediate_size) ) + gemm_out_a = workspace[temp_offset + silu_bytes : temp_offset + silu_bytes + gemm_a_bytes] + gemm_out_a = gemm_out_a.view(hidden_dtype).view(chunk_rows, intermediate_twice) deepgemm_grouped_fp8_nt_contiguous( (recv_x[0][chunk_start:chunk_end], recv_x[1][chunk_start:chunk_end]), (w1, w1_scale), @@ -457,13 +454,11 @@ def expanded_moe_chunked_reduce( silu_and_mul_fwd(gemm_out_a, silu_out) del gemm_out_a - quant_offset = temp_offset + temp_bytes - quant_bytes - scale_bytes - qsilu_workspace = ( - workspace.narrow(0, quant_offset, quant_bytes).view(w2.dtype).view(chunk_rows, intermediate_size) - ) - scale_workspace = ( - workspace.narrow(0, quant_offset + quant_bytes, scale_bytes).view(torch.float32).view(scale_storage_shape) - ) + quant_offset = temp_offset + temp_bytes - q_data_bytes - scale_bytes + qsilu_workspace = workspace[quant_offset : quant_offset + q_data_bytes] + qsilu_workspace = qsilu_workspace.view(w2.dtype).view(chunk_rows, intermediate_size) + scale_workspace = workspace[quant_offset + q_data_bytes : quant_offset + q_data_bytes + scale_bytes] + scale_workspace = scale_workspace.view(torch.float32).view(scale_storage_shape) def workspace_quant_alloc(shape, dtype, device): if tuple(shape) == tuple(qsilu_workspace.shape) and dtype == qsilu_workspace.dtype: @@ -480,7 +475,9 @@ def workspace_quant_alloc(shape, dtype, device): scale_tma_aligned=True, alloc_func=workspace_quant_alloc, ) - gemm_out_b = workspace.narrow(0, temp_offset, gemm_b_bytes).view(hidden_dtype).view(chunk_rows, hidden_size) + gemm_out_b = ( + workspace[temp_offset : temp_offset + gemm_b_bytes].view(hidden_dtype).view(chunk_rows, hidden_size) + ) deepgemm_grouped_fp8_nt_contiguous( (qsilu_out, qsilu_out_scale), (w2, w2_scale), @@ -488,13 +485,7 @@ def workspace_quant_alloc(shape, dtype, device): m_indices[chunk_start:chunk_end], ) del qsilu_out, qsilu_out_scale, silu_out - ep_accumulate_expanded_chunk( - gemm_out_b, - chunk_start, - recv_topk_weights, - recv_src_metadata, - gather_out, - ) + ep_accumulate_expanded_chunk(gemm_out_b, chunk_start, recv_topk_weights, recv_src_metadata, gather_out) ep_compact_expanded_metadata(recv_src_metadata) return gather_out diff --git a/lightllm/distributed/communication_op.py b/lightllm/distributed/communication_op.py index 83dd8932c5..6f89be5bf3 100644 --- a/lightllm/distributed/communication_op.py +++ b/lightllm/distributed/communication_op.py @@ -109,7 +109,7 @@ def __init__(self): self.groups = [] self.ep_buffer = None self.ep_low_latency_buffer = None - self.ep_prefill_workspace = None + self.prefill_moe_workspace = None self.ep_mega_moe_buffer = None self.ep_num_sms = None @@ -146,7 +146,7 @@ def new_deepep_group( if not enable_ep_moe: self.ep_buffer = None self.ep_low_latency_buffer = None - self.ep_prefill_workspace = None + self.prefill_moe_workspace = None self.ep_mega_moe_buffer = None self.ep_num_sms = None return @@ -177,7 +177,8 @@ def new_deepep_group( low_latency_mode=True, num_qps_per_rank=(self.ll_num_experts // global_world_size), ) - self.ep_prefill_workspace = self.ep_low_latency_buffer.get_local_buffer_tensor( + # 当前rank的low-latency RDMA通信空间在prefill阶段处于空闲状态,将其复用为prefill MoE计算的临时工作区,降低峰值显存占用。 + self.prefill_moe_workspace = self.ep_low_latency_buffer.get_local_buffer_tensor( torch.uint8, use_rdma_buffer=True ) if is_sm100_gpu(): @@ -215,7 +216,8 @@ def _set_num_sms_for_deep_gemm(self, deepep_sms: int): def clear_deepep_buffer(self): """ - Prefill after using ElasticBuffer may leave the legacy low-latency buffer dirty for decode. + Prefill MoE compute reuses the low-latency RDMA buffer as workspace. + Clean it before the buffer is used by low-latency decode kernels. """ if self.ep_low_latency_buffer is not None: self.ep_low_latency_buffer.clean_low_latency_buffer( From 84d5277e26aa6447f74061647a09134a22bd6125 Mon Sep 17 00:00:00 2001 From: niushengxiao Date: Mon, 20 Jul 2026 19:32:58 +0800 Subject: [PATCH 3/3] feat: remove redundant code --- .../fused_moe/impl/deepgemm_impl.py | 7 +- .../deepep_expanded_layout_kernels.py | 264 +++++++++++ .../fused_moe/deepep_scatter_gather.py | 437 ------------------ .../fused_moe/grouped_fused_moe_ep.py | 49 +- unit_tests/common/fused_moe/test_deepep.py | 77 --- 5 files changed, 292 insertions(+), 542 deletions(-) create mode 100644 lightllm/common/basemodel/triton_kernel/fused_moe/deepep_expanded_layout_kernels.py delete mode 100644 lightllm/common/basemodel/triton_kernel/fused_moe/deepep_scatter_gather.py diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py index f419fdd8d2..4804f0ab25 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py @@ -12,7 +12,7 @@ get_ep_num_sms, masked_group_gemm, get_prefill_moe_workspace, - expanded_moe_chunked_reduce, + chunked_expanded_moe_forward, quantize_fused_experts_input, ) from lightllm.common.basemodel.triton_kernel.redundancy_topk_ids_repair import redundancy_topk_ids_repair @@ -221,7 +221,7 @@ def prefilled_group_gemm( w13_weight, w13_scale = w13.weight, w13.weight_scale w2_weight, w2_scale = w2.weight, w2.weight_scale assert recv_topk_idx is None - gather_out = expanded_moe_chunked_reduce( + gather_out = chunked_expanded_moe_forward( num_recv_tokens_per_expert_list, num_unaligned_recv_tokens_per_expert, recv_x, @@ -256,8 +256,7 @@ def combine( handle: Any, overlap_event: Optional[Any] = None, ): - # The prefill kernel keeps expanded routing metadata while pointing its - # single valid slot at each pre-reduced dense row. + # normal combine combined_x, _, event = dist_group_manager.ep_buffer.combine( gemm_out_b, handle, diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/deepep_expanded_layout_kernels.py b/lightllm/common/basemodel/triton_kernel/fused_moe/deepep_expanded_layout_kernels.py new file mode 100644 index 0000000000..ad9829dcea --- /dev/null +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/deepep_expanded_layout_kernels.py @@ -0,0 +1,264 @@ +import torch +import triton +import triton.language as tl + + +@triton.jit +def _ep_build_m_indices_kernel( + num_unaligned_recv_tokens_per_expert, + expert_start_loc, + m_indices, + num_experts: tl.constexpr, + BLOCK_E: tl.constexpr, + BLOCK_EXPERT_NUM: tl.constexpr, +): + cur_expert = tl.program_id(0) + + offset_cumsum = tl.arange(0, BLOCK_EXPERT_NUM) + tokens_per_expert = tl.load( + num_unaligned_recv_tokens_per_expert + offset_cumsum, + mask=offset_cumsum < num_experts, + other=0, + ) + tokens_per_expert = (tokens_per_expert + BLOCK_E - 1) // BLOCK_E * BLOCK_E + cur_expert_start = tl.sum(tl.where(offset_cumsum < cur_expert, tokens_per_expert, 0)) + cur_expert_token_num = tl.load(num_unaligned_recv_tokens_per_expert + cur_expert) + cur_expert_token_num = (cur_expert_token_num + BLOCK_E - 1) // BLOCK_E * BLOCK_E + tl.store(expert_start_loc + cur_expert, cur_expert_start) + + m_indices_start_ptr = m_indices + cur_expert_start + off_expert = tl.arange(0, BLOCK_E) + + for start_m in tl.range(0, cur_expert_token_num, BLOCK_E, num_stages=4): + tl.store( + m_indices_start_ptr + start_m + off_expert, + cur_expert, + ) + + +@torch.no_grad() +def ep_build_m_indices( + num_unaligned_recv_tokens_per_expert: torch.Tensor, # [num_local_experts] + m_indices: torch.Tensor, # [num_expanded_tokens] +): + """Build the 128-aligned expert layout used by contiguous grouped GEMM. + + Each expert's actual token count is rounded up to 128. ``m_indices`` is + filled in-place with the owning expert ID for every real and padding row. + + Returns: + ``expert_start_loc`` with shape ``[num_local_experts]``. Each value is + the expert's starting row in the expanded tensors. + """ + block_e = 128 + num_experts = num_unaligned_recv_tokens_per_expert.shape[0] + assert m_indices.shape[0] % block_e == 0 + + expert_start_loc = torch.empty_like(num_unaligned_recv_tokens_per_expert) + _ep_build_m_indices_kernel[(num_experts,)]( + num_unaligned_recv_tokens_per_expert, + expert_start_loc, + m_indices, + num_experts=num_experts, + num_warps=8, + BLOCK_E=block_e, + BLOCK_EXPERT_NUM=triton.next_power_of_2(num_experts), + ) + return expert_start_loc + + +@triton.jit +def _ep_zero_padding_kernel( + recv_x, + recv_x_stride_m, + recv_x_stride_k, + recv_x_scale, + recv_x_scale_stride_m, + recv_x_scale_stride_k, + recv_topk_weights, + num_unaligned_recv_tokens_per_expert, + expert_start_loc, + hidden_size: tl.constexpr, + scale_hidden_size: tl.constexpr, + BLOCK_M: tl.constexpr, + BLOCK_K: tl.constexpr, + BLOCK_SCALE_K: tl.constexpr, +): + expert_id = tl.program_id(0) + pad_block_id = tl.program_id(1) + hidden_block_id = tl.program_id(2) + expert_start = tl.load(expert_start_loc + expert_id) + actual_count = tl.load(num_unaligned_recv_tokens_per_expert + expert_id) + aligned_count = (actual_count + 127) // 128 * 128 + pad_offsets = pad_block_id * BLOCK_M + tl.arange(0, BLOCK_M) + row_offsets = (expert_start + actual_count + pad_offsets).to(tl.int64) + row_mask = pad_offsets < aligned_count - actual_count + + hidden_offsets = hidden_block_id * BLOCK_K + tl.arange(0, BLOCK_K) + x_ptrs = recv_x + row_offsets[:, None] * recv_x_stride_m + hidden_offsets[None, :] * recv_x_stride_k + tl.store(x_ptrs, 0.0, mask=row_mask[:, None] & (hidden_offsets[None, :] < hidden_size)) + if hidden_block_id == 0: + scale_offsets = tl.arange(0, BLOCK_SCALE_K) + scale_ptrs = ( + recv_x_scale + row_offsets[:, None] * recv_x_scale_stride_m + scale_offsets[None, :] * recv_x_scale_stride_k + ) + tl.store(scale_ptrs, 0.0, mask=row_mask[:, None] & (scale_offsets[None, :] < scale_hidden_size)) + tl.store(recv_topk_weights + row_offsets, 0.0, mask=row_mask) + + +@torch.no_grad() +def ep_zero_padding( + recv_x: torch.Tensor, # [num_expanded_tokens, hidden_size] + recv_x_scale: torch.Tensor, # [num_expanded_tokens, scale_hidden_size] + recv_topk_weights: torch.Tensor, # [num_expanded_tokens] + num_unaligned_recv_tokens_per_expert: torch.Tensor, # [num_local_experts] + expert_start_loc: torch.Tensor, # [num_local_experts] +): + """Zero the alignment-padding rows in DeepEP's expanded receive layout. + + For every expert, rows from its actual token count up to its 128-aligned + count are cleared in-place in the FP8 activations, activation scales, and + routing weights. ``recv_x_scale`` may use a column-major physical layout; + its logical shape remains ``[num_expanded_tokens, scale_hidden_size]``. + """ + block_m = 8 + block_k = 256 + scale_hidden_size = recv_x_scale.shape[1] + grid = ( + num_unaligned_recv_tokens_per_expert.shape[0], + triton.cdiv(127, block_m), + triton.cdiv(recv_x.shape[1], block_k), + ) + _ep_zero_padding_kernel[grid]( + recv_x, + recv_x.stride(0), + recv_x.stride(1), + recv_x_scale, + recv_x_scale.stride(0), + recv_x_scale.stride(1), + recv_topk_weights, + num_unaligned_recv_tokens_per_expert, + expert_start_loc, + hidden_size=recv_x.shape[1], + scale_hidden_size=scale_hidden_size, + BLOCK_M=block_m, + BLOCK_K=block_k, + BLOCK_SCALE_K=triton.next_power_of_2(scale_hidden_size), + num_warps=4, + ) + + +@triton.jit +def _ep_gather_chunk_kernel( + total_recv_tokens, + chunk, + chunk_stride_m, + chunk_stride_k, + chunk_start, + chunk_end, + weights, + recv_src_metadata, + metadata_stride_m, + metadata_stride_k, + output, + output_stride_m, + output_stride_k, + TOPK: tl.constexpr, + BLOCK_D: tl.constexpr, +): + hidden_block_id = tl.program_id(0) + start_recv_token_id = tl.program_id(1) + recv_token_grid_size = tl.num_programs(1) + hidden_offsets = hidden_block_id * BLOCK_D + tl.arange(0, BLOCK_D) + + for recv_token_id in range(start_recv_token_id, total_recv_tokens, recv_token_grid_size): + output_ptrs = output + recv_token_id * output_stride_m + hidden_offsets * output_stride_k + accumulator = tl.load(output_ptrs).to(tl.float32) + for topk_id in range(TOPK): + slot = tl.load(recv_src_metadata + recv_token_id * metadata_stride_m + (topk_id + 2) * metadata_stride_k) + if slot >= chunk_start and slot < chunk_end: + local_row = (slot - chunk_start).to(tl.int64) + value = tl.load(chunk + local_row * chunk_stride_m + hidden_offsets * chunk_stride_k) + weight = tl.load(weights + slot) + accumulator += value.to(tl.float32) * weight + tl.store(output_ptrs, accumulator) + + +@torch.no_grad() +def ep_gather_chunk( + chunk: torch.Tensor, # [chunk_rows, hidden_size] + chunk_start: int, # scalar expanded-row offset + weights: torch.Tensor, # [num_expanded_tokens] + recv_src_metadata: torch.Tensor, # [num_recv_tokens, topk + 2] + output: torch.Tensor, # [num_recv_tokens, hidden_size] +): + """Accumulate one expanded W2-output chunk into dense receive-token rows. + + The last ``topk`` columns of ``recv_src_metadata`` map each dense receive + token to global expanded-row IDs. Entries covered by this chunk are read, + multiplied by ``weights``, and accumulated in-place into ``output``. This + allows multiple chunks to contribute to the same dense output tensor. + """ + topk = recv_src_metadata.shape[1] - 2 + block_d = 1024 + assert chunk.shape[1] == output.shape[1] and output.shape[1] % block_d == 0 + grid = (triton.cdiv(output.shape[1], block_d), min(output.shape[0], 1024)) + _ep_gather_chunk_kernel[grid]( + output.shape[0], + chunk, + chunk.stride(0), + chunk.stride(1), + chunk_start, + chunk_start + chunk.shape[0], + weights, + recv_src_metadata, + recv_src_metadata.stride(0), + recv_src_metadata.stride(1), + output, + output.stride(0), + output.stride(1), + TOPK=topk, + BLOCK_D=block_d, + num_warps=2, + ) + + +@triton.jit +def _ep_compact_metadata_kernel( + recv_src_metadata, + metadata_stride_m, + metadata_stride_k, + TOPK: tl.constexpr, + BLOCK_TOPK: tl.constexpr, +): + recv_token_id = tl.program_id(0) + topk_id = tl.arange(0, BLOCK_TOPK) + slot = tl.where(topk_id == 0, recv_token_id, -1) + tl.store( + recv_src_metadata + recv_token_id * metadata_stride_m + (topk_id + 2) * metadata_stride_k, + slot, + mask=topk_id < TOPK, + ) + + +@torch.no_grad() +def ep_compact_metadata( + recv_src_metadata: torch.Tensor, # [num_recv_tokens, topk + 2] +): + """Rewrite expanded routing metadata for a pre-reduced dense tensor. + + The operation preserves the first two metadata columns and updates the + final ``topk`` columns in-place to ``[recv_token_id, -1, ...]``. DeepEP + combine can then read each already-reduced dense row exactly once. + """ + topk = recv_src_metadata.shape[1] - 2 + if recv_src_metadata.shape[0] == 0: + return + _ep_compact_metadata_kernel[(recv_src_metadata.shape[0],)]( + recv_src_metadata, + recv_src_metadata.stride(0), + recv_src_metadata.stride(1), + TOPK=topk, + BLOCK_TOPK=triton.next_power_of_2(topk), + num_warps=1, + ) diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/deepep_scatter_gather.py b/lightllm/common/basemodel/triton_kernel/fused_moe/deepep_scatter_gather.py deleted file mode 100644 index 9b292e43cf..0000000000 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/deepep_scatter_gather.py +++ /dev/null @@ -1,437 +0,0 @@ -import random -import torch -import torch.nn.functional as F -import triton -import triton.language as tl -from typing import Dict - - -@triton.jit -def _fwd_kernel_ep_scatter_1( - num_recv_tokens_per_expert, - expert_start_loc, - m_indices, - num_experts: tl.constexpr, - BLOCK_E: tl.constexpr, - BLOCK_EXPERT_NUM: tl.constexpr, - ALIGN_COUNTS: tl.constexpr, -): - cur_expert = tl.program_id(0) - - offset_cumsum = tl.arange(0, BLOCK_EXPERT_NUM) - tokens_per_expert = tl.load(num_recv_tokens_per_expert + offset_cumsum, mask=offset_cumsum < num_experts, other=0) - if ALIGN_COUNTS: - tokens_per_expert = (tokens_per_expert + BLOCK_E - 1) // BLOCK_E * BLOCK_E - cur_expert_start = tl.sum(tl.where(offset_cumsum < cur_expert, tokens_per_expert, 0)) - cur_expert_token_num = tl.load(num_recv_tokens_per_expert + cur_expert) - if ALIGN_COUNTS: - cur_expert_token_num = (cur_expert_token_num + BLOCK_E - 1) // BLOCK_E * BLOCK_E - tl.store(expert_start_loc + cur_expert, cur_expert_start) - - m_indices_start_ptr = m_indices + cur_expert_start - off_expert = tl.arange(0, BLOCK_E) - - for start_m in tl.range(0, cur_expert_token_num, BLOCK_E, num_stages=4): - tl.store( - m_indices_start_ptr + start_m + off_expert, - cur_expert, - ) - - -@triton.jit -def _fwd_kernel_ep_scatter_2( - total_token_num, - expert_start_loc, - recv_x, - recv_x_stride0, - recv_x_stride1, - recv_x_scale, - recv_x_scale_stride0, - recv_x_scale_stride1, - recv_topk, - recv_topk_stride0, - recv_topk_stride1, - output_tensor, - output_tensor_stride0, - output_tensor_stride1, - output_tensor_scale, - output_tensor_scale_stride0, - output_tensor_scale_stride1, - output_index, - output_index_stride0, - output_index_stride1, - topk_num: tl.constexpr, - HIDDEN_SIZE: tl.constexpr, - HIDDEN_SIZE_PAD: tl.constexpr, - SCALE_HIDDEN_SIZE: tl.constexpr, - SCALE_HIDDEN_SIZE_PAD: tl.constexpr, -): - start_token_id = tl.program_id(0) - grid_num = tl.num_programs(0) - - offset_in = tl.arange(0, HIDDEN_SIZE_PAD) - mask = offset_in < HIDDEN_SIZE - - offset_in_s = tl.arange(0, SCALE_HIDDEN_SIZE_PAD) - mask_s = offset_in_s < SCALE_HIDDEN_SIZE - for token_id in range(start_token_id, total_token_num, grid_num): - to_copy = tl.load(recv_x + token_id * recv_x_stride0 + offset_in, mask=mask) - to_copy_s = tl.load(recv_x_scale + token_id * recv_x_scale_stride0 + offset_in_s, mask=mask_s) - - for topk_index in tl.range(0, topk_num, 1, num_stages=4): - expert_id = tl.load(recv_topk + token_id * recv_topk_stride0 + topk_index) - if expert_id >= 0: - dest_token_index = tl.atomic_add(expert_start_loc + expert_id, 1) - dest_token_index = dest_token_index.to(tl.int64) - tl.store(output_index + token_id * output_index_stride0 + topk_index, dest_token_index) - output_tensor_ptr = output_tensor + dest_token_index * output_tensor_stride0 - output_tensor_scale_ptr = output_tensor_scale + dest_token_index * output_tensor_scale_stride0 - tl.store(output_tensor_ptr + offset_in, to_copy, mask=mask) - tl.store(output_tensor_scale_ptr + offset_in_s, to_copy_s, mask=mask_s) - - -@torch.no_grad() -def ep_scatter( - recv_x: torch.Tensor, - recv_x_scale: torch.Tensor, - recv_topk: torch.Tensor, - num_recv_tokens_per_expert: torch.Tensor, - expert_start_loc: torch.Tensor, - output_tensor: torch.Tensor, - output_tensor_scale: torch.Tensor, - m_indices: torch.Tensor, - output_index: torch.Tensor, -): - BLOCK_E = 128 # token num of per expert is aligned to 128 - BLOCK_D = 128 # block size of quantization - num_warps = 8 - num_experts = num_recv_tokens_per_expert.shape[0] # 获取num_recv_tokens_per_expert的元素个数 - hidden_size = recv_x.shape[1] - # grid = (triton.cdiv(hidden_size, BLOCK_D), num_experts) - grid = num_experts - - assert m_indices.shape[0] % BLOCK_E == 0 - - _fwd_kernel_ep_scatter_1[(grid,)]( - num_recv_tokens_per_expert, - expert_start_loc, - m_indices, - num_experts=num_experts, - num_warps=num_warps, - BLOCK_E=BLOCK_E, - BLOCK_EXPERT_NUM=triton.next_power_of_2(num_experts), - ALIGN_COUNTS=False, - ) - - grid = min(recv_topk.shape[0], 1024 * 8) - - _fwd_kernel_ep_scatter_2[(grid,)]( - recv_topk.shape[0], - expert_start_loc, - recv_x, - recv_x.stride(0), - recv_x.stride(1), - recv_x_scale, - recv_x_scale.stride(0), - recv_x_scale.stride(1), - recv_topk, - recv_topk.stride(0), - recv_topk.stride(1), - output_tensor, - output_tensor.stride(0), - output_tensor.stride(1), - output_tensor_scale, - output_tensor_scale.stride(0), - output_tensor_scale.stride(1), - output_index, - output_index.stride(0), - output_index.stride(1), - topk_num=recv_topk.shape[1], - num_warps=num_warps, - HIDDEN_SIZE=hidden_size, - HIDDEN_SIZE_PAD=triton.next_power_of_2(hidden_size), - SCALE_HIDDEN_SIZE=hidden_size // BLOCK_D, - SCALE_HIDDEN_SIZE_PAD=triton.next_power_of_2(hidden_size // BLOCK_D), - ) - return - - -@torch.no_grad() -def ep_fill_m_indices( - num_unaligned_recv_tokens_per_expert: torch.Tensor, - m_indices: torch.Tensor, -): - """Build aligned expert offsets and DeepGEMM's expert index vector.""" - block_e = 128 - num_experts = num_unaligned_recv_tokens_per_expert.shape[0] - assert m_indices.shape[0] % block_e == 0 - - expert_start_loc = torch.empty_like(num_unaligned_recv_tokens_per_expert) - _fwd_kernel_ep_scatter_1[(num_experts,)]( - num_unaligned_recv_tokens_per_expert, - expert_start_loc, - m_indices, - num_experts=num_experts, - num_warps=8, - BLOCK_E=block_e, - BLOCK_EXPERT_NUM=triton.next_power_of_2(num_experts), - ALIGN_COUNTS=True, - ) - return expert_start_loc - - -@triton.jit -def _zero_expanded_padding_kernel( - recv_x, - recv_x_stride_m, - recv_x_stride_k, - recv_x_scale, - recv_x_scale_stride_m, - recv_x_scale_stride_k, - recv_topk_weights, - num_unaligned_recv_tokens_per_expert, - expert_start_loc, - hidden_size: tl.constexpr, - scale_hidden_size: tl.constexpr, - BLOCK_M: tl.constexpr, - BLOCK_K: tl.constexpr, - BLOCK_SCALE_K: tl.constexpr, -): - expert_id = tl.program_id(0) - pad_block_id = tl.program_id(1) - hidden_block_id = tl.program_id(2) - expert_start = tl.load(expert_start_loc + expert_id) - actual_count = tl.load(num_unaligned_recv_tokens_per_expert + expert_id) - aligned_count = (actual_count + 127) // 128 * 128 - pad_offsets = pad_block_id * BLOCK_M + tl.arange(0, BLOCK_M) - row_offsets = (expert_start + actual_count + pad_offsets).to(tl.int64) - row_mask = pad_offsets < aligned_count - actual_count - - hidden_offsets = hidden_block_id * BLOCK_K + tl.arange(0, BLOCK_K) - x_ptrs = recv_x + row_offsets[:, None] * recv_x_stride_m + hidden_offsets[None, :] * recv_x_stride_k - tl.store(x_ptrs, 0.0, mask=row_mask[:, None] & (hidden_offsets[None, :] < hidden_size)) - if hidden_block_id == 0: - scale_offsets = tl.arange(0, BLOCK_SCALE_K) - scale_ptrs = ( - recv_x_scale + row_offsets[:, None] * recv_x_scale_stride_m + scale_offsets[None, :] * recv_x_scale_stride_k - ) - tl.store(scale_ptrs, 0.0, mask=row_mask[:, None] & (scale_offsets[None, :] < scale_hidden_size)) - tl.store(recv_topk_weights + row_offsets, 0.0, mask=row_mask) - - -@torch.no_grad() -def ep_zero_expanded_padding( - recv_x: torch.Tensor, - recv_x_scale: torch.Tensor, - recv_topk_weights: torch.Tensor, - num_unaligned_recv_tokens_per_expert: torch.Tensor, - expert_start_loc: torch.Tensor, -): - block_m = 8 - block_k = 256 - scale_hidden_size = recv_x_scale.shape[1] - grid = ( - num_unaligned_recv_tokens_per_expert.shape[0], - triton.cdiv(127, block_m), - triton.cdiv(recv_x.shape[1], block_k), - ) - _zero_expanded_padding_kernel[grid]( - recv_x, - recv_x.stride(0), - recv_x.stride(1), - recv_x_scale, - recv_x_scale.stride(0), - recv_x_scale.stride(1), - recv_topk_weights, - num_unaligned_recv_tokens_per_expert, - expert_start_loc, - hidden_size=recv_x.shape[1], - scale_hidden_size=scale_hidden_size, - BLOCK_M=block_m, - BLOCK_K=block_k, - BLOCK_SCALE_K=triton.next_power_of_2(scale_hidden_size), - num_warps=4, - ) - - -@triton.jit -def _accumulate_expanded_chunk_kernel( - total_recv_tokens, - chunk, - chunk_stride_m, - chunk_stride_k, - chunk_start, - chunk_end, - weights, - recv_src_metadata, - metadata_stride_m, - metadata_stride_k, - output, - output_stride_m, - output_stride_k, - TOPK: tl.constexpr, - BLOCK_D: tl.constexpr, -): - hidden_block_id = tl.program_id(0) - start_recv_token_id = tl.program_id(1) - recv_token_grid_size = tl.num_programs(1) - hidden_offsets = hidden_block_id * BLOCK_D + tl.arange(0, BLOCK_D) - - for recv_token_id in range(start_recv_token_id, total_recv_tokens, recv_token_grid_size): - output_ptrs = output + recv_token_id * output_stride_m + hidden_offsets * output_stride_k - accumulator = tl.load(output_ptrs).to(tl.float32) - for topk_id in range(TOPK): - slot = tl.load(recv_src_metadata + recv_token_id * metadata_stride_m + (topk_id + 2) * metadata_stride_k) - if slot >= chunk_start and slot < chunk_end: - local_row = (slot - chunk_start).to(tl.int64) - value = tl.load(chunk + local_row * chunk_stride_m + hidden_offsets * chunk_stride_k) - weight = tl.load(weights + slot) - accumulator += value.to(tl.float32) * weight - tl.store(output_ptrs, accumulator) - - -@torch.no_grad() -def ep_accumulate_expanded_chunk( - chunk: torch.Tensor, - chunk_start: int, - weights: torch.Tensor, - recv_src_metadata: torch.Tensor, - output: torch.Tensor, -): - """Accumulate one contiguous expanded W2 chunk into dense receive-token rows.""" - topk = recv_src_metadata.shape[1] - 2 - block_d = 1024 - assert chunk.shape[1] == output.shape[1] and output.shape[1] % block_d == 0 - grid = (triton.cdiv(output.shape[1], block_d), min(output.shape[0], 1024)) - _accumulate_expanded_chunk_kernel[grid]( - output.shape[0], - chunk, - chunk.stride(0), - chunk.stride(1), - chunk_start, - chunk_start + chunk.shape[0], - weights, - recv_src_metadata, - recv_src_metadata.stride(0), - recv_src_metadata.stride(1), - output, - output.stride(0), - output.stride(1), - TOPK=topk, - BLOCK_D=block_d, - num_warps=2, - ) - - -@triton.jit -def _compact_expanded_metadata_kernel( - recv_src_metadata, - metadata_stride_m, - metadata_stride_k, - TOPK: tl.constexpr, - BLOCK_TOPK: tl.constexpr, -): - recv_token_id = tl.program_id(0) - topk_id = tl.arange(0, BLOCK_TOPK) - slot = tl.where(topk_id == 0, recv_token_id, -1) - tl.store( - recv_src_metadata + recv_token_id * metadata_stride_m + (topk_id + 2) * metadata_stride_k, - slot, - mask=topk_id < TOPK, - ) - - -@torch.no_grad() -def ep_compact_expanded_metadata(recv_src_metadata: torch.Tensor): - """Point expanded combine metadata at pre-reduced dense token rows.""" - topk = recv_src_metadata.shape[1] - 2 - if recv_src_metadata.shape[0] == 0: - return - _compact_expanded_metadata_kernel[(recv_src_metadata.shape[0],)]( - recv_src_metadata, - recv_src_metadata.stride(0), - recv_src_metadata.stride(1), - TOPK=topk, - BLOCK_TOPK=triton.next_power_of_2(topk), - num_warps=1, - ) - - -@triton.jit -def _fwd_kernel_ep_gather( - total_token_num, - input_tensor, - input_tensor_stride0, - input_tensor_stride1, - recv_topk_ids, - recv_topk_ids_stride0, - recv_topk_ids_stride1, - recv_topk_weight, - recv_topk_weight_stride0, - recv_topk_weight_stride1, - input_index, - input_index_stride0, - input_index_stride1, - output_tensor, - output_tensor_stride0, - output_tensor_stride1, - topk_num: tl.constexpr, - BLOCK_D: tl.constexpr, -): - cur_block = tl.program_id(0) - start_cur_token = tl.program_id(1) - grid_num = tl.num_programs(1) - - for cur_token in range(start_cur_token, total_token_num, grid_num): - off_d = tl.arange(0, BLOCK_D) - accumulator = tl.zeros([BLOCK_D], dtype=tl.float32) - for topk_index in range(0, topk_num): - expert_id = tl.load(recv_topk_ids + cur_token * recv_topk_ids_stride0 + topk_index) - if expert_id >= 0: - source_token_index = tl.load(input_index + cur_token * input_index_stride0 + topk_index) - acc_weight = tl.load(recv_topk_weight + cur_token * recv_topk_weight_stride0 + topk_index) - tmp = tl.load(input_tensor + source_token_index * input_tensor_stride0 + cur_block * BLOCK_D + off_d) - accumulator += tmp.to(tl.float32) * acc_weight - - tl.store( - output_tensor + cur_token * output_tensor_stride0 + cur_block * BLOCK_D + off_d, - accumulator.to(output_tensor.dtype.element_ty), - ) - - -@torch.no_grad() -def ep_gather( - input_tensor: torch.Tensor, - recv_topk_ids: torch.Tensor, - recv_topk_weight: torch.Tensor, - input_index: torch.Tensor, - output_tensor: torch.Tensor, -): - BLOCK_D = 1024 # block size of quantization - num_warps = 2 - num_tokens = output_tensor.shape[0] - hidden_size = input_tensor.shape[1] - assert hidden_size % BLOCK_D == 0 - grid = (triton.cdiv(hidden_size, BLOCK_D), min(num_tokens, 1024)) - _fwd_kernel_ep_gather[grid]( - num_tokens, - input_tensor, - input_tensor.stride(0), - input_tensor.stride(1), - recv_topk_ids, - recv_topk_ids.stride(0), - recv_topk_ids.stride(1), - recv_topk_weight, - recv_topk_weight.stride(0), - recv_topk_weight.stride(1), - input_index, - input_index.stride(0), - input_index.stride(1), - output_tensor, - output_tensor.stride(0), - output_tensor.stride(1), - topk_num=recv_topk_ids.shape[1], - num_warps=num_warps, - BLOCK_D=BLOCK_D, - ) - return diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py index db32769064..87d3f74ccc 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py @@ -12,11 +12,11 @@ from lightllm.common.basemodel.triton_kernel.quantization.fp8act_quant_kernel import ( per_token_group_quant_fp8, ) -from lightllm.common.basemodel.triton_kernel.fused_moe.deepep_scatter_gather import ( - ep_accumulate_expanded_chunk, - ep_compact_expanded_metadata, - ep_fill_m_indices, - ep_zero_expanded_padding, +from lightllm.common.basemodel.triton_kernel.fused_moe.deepep_expanded_layout_kernels import ( + ep_build_m_indices, + ep_compact_metadata, + ep_gather_chunk, + ep_zero_padding, ) from lightllm.utils.envs_utils import ( get_deepep_num_max_dispatch_tokens_per_rank_prefill, @@ -290,7 +290,7 @@ def fused_experts_impl( # needed once the received tensors have been produced. del qinput_tensor, input_scale - gather_out = expanded_moe_chunked_reduce( + gather_out = chunked_expanded_moe_forward( handle.num_recv_tokens_per_expert_list, handle.num_unaligned_recv_tokens_per_expert, recv_x, @@ -306,8 +306,7 @@ def fused_experts_impl( ) del recv_x - # W2 chunks were reduced to the deduplicated receive-token layout. Keep - # the expanded handle for routing, but point its slots at the dense rows. + # normal combine combined_x, _, event = buffer.combine( gather_out, handle, @@ -362,19 +361,21 @@ def get_prefill_moe_workspace( return workspace -def expanded_moe_chunked_reduce( - num_recv_tokens_per_expert_list: List[int], - num_unaligned_recv_tokens_per_expert: torch.Tensor, - recv_x: Tuple[torch.Tensor, torch.Tensor], - recv_topk_weights: torch.Tensor, - recv_src_metadata: torch.Tensor, - w1: torch.Tensor, - w1_scale: torch.Tensor, - w2: torch.Tensor, - w2_scale: torch.Tensor, +def chunked_expanded_moe_forward( + num_recv_tokens_per_expert_list: List[int], # [num_local_experts], 128-aligned token counts + num_unaligned_recv_tokens_per_expert: torch.Tensor, # [num_local_experts], actual token counts + recv_x: Tuple[ + torch.Tensor, torch.Tensor # [fp8, scale] + ], # ([num_expanded_tokens, hidden_size], [num_expanded_tokens, hidden_size // block_size_k]) + recv_topk_weights: torch.Tensor, # [num_expanded_tokens] + recv_src_metadata: torch.Tensor, # [num_recv_tokens, topk + 2] + w1: torch.Tensor, # [num_local_experts, 2 * intermediate_size, hidden_size] + w1_scale: torch.Tensor, # [num_local_experts, 2 * intermediate_size // block_size_k, hidden_size // block_size_k] + w2: torch.Tensor, # [num_local_experts, hidden_size, intermediate_size] + w2_scale: torch.Tensor, # [num_local_experts, hidden_size // block_size_k, intermediate_size // block_size_k] block_size_k: int, - workspace: torch.Tensor, - hidden_dtype: torch.dtype, + workspace: torch.Tensor, # [workspace_bytes], uint8 + hidden_dtype: torch.dtype, # scalar dtype descriptor ): """Run bounded expanded MoE and rewrite metadata for dense DeepEP combine.""" alignment = 128 @@ -390,8 +391,8 @@ def expanded_moe_chunked_reduce( return torch.empty((0, hidden_size), device=recv_x[0].device, dtype=hidden_dtype) m_indices = torch.empty(all_tokens, device=recv_x[0].device, dtype=torch.int32) - expert_start_loc = ep_fill_m_indices(num_unaligned_recv_tokens_per_expert, m_indices) - ep_zero_expanded_padding( + expert_start_loc = ep_build_m_indices(num_unaligned_recv_tokens_per_expert, m_indices) + ep_zero_padding( recv_x[0], recv_x[1], recv_topk_weights, @@ -485,9 +486,9 @@ def workspace_quant_alloc(shape, dtype, device): m_indices[chunk_start:chunk_end], ) del qsilu_out, qsilu_out_scale, silu_out - ep_accumulate_expanded_chunk(gemm_out_b, chunk_start, recv_topk_weights, recv_src_metadata, gather_out) + ep_gather_chunk(gemm_out_b, chunk_start, recv_topk_weights, recv_src_metadata, gather_out) - ep_compact_expanded_metadata(recv_src_metadata) + ep_compact_metadata(recv_src_metadata) return gather_out diff --git a/unit_tests/common/fused_moe/test_deepep.py b/unit_tests/common/fused_moe/test_deepep.py index 45778244b7..ecc92af1c0 100644 --- a/unit_tests/common/fused_moe/test_deepep.py +++ b/unit_tests/common/fused_moe/test_deepep.py @@ -6,10 +6,8 @@ import torch import torch.distributed as dist import deep_ep -import random import numpy as np from lightllm.common.fused_moe.grouped_fused_moe_ep import fused_experts_impl -from lightllm.common.fused_moe.deepep_scatter_gather import ep_scatter, ep_gather from typing import Tuple from lightllm.utils.log_utils import init_logger @@ -309,80 +307,5 @@ def test_end2end(): torch.multiprocessing.spawn(case1, args=(num_processes,), nprocs=num_processes) -def test_scatter_gather(): - block_size = 128 - num_recv_tokens_per_expert_list = [0] * 32 - num_recv_tokens_per_expert_list[6] = 128 - num_recv_tokens_per_expert_list[7] = 128 - num_recv_tokens_per_expert_list[8] = 128 - num_recv_tokens_per_expert = torch.tensor(num_recv_tokens_per_expert_list, dtype=torch.int, device="cuda") - - all_tokens = sum(num_recv_tokens_per_expert_list) - m_indices_ref = torch.empty(all_tokens, device="cuda", dtype=torch.int32) - m_indices = torch.empty(all_tokens, device="cuda", dtype=torch.int32) - - recv_x = torch.randn((7, 4096), device="cuda", dtype=torch.float32).to(torch.float8_e4m3fn) - recv_x_scale = torch.randn((7, 4096 // block_size), device="cuda", dtype=torch.float32) - - recv_topk_id = torch.ones((7, 8), device="cuda", dtype=torch.int32) * -1 - recv_topk_weights = torch.zeros((7, 8), device="cuda", dtype=torch.float) - for i in range(7): - for j in range(4): - idx = random.randint(0, 7) - expert_id = random.randint(6, 8) - recv_topk_id[i][idx] = expert_id - recv_topk_weights[i][idx] = random.randint(0, 10) / 10.0 - - output_indexs = torch.zeros_like(recv_topk_id) - output_tensor = torch.zeros((all_tokens, 4096), device="cuda", dtype=torch.float32).to(torch.float8_e4m3fn) - output_tensor_ref = torch.zeros((all_tokens, 4096), device="cuda", dtype=torch.float32).to(torch.float8_e4m3fn) - - output_tensor_scale = torch.zeros((all_tokens, 4096 // block_size), device="cuda", dtype=torch.float32) - output_tensor_scale_ref = torch.zeros((all_tokens, 4096 // block_size), device="cuda", dtype=torch.float32) - - expert_start_loc = torch.cumsum(torch.tensor([0] + num_recv_tokens_per_expert_list[:-1], device="cuda"), dim=0) - - cur = 0 - for i, k in enumerate(num_recv_tokens_per_expert_list): - m_indices_ref[cur : cur + k] = i - cur += k - - ep_scatter( - recv_x, - recv_x_scale, - recv_topk_id, - num_recv_tokens_per_expert, - expert_start_loc, - output_tensor, - output_tensor_scale, - m_indices, - output_indexs, - ) - assert torch.allclose(m_indices, m_indices_ref, atol=1e-2, rtol=0) - - for i in range(recv_topk_id.shape[0]): - for j in range(recv_topk_id.shape[1]): - if recv_topk_id[i][j] >= 0: - dst = output_indexs[i][j] - output_tensor_ref[dst][:] = recv_x[i][:] - output_tensor_scale_ref[dst][:] = recv_x_scale[i][:] - - assert torch.allclose(output_tensor.to(torch.float), output_tensor_ref.to(torch.float), atol=1e-2, rtol=0) - assert torch.allclose(output_tensor_scale, output_tensor_scale_ref, atol=1e-2, rtol=0) - - #### gather - - gather_out_ref = torch.zeros_like(recv_x, device="cuda", dtype=torch.bfloat16) - gather_out = torch.empty_like(recv_x, device="cuda", dtype=torch.bfloat16) - gather_input = torch.zeros((all_tokens, 4096), device="cuda", dtype=torch.bfloat16) - for i in range(recv_topk_id.shape[0]): - for j in range(recv_topk_id.shape[1]): - if recv_topk_id[i][j] >= 0: - dst = output_indexs[i][j] - gather_out_ref[i][:] += gather_input[dst][:] * recv_topk_weights[i][j] - ep_gather(gather_input, recv_topk_id, recv_topk_weights, output_indexs, gather_out) - assert torch.allclose(gather_out, gather_out_ref, atol=1e-2, rtol=0) - - if __name__ == "__main__": pytest.main()