From a84cdb80a100f0163d11dc94030a14befef1d65b Mon Sep 17 00:00:00 2001 From: niushengxiao Date: Fri, 17 Jul 2026 16:57:37 +0800 Subject: [PATCH 01/20] 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 115b0c48ea..52119031a7 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 @@ -93,8 +92,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 c20acb12f7..4f9196c8bc 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 @@ -226,20 +226,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 cfc82facee..9f024d5c5d 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 @@ -182,6 +176,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(): @@ -214,87 +210,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( @@ -315,7 +259,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 4671329840..6d3c7810e9 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 2ac834d89e..ceb040fe30 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 = "fp8w8a8-b128-vllm" 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 07a69db0fe..e679f98bf8 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 @@ -169,22 +171,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 dae79cc8a6..e6e303d7c5 100644 --- a/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py @@ -501,7 +501,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 @@ -527,7 +534,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 7edfd5a6f9..26b4a11861 100644 --- a/lightllm/models/qwen3_moe/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/qwen3_moe/layer_infer/transformer_layer_infer.py @@ -315,7 +315,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 @@ -341,7 +348,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 81482835889df9700731c7ba6752f77aa93fbd70 Mon Sep 17 00:00:00 2001 From: niushengxiao Date: Mon, 20 Jul 2026 18:03:05 +0800 Subject: [PATCH 02/20] 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 6d3c7810e9..4e4cf551a2 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 e679f98bf8..ed8897b03d 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 @@ -184,7 +184,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(): @@ -222,7 +223,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 9cfb7834c1ea67d4cf46ce910b033e199e9547b9 Mon Sep 17 00:00:00 2001 From: niushengxiao Date: Mon, 20 Jul 2026 19:32:58 +0800 Subject: [PATCH 03/20] 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 9f024d5c5d..7f94fe413f 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 @@ -224,7 +224,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, @@ -259,8 +259,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 4e4cf551a2..1b3903a4f8 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() From 7dfb5d80f1a6d26f69d483616232894e4d2724d7 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Tue, 4 Aug 2026 06:36:24 +0000 Subject: [PATCH 04/20] fix: configure expandable segments before spawning workers --- lightllm/__init__.py | 4 ---- lightllm/server/api_start.py | 5 +++++ 2 files changed, 5 insertions(+), 4 deletions(-) diff --git a/lightllm/__init__.py b/lightllm/__init__.py index bc09ec5a17..e9ba6f3041 100644 --- a/lightllm/__init__.py +++ b/lightllm/__init__.py @@ -2,7 +2,3 @@ if is_musa(): import torchada # noqa: F401 -else: - import torch - - torch._C._accelerator_setAllocatorSettings("expandable_segments:True") diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index 67ae286a7c..d5cf0df347 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -26,11 +26,16 @@ auto_set_response_parsers, ) from lightllm.utils.dist_check_utils import auto_configure_allreduce_flags_from_args +from lightllm.utils.device_utils import is_musa logger = init_logger(__name__) def _set_envs_and_config(args: StartArgs): + if not is_musa(): + # 减少动态 batch/序列长度引起的 CUDA 显存碎片;该配置会被所有子进程继承, + # CUDA IPC 或自定义 allocator 场景需关注 expandable_segments 的兼容性。 + os.environ["PYTORCH_ALLOC_CONF"] = "expandable_segments:True" mp.set_start_method("spawn", force=True) From 90326cefed6ffd49ebab02b951ffacc76be01f0e Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Tue, 4 Aug 2026 07:03:48 +0000 Subject: [PATCH 05/20] fix --- lightllm/server/api_start.py | 5 ----- 1 file changed, 5 deletions(-) diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index d5cf0df347..67ae286a7c 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -26,16 +26,11 @@ auto_set_response_parsers, ) from lightllm.utils.dist_check_utils import auto_configure_allreduce_flags_from_args -from lightllm.utils.device_utils import is_musa logger = init_logger(__name__) def _set_envs_and_config(args: StartArgs): - if not is_musa(): - # 减少动态 batch/序列长度引起的 CUDA 显存碎片;该配置会被所有子进程继承, - # CUDA IPC 或自定义 allocator 场景需关注 expandable_segments 的兼容性。 - os.environ["PYTORCH_ALLOC_CONF"] = "expandable_segments:True" mp.set_start_method("spawn", force=True) From dfbbcf0803014d904f4ed8b5eb5f892dfde1bd4c Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Tue, 4 Aug 2026 07:43:23 +0000 Subject: [PATCH 06/20] refactor: group MoE workspace configuration --- .../meta_weights/fused_moe/fused_moe_weight.py | 7 +++---- .../fused_moe/impl/deepgemm_impl.py | 6 +++--- .../fused_moe/grouped_fused_moe_ep.py | 18 ++++++++++++------ .../layer_infer/transformer_layer_infer.py | 11 ++++++----- .../layer_infer/transformer_layer_infer.py | 11 ++++++----- 5 files changed, 30 insertions(+), 23 deletions(-) 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 4f9196c8bc..f906650256 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 @@ -14,6 +14,7 @@ from lightllm.utils.envs_utils import get_redundancy_expert_ids, get_redundancy_expert_num, get_env_start_args from lightllm.utils.dist_utils import get_global_world_size, get_global_rank from lightllm.utils.log_utils import init_logger +from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_fused_moe_ep import MoeWorkspaceConfig logger = init_logger(__name__) @@ -232,8 +233,7 @@ def prefilled_group_gemm( recv_topk_idx: torch.Tensor, recv_topk_weights: torch.Tensor, hidden_dtype=torch.bfloat16, - workspace_index: int = 0, - workspace_count: int = 1, + workspace_config: MoeWorkspaceConfig = MoeWorkspaceConfig(), ): 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( @@ -246,8 +246,7 @@ def prefilled_group_gemm( w13=self.w13, w2=self.w2, hidden_dtype=hidden_dtype, - workspace_index=workspace_index, - workspace_count=workspace_count, + workspace_config=workspace_config, ) 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 7f94fe413f..121f9aa4c0 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 @@ -11,6 +11,7 @@ fused_experts, get_ep_num_sms, masked_group_gemm, + MoeWorkspaceConfig, get_prefill_moe_workspace, chunked_expanded_moe_forward, quantize_fused_experts_input, @@ -218,8 +219,7 @@ def prefilled_group_gemm( w13: WeightPack, w2: WeightPack, hidden_dtype=torch.bfloat16, - workspace_index: int = 0, - workspace_count: int = 1, + workspace_config: MoeWorkspaceConfig = MoeWorkspaceConfig(), ): w13_weight, w13_scale = w13.weight, w13.weight_scale w2_weight, w2_scale = w2.weight, w2.weight_scale @@ -235,7 +235,7 @@ def prefilled_group_gemm( w2_weight, w2_scale, self.quant_method.block_size, - get_prefill_moe_workspace(workspace_index, workspace_count), + get_prefill_moe_workspace(workspace_config), hidden_dtype, ) del recv_x 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 1b3903a4f8..39296ec8e3 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,8 +1,9 @@ """Fused MoE kernel.""" + import torch import triton import triton.language as tl -from typing import Any, Callable, Dict, List, Optional, Tuple +from typing import Any, Callable, Dict, List, NamedTuple, 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 @@ -29,6 +30,12 @@ _MEGA_MOE_STATES: Dict[Tuple[int, int, int, int], Dict[str, Any]] = {} SUPPORTED_EP_EXPERT_DTYPES = ("fp8w8a8-b128-deepgemm", "fp4fp8-b32-deepgemm") + +class MoeWorkspaceConfig(NamedTuple): + index: int = 0 + count: int = 1 + + try: from deep_ep import Buffer, EventOverlap import deep_gemm @@ -351,13 +358,12 @@ def deepgemm_grouped_fp8_nt_contiguous( def get_prefill_moe_workspace( - workspace_index: int = 0, - workspace_count: int = 1, + workspace_config: MoeWorkspaceConfig = MoeWorkspaceConfig(), ): 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) + assert 0 <= workspace_config.index < workspace_config.count + workspace_size = workspace.numel() // workspace_config.count + workspace = workspace.narrow(0, workspace_config.index * workspace_size, workspace_size) return workspace diff --git a/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py index e6e303d7c5..2d08bafd59 100644 --- a/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py @@ -7,7 +7,10 @@ from lightllm.models.llama.layer_infer.transformer_layer_infer import LlamaTransformerLayerInfer from lightllm.models.deepseek2.triton_kernel.rotary_emb import rotary_emb_fwd from lightllm.models.deepseek2.infer_struct import Deepseek2InferStateInfo -from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_fused_moe_ep import use_sm100_mega_moe +from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_fused_moe_ep import ( + MoeWorkspaceConfig, + use_sm100_mega_moe, +) from functools import partial from lightllm.models.llama.yarn_rotary_utils import get_deepseek_mscale from lightllm.utils.envs_utils import get_env_start_args @@ -507,8 +510,7 @@ def overlap_tpsp_context_forward( _0_recv_x, _0_recv_topk_idx, _0_recv_topk_weight, - workspace_index=0, - workspace_count=2, + workspace_config=MoeWorkspaceConfig(index=0, count=2), ) # 1 dispatch execute @@ -540,8 +542,7 @@ def overlap_tpsp_context_forward( _1_recv_x, _1_recv_topk_idx, _1_recv_topk_weight, - workspace_index=1, - workspace_count=2, + workspace_config=MoeWorkspaceConfig(index=1, 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 26b4a11861..09be0e2519 100644 --- a/lightllm/models/qwen3_moe/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/qwen3_moe/layer_infer/transformer_layer_infer.py @@ -6,7 +6,10 @@ from lightllm.models.llama.layer_infer.transformer_layer_infer import LlamaTransformerLayerInfer from lightllm.models.llama.infer_struct import LlamaInferStateInfo from lightllm.models.llama.triton_kernel.rotary_emb import rotary_emb_fwd -from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_fused_moe_ep import use_sm100_mega_moe +from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_fused_moe_ep import ( + MoeWorkspaceConfig, + use_sm100_mega_moe, +) from lightllm.utils.dist_utils import get_global_world_size from lightllm.utils.envs_utils import get_env_start_args @@ -321,8 +324,7 @@ def overlap_tpsp_context_forward( _0_recv_x, _0_recv_topk_idx, _0_recv_topk_weight, - workspace_index=0, - workspace_count=2, + workspace_config=MoeWorkspaceConfig(index=0, count=2), ) # 1 dispatch execute @@ -354,8 +356,7 @@ def overlap_tpsp_context_forward( _1_recv_x, _1_recv_topk_idx, _1_recv_topk_weight, - workspace_index=1, - workspace_count=2, + workspace_config=MoeWorkspaceConfig(index=1, count=2), ) # wait 0 combine From 483eb36949390d61a45d2fafb557e912cb552ce8 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 5 Aug 2026 03:06:14 +0000 Subject: [PATCH 07/20] refactor: clarify expert quantization priority --- lightllm/common/quantization/__init__.py | 30 +++++++++++++++++------- 1 file changed, 22 insertions(+), 8 deletions(-) diff --git a/lightllm/common/quantization/__init__.py b/lightllm/common/quantization/__init__.py index ceb040fe30..6ce4509776 100644 --- a/lightllm/common/quantization/__init__.py +++ b/lightllm/common/quantization/__init__.py @@ -22,8 +22,10 @@ class Quantcfg: def __init__(self, network_config, quant_type="none", custom_cfg_path=None, expert_dtype=None): self.layer_num = network_config["n_layer"] self.quant_type = quant_type - self.expert_dtype = expert_dtype - self.network_config_ = network_config + self.start_args_expert_dtype = expert_dtype + self.config_expert_dtype = network_config.get("expert_dtype", None) + # Parse quant_cfg first so model config only fills missing per-layer fused_moe entries; + # an explicit startup argument is applied afterward and overrides every layer. self._parse_custom_cfg(custom_cfg_path) self._parse_network_config(network_config) @@ -51,19 +53,30 @@ 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) + expert_dtype = self.start_args_expert_dtype or self.config_expert_dtype 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 + layer_quant_cfg = self.quant_cfg[layer_num] + if self.start_args_expert_dtype is not None: + # 优先级 1:启动参数显式指定 expert_dtype,覆盖其他来源的配置。 + layer_quant_cfg["fused_moe"] = target + elif "fused_moe" in layer_quant_cfg: + # 优先级 2:未指定启动参数时,保留 quant_cfg 中已有的逐层配置。 + continue else: - self.quant_cfg[layer_num].setdefault("fused_moe", target) - logger.info(f"select fused_moe quant way from expert_dtype=`{expert_dtype}`: {target}") + # 优先级 3:前两个来源均未指定时,使用 config.json 中的 expert_dtype。 + layer_quant_cfg["fused_moe"] = target + + source = "startup arguments" if self.start_args_expert_dtype is not None else "model config" + logger.info( + f"select fused_moe quant way from {source} expert_dtype=`{expert_dtype}`: {target}; " + "priority is startup arguments > quant_cfg > model config." + ) def _mapping_quant_method(self): if self.hf_quantization_method == "fp8": @@ -76,6 +89,7 @@ def _mapping_quant_method(self): else: self.quant_type = "fp8w8a8-b128-vllm" logger.info(f"select fp8w8a8-b128 quant way: {self.quant_type}") + self._mapping_expert_quant_method() elif self.hf_quantization_method == "awq": self.quant_type = "awq" From b6a4e3b3506a3c32a0417d4498e14ff85e09ead3 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 5 Aug 2026 06:43:32 +0000 Subject: [PATCH 08/20] fix --- .dockerignore | 30 ------------------------------ 1 file changed, 30 deletions(-) delete mode 100644 .dockerignore diff --git a/.dockerignore b/.dockerignore deleted file mode 100644 index 1ac2bb0d48..0000000000 --- a/.dockerignore +++ /dev/null @@ -1,30 +0,0 @@ -.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 From a6a8c3f46570fb42023d0000cd3d20abda9637b6 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 5 Aug 2026 07:31:41 +0000 Subject: [PATCH 09/20] refactor: centralize DeepEP prefill workspace slicing --- .../fused_moe/fused_moe_weight.py | 5 ++- .../fused_moe/impl/deepgemm_impl.py | 6 ++-- .../fused_moe/grouped_fused_moe_ep.py | 19 ++-------- lightllm/distributed/communication_op.py | 36 +++++++++++++++---- .../layer_infer/transformer_layer_infer.py | 5 ++- .../layer_infer/transformer_layer_infer.py | 5 ++- 6 files changed, 40 insertions(+), 36 deletions(-) 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 f906650256..7f369c4fd8 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 @@ -14,7 +14,6 @@ from lightllm.utils.envs_utils import get_redundancy_expert_ids, get_redundancy_expert_num, get_env_start_args from lightllm.utils.dist_utils import get_global_world_size, get_global_rank from lightllm.utils.log_utils import init_logger -from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_fused_moe_ep import MoeWorkspaceConfig logger = init_logger(__name__) @@ -233,7 +232,7 @@ def prefilled_group_gemm( recv_topk_idx: torch.Tensor, recv_topk_weights: torch.Tensor, hidden_dtype=torch.bfloat16, - workspace_config: MoeWorkspaceConfig = MoeWorkspaceConfig(), + microbatch_index: int = 0, ): 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( @@ -246,7 +245,7 @@ def prefilled_group_gemm( w13=self.w13, w2=self.w2, hidden_dtype=hidden_dtype, - workspace_config=workspace_config, + microbatch_index=microbatch_index, ) 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 121f9aa4c0..7976f0bcef 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 @@ -11,8 +11,6 @@ fused_experts, get_ep_num_sms, masked_group_gemm, - MoeWorkspaceConfig, - get_prefill_moe_workspace, chunked_expanded_moe_forward, quantize_fused_experts_input, ) @@ -219,7 +217,7 @@ def prefilled_group_gemm( w13: WeightPack, w2: WeightPack, hidden_dtype=torch.bfloat16, - workspace_config: MoeWorkspaceConfig = MoeWorkspaceConfig(), + microbatch_index: int = 0, ): w13_weight, w13_scale = w13.weight, w13.weight_scale w2_weight, w2_scale = w2.weight, w2.weight_scale @@ -235,7 +233,7 @@ def prefilled_group_gemm( w2_weight, w2_scale, self.quant_method.block_size, - get_prefill_moe_workspace(workspace_config), + dist_group_manager.get_deep_ep_prefill_moe_workspace(microbatch_index), hidden_dtype, ) del recv_x 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 39296ec8e3..6a3991384b 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 @@ -3,7 +3,7 @@ import torch import triton import triton.language as tl -from typing import Any, Callable, Dict, List, NamedTuple, Optional, Tuple +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 @@ -31,11 +31,6 @@ SUPPORTED_EP_EXPERT_DTYPES = ("fp8w8a8-b128-deepgemm", "fp4fp8-b32-deepgemm") -class MoeWorkspaceConfig(NamedTuple): - index: int = 0 - count: int = 1 - - try: from deep_ep import Buffer, EventOverlap import deep_gemm @@ -308,7 +303,7 @@ def fused_experts_impl( w2, w2_scale, block_size_k, - get_prefill_moe_workspace(), + dist_group_manager.get_deep_ep_prefill_moe_workspace(), hidden_states.dtype, ) del recv_x @@ -357,16 +352,6 @@ 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_config: MoeWorkspaceConfig = MoeWorkspaceConfig(), -): - workspace = dist_group_manager.prefill_moe_workspace - assert 0 <= workspace_config.index < workspace_config.count - workspace_size = workspace.numel() // workspace_config.count - workspace = workspace.narrow(0, workspace_config.index * workspace_size, workspace_size) - return workspace - - 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 diff --git a/lightllm/distributed/communication_op.py b/lightllm/distributed/communication_op.py index ed8897b03d..2eb2ea5a79 100644 --- a/lightllm/distributed/communication_op.py +++ b/lightllm/distributed/communication_op.py @@ -109,7 +109,6 @@ def __init__(self): self.groups = [] self.ep_buffer = None self.ep_low_latency_buffer = None - self.prefill_moe_workspace = None self.ep_mega_moe_buffer = None self.ep_num_sms = None @@ -146,7 +145,6 @@ def new_deepep_group( if not enable_ep_moe: self.ep_buffer = None self.ep_low_latency_buffer = None - self.prefill_moe_workspace = None self.ep_mega_moe_buffer = None self.ep_num_sms = None return @@ -184,10 +182,6 @@ def new_deepep_group( low_latency_mode=True, num_qps_per_rank=(self.ll_num_experts // global_world_size), ) - # 当前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(): if moe_intermediate_size is None: raise ValueError("SM100 Mega MoE requires moe_intermediate_size or intermediate_size in model config") @@ -221,6 +215,36 @@ def _set_num_sms_for_deep_gemm(self, deepep_sms: int): except BaseException as e: logger.warning(f"set num sms for deep_gemm failed: {e}") + def get_deep_ep_prefill_moe_workspace(self, microbatch_index: int = 0) -> torch.Tensor: + """Return a slice of the workspace reused by DeepEP prefill MoE kernels. + + DeepEP's low-latency RDMA buffer is idle during prefill, so its local + storage is reused as temporary workspace for the expanded MoE compute + path to reduce peak GPU memory. With one communication group, the + default ``microbatch_index=0`` receives the whole workspace. With + multiple groups, the workspace is split into ``len(self.groups)`` + equal slices and each in-flight microbatch uses the slice matching its + group index. + + Args: + microbatch_index: Zero-based microbatch and communication-group + index assigned to this in-flight prefill computation. + + Returns: + A one-dimensional uint8 tensor view over the selected workspace + slice; no additional GPU memory is allocated. + + This workspace is only valid after the DeepEP group has been + initialized. The same returned slice must not be used concurrently by + overlapping calls. + """ + assert self.ep_low_latency_buffer is not None, "DeepEP low-latency buffer is not initialized" + workspace = self.ep_low_latency_buffer.get_local_buffer_tensor(torch.uint8, use_rdma_buffer=True) + microbatch_count = len(self.groups) + assert 0 <= microbatch_index < microbatch_count + workspace_size = workspace.numel() // microbatch_count + return workspace.narrow(0, microbatch_index * workspace_size, workspace_size) + def clear_deepep_buffer(self): """ Prefill MoE compute reuses the low-latency RDMA buffer as workspace. diff --git a/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py index 2d08bafd59..3254031056 100644 --- a/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py @@ -8,7 +8,6 @@ from lightllm.models.deepseek2.triton_kernel.rotary_emb import rotary_emb_fwd from lightllm.models.deepseek2.infer_struct import Deepseek2InferStateInfo from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_fused_moe_ep import ( - MoeWorkspaceConfig, use_sm100_mega_moe, ) from functools import partial @@ -510,7 +509,7 @@ def overlap_tpsp_context_forward( _0_recv_x, _0_recv_topk_idx, _0_recv_topk_weight, - workspace_config=MoeWorkspaceConfig(index=0, count=2), + microbatch_index=0, ) # 1 dispatch execute @@ -542,7 +541,7 @@ def overlap_tpsp_context_forward( _1_recv_x, _1_recv_topk_idx, _1_recv_topk_weight, - workspace_config=MoeWorkspaceConfig(index=1, count=2), + microbatch_index=1, ) # 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 09be0e2519..7311c4d141 100644 --- a/lightllm/models/qwen3_moe/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/qwen3_moe/layer_infer/transformer_layer_infer.py @@ -7,7 +7,6 @@ from lightllm.models.llama.infer_struct import LlamaInferStateInfo from lightllm.models.llama.triton_kernel.rotary_emb import rotary_emb_fwd from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_fused_moe_ep import ( - MoeWorkspaceConfig, use_sm100_mega_moe, ) from lightllm.utils.dist_utils import get_global_world_size @@ -324,7 +323,7 @@ def overlap_tpsp_context_forward( _0_recv_x, _0_recv_topk_idx, _0_recv_topk_weight, - workspace_config=MoeWorkspaceConfig(index=0, count=2), + microbatch_index=0, ) # 1 dispatch execute @@ -356,7 +355,7 @@ def overlap_tpsp_context_forward( _1_recv_x, _1_recv_topk_idx, _1_recv_topk_weight, - workspace_config=MoeWorkspaceConfig(index=1, count=2), + microbatch_index=1, ) # wait 0 combine From 14abe058af0d35a0439f06a51ad35787d4ef246e Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 5 Aug 2026 08:22:54 +0000 Subject: [PATCH 10/20] refactor: simplify DeepEP padding handling --- .../deepep_expanded_layout_kernels.py | 83 ++++++++++--------- .../fused_moe/grouped_fused_moe_ep.py | 10 ++- 2 files changed, 49 insertions(+), 44 deletions(-) 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 index ad9829dcea..c63d091864 100644 --- 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 @@ -6,8 +6,8 @@ @triton.jit def _ep_build_m_indices_kernel( num_unaligned_recv_tokens_per_expert, - expert_start_loc, m_indices, + padding_mask, num_experts: tl.constexpr, BLOCK_E: tl.constexpr, BLOCK_EXPERT_NUM: tl.constexpr, @@ -20,51 +20,58 @@ def _ep_build_m_indices_kernel( mask=offset_cumsum < num_experts, other=0, ) - tokens_per_expert = (tokens_per_expert + BLOCK_E - 1) // BLOCK_E * BLOCK_E + tokens_per_expert = tl.cdiv(tokens_per_expert, 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) + cur_expert_aligned_token_num = tl.cdiv(cur_expert_token_num, BLOCK_E) * BLOCK_E m_indices_start_ptr = m_indices + cur_expert_start + padding_mask_start_ptr = padding_mask + 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): + for start_m in tl.range(0, cur_expert_aligned_token_num, BLOCK_E, num_stages=4): tl.store( m_indices_start_ptr + start_m + off_expert, cur_expert, ) + tl.store( + padding_mask_start_ptr + start_m + off_expert, + tl.where(start_m + off_expert >= cur_expert_token_num, 1, 0), + ) @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] + padding_mask: torch.Tensor, # [num_expanded_tokens] + expert_alignment: int, ): - """Build the 128-aligned expert layout used by contiguous grouped GEMM. + """Build the 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. + Each expert's actual token count is rounded up to ``expert_alignment``. + ``m_indices`` is filled in-place with the owning expert ID for every real + and padding row. The alignment must match the value used by DeepEP + dispatch. - Returns: - ``expert_start_loc`` with shape ``[num_local_experts]``. Each value is - the expert's starting row in the expanded tensors. + ``padding_mask`` is filled in-place with ``1`` for alignment-padding rows + and ``0`` for real token rows. """ - block_e = 128 + assert expert_alignment >= 8, "expert_alignment must be at least the zero-padding BLOCK_M (8)" + assert triton.next_power_of_2(expert_alignment) == expert_alignment, "expert_alignment must be a power of two" num_experts = num_unaligned_recv_tokens_per_expert.shape[0] - assert m_indices.shape[0] % block_e == 0 + assert m_indices.shape[0] % expert_alignment == 0 + assert padding_mask.dtype == torch.int32 and padding_mask.shape == m_indices.shape - 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, + padding_mask, num_experts=num_experts, num_warps=8, - BLOCK_E=block_e, + BLOCK_E=expert_alignment, BLOCK_EXPERT_NUM=triton.next_power_of_2(num_experts), ) - return expert_start_loc @triton.jit @@ -76,33 +83,30 @@ def _ep_zero_padding_kernel( recv_x_scale_stride_m, recv_x_scale_stride_k, recv_topk_weights, - num_unaligned_recv_tokens_per_expert, - expert_start_loc, + padding_mask, 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 + row_block_id = tl.program_id(0) + hidden_block_id = tl.program_id(1) + row_offsets = row_block_id * BLOCK_M + tl.arange(0, BLOCK_M) + row_mask = tl.load(padding_mask + row_offsets) == 1 + row_offsets = row_offsets.to(tl.int64) hidden_offsets = hidden_block_id * BLOCK_K + tl.arange(0, BLOCK_K) + hidden_mask = hidden_offsets < hidden_size 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)) + tl.store(x_ptrs, 0.0, mask=row_mask[:, None] & hidden_mask[None, :]) if hidden_block_id == 0: scale_offsets = tl.arange(0, BLOCK_SCALE_K) + scale_mask = scale_offsets < scale_hidden_size 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(scale_ptrs, 0.0, mask=row_mask[:, None] & scale_mask[None, :]) tl.store(recv_topk_weights + row_offsets, 0.0, mask=row_mask) @@ -111,22 +115,22 @@ 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] + padding_mask: torch.Tensor, # [num_expanded_tokens], 1 for padding rows ): """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]``. + Rows marked by ``padding_mask`` 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]``. """ + assert padding_mask.dtype == torch.int32 and padding_mask.shape == recv_topk_weights.shape block_m = 8 block_k = 256 + assert padding_mask.shape[0] % block_m == 0, "padding_mask rows must be divisible by BLOCK_M (8)" scale_hidden_size = recv_x_scale.shape[1] grid = ( - num_unaligned_recv_tokens_per_expert.shape[0], - triton.cdiv(127, block_m), + triton.cdiv(padding_mask.shape[0], block_m), triton.cdiv(recv_x.shape[1], block_k), ) _ep_zero_padding_kernel[grid]( @@ -137,8 +141,7 @@ def ep_zero_padding( recv_x_scale.stride(0), recv_x_scale.stride(1), recv_topk_weights, - num_unaligned_recv_tokens_per_expert, - expert_start_loc, + padding_mask, hidden_size=recv_x.shape[1], scale_hidden_size=scale_hidden_size, BLOCK_M=block_m, 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 6a3991384b..55c06c537b 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 @@ -382,15 +382,17 @@ def chunked_expanded_moe_forward( 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_build_m_indices(num_unaligned_recv_tokens_per_expert, m_indices) + # 与 m_indices 一一对应:0 表示真实 token,1 表示 expert 对齐产生的 padding 行。 + # padding 行必须在 grouped GEMM 前清零,避免无效数据参与计算。 + padding_mask = torch.empty_like(m_indices) + ep_build_m_indices(num_unaligned_recv_tokens_per_expert, m_indices, padding_mask, alignment) ep_zero_padding( recv_x[0], recv_x[1], recv_topk_weights, - num_unaligned_recv_tokens_per_expert, - expert_start_loc, + padding_mask, ) - del expert_start_loc + del padding_mask gather_rows = recv_src_metadata.shape[0] gather_bytes = gather_rows * hidden_size * hidden_dtype.itemsize From 70e40bdb552e99edf8892218c90fa85288addd19 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 5 Aug 2026 09:48:54 +0000 Subject: [PATCH 11/20] feat: add reusable tensor buffer manager --- lightllm/utils/tensor_buffer_manager.py | 187 +++++++++++++++++++++++ test/utils/test_tensor_buffer_manager.py | 133 ++++++++++++++++ 2 files changed, 320 insertions(+) create mode 100644 lightllm/utils/tensor_buffer_manager.py create mode 100644 test/utils/test_tensor_buffer_manager.py diff --git a/lightllm/utils/tensor_buffer_manager.py b/lightllm/utils/tensor_buffer_manager.py new file mode 100644 index 0000000000..6565f9ad4a --- /dev/null +++ b/lightllm/utils/tensor_buffer_manager.py @@ -0,0 +1,187 @@ +import math +from dataclasses import dataclass +from typing import Dict, Iterable, Union + +import torch + + +Shape = Union[int, torch.Size, Iterable[int]] + + +class TensorBufferManager: + """从一块连续的底层 buffer 中申请和复用 tensor。 + + ``alloc`` 返回与底层 buffer 共享 storage 的连续 tensor view;``free`` + 将其占用的区间归还给管理器。每次申请的起始地址和预留空间都会按 + ``alignment_bytes`` 对齐,默认对齐到 256 bytes。 + + 调用方必须使用 ``free`` 释放 ``alloc`` 原样返回的 tensor。释放后不能 + 继续使用该 tensor、它的别名,或者尚未完成的异步操作;本管理器不负责 + CUDA stream 间的生命周期同步。 + + 使用示例:: + + buffer = torch.empty(1024 * 1024, dtype=torch.uint8, device="cuda") + manager = TensorBufferManager(buffer) + tensor = manager.alloc((128, 256), torch.bfloat16) + manager.free(tensor) + """ + + def __init__(self, buffer: torch.Tensor, alignment_bytes: int = 256): + """初始化连续内存池,并将可用区间调整到指定的地址对齐边界。""" + # 底层 buffer 必须是一块不参与 autograd 的连续内存,才能安全地按字节切分。 + if not isinstance(buffer, torch.Tensor): + raise TypeError("buffer must be a torch.Tensor") + if not buffer.is_contiguous(): + raise ValueError("buffer must be contiguous") + if buffer.requires_grad: + raise ValueError("buffer must not require gradients") + if alignment_bytes <= 0 or alignment_bytes & (alignment_bytes - 1): + raise ValueError("alignment_bytes must be a positive power of two") + + self._alignment_bytes = alignment_bytes + # 统一转换成一维 byte tensor,后续 offset 和 size 都使用字节作为单位。 + byte_buffer = buffer.view(torch.uint8).view(-1) + + # CPU buffer 的起始地址不一定满足指定对齐要求。构造阶段一次性切掉未对齐的 + # 前缀,使后续所有 allocation offset 都能从对齐后的地址 0 开始计算。 + address_remainder = buffer.data_ptr() % self._alignment_bytes + if address_remainder != 0: + aligned_offset = self._alignment_bytes - address_remainder + byte_buffer = byte_buffer[aligned_offset:] + + assert byte_buffer.numel() > 0, "buffer has no usable bytes after alignment" + self._byte_buffer = byte_buffer + # 初始状态下,整个对齐后的 byte buffer 都是一个连续空闲块。 + self._free_blocks = [_FreeBlock(offset=0, size=byte_buffer.numel())] + # 使用 data_ptr 定位 allocation,同时保存 tensor 对象用于释放时的身份校验。 + self._allocations: Dict[int, _Allocation] = {} + + def alloc(self, shape: Shape, dtype: torch.dtype) -> torch.Tensor: + """从内存池申请指定 shape 和 dtype 的连续 tensor view。""" + # tensor 的实际占用空间统一换算成字节数。 + shape = self._normalize_shape(shape) + element_size = self._element_size(dtype) + tensor_bytes = math.prod(shape) * element_size + + # 空 tensor 不占用底层 buffer,无需创建 allocation 记录。 + if tensor_bytes == 0: + return torch.empty(shape, dtype=dtype, device=self._byte_buffer.device) + + # 实际 tensor 只使用 tensor_bytes;内存池按对齐后的 reserved_bytes 划出区间。 + reserved_bytes = self._align_up(tensor_bytes) + offset = self._take_free_block(reserved_bytes) + # 先截取精确字节范围,再转换为调用方需要的 dtype 和 shape。 + tensor = self._byte_buffer[offset : offset + tensor_bytes].view(dtype).view(shape) + self._remember_allocation(tensor, offset, reserved_bytes) + return tensor + + def free(self, tensor: torch.Tensor) -> None: + """释放 ``alloc`` 返回的 tensor,并回收它的完整对齐预留区间。""" + # 空 tensor 没有占用内存池空间,可以直接忽略。 + if tensor.numel() == 0: + return + + # 先按地址查找,再校验 tensor 身份,避免旧 tensor 误释放地址相同的新 allocation。 + data_ptr = tensor.data_ptr() + allocation = self._allocations.get(data_ptr) + if allocation is None or allocation.tensor is not tensor: + raise ValueError("tensor was not allocated by this manager or has already been freed") + + del self._allocations[data_ptr] + # 归还的是 reserved_bytes,而不是 tensor 的实际字节数,确保对齐填充也被回收。 + self._free_blocks.append(_FreeBlock(allocation.offset, allocation.reserved_bytes)) + self._merge_adjacent_free_blocks() + + def _take_free_block(self, required_bytes: int) -> int: + """使用 first-fit 策略取出一个连续空闲区间,并返回其起始 offset。""" + for index, block in enumerate(self._free_blocks): + if block.size < required_bytes: + continue + + allocation_offset = block.offset + # 大小完全相同则删除空闲块,否则从空闲块头部切出所需空间。 + if block.size == required_bytes: + del self._free_blocks[index] + else: + block.offset += required_bytes + block.size -= required_bytes + return allocation_offset + + # 总空闲空间足够但最大连续块不足时,也会在这里报告碎片化信息。 + free_bytes = sum(block.size for block in self._free_blocks) + largest_block = max((block.size for block in self._free_blocks), default=0) + raise MemoryError( + f"tensor buffer has no contiguous block for {required_bytes} bytes; " + f"free={free_bytes} bytes, largest_free_block={largest_block} bytes" + ) + + def _merge_adjacent_free_blocks(self) -> None: + """按 offset 排序空闲块,并将地址上相邻的区间合并。""" + self._free_blocks.sort(key=lambda block: block.offset) + if len(self._free_blocks) < 2: + return + + merged_blocks = [self._free_blocks[0]] + for current_block in self._free_blocks[1:]: + previous_block = merged_blocks[-1] + + # 相邻区间直接合并;存在间隔则保留;发生重叠说明内部状态已损坏。 + if previous_block.end == current_block.offset: + previous_block.size += current_block.size + elif previous_block.end < current_block.offset: + merged_blocks.append(current_block) + else: + raise RuntimeError("tensor buffer contains overlapping free blocks") + + self._free_blocks = merged_blocks + + def _remember_allocation(self, tensor: torch.Tensor, offset: int, reserved_bytes: int) -> None: + """以 tensor 地址为 key,记录其原始对象和对应的内存池区间。""" + self._allocations[tensor.data_ptr()] = _Allocation(tensor, offset, reserved_bytes) + + @staticmethod + def _normalize_shape(shape: Shape) -> torch.Size: + """将整数或整数序列统一转换成不包含负数维度的 ``torch.Size``。""" + if isinstance(shape, int): + shape = (shape,) + try: + shape = torch.Size(shape) + except TypeError as exc: + raise TypeError("shape must be an int or an iterable of ints") from exc + if any(dim < 0 for dim in shape): + raise ValueError("shape dimensions must be non-negative") + return shape + + @staticmethod + def _element_size(dtype: torch.dtype) -> int: + """校验 dtype,并返回该 dtype 单个元素占用的字节数。""" + if not isinstance(dtype, torch.dtype): + raise TypeError("dtype must be a torch.dtype") + return torch.empty((), dtype=dtype).element_size() + + def _align_up(self, size_bytes: int) -> int: + """将字节数向上取整到 ``alignment_bytes`` 的整数倍。""" + return (size_bytes + self._alignment_bytes - 1) // self._alignment_bytes * self._alignment_bytes + + +@dataclass +class _FreeBlock: + """描述内存池中的一段连续空闲字节区间。""" + + offset: int + size: int + + @property + def end(self) -> int: + """返回空闲区间的右边界,不包含该位置。""" + return self.offset + self.size + + +@dataclass(frozen=True) +class _Allocation: + """记录已申请 tensor 及其在内存池中的完整预留区间。""" + + tensor: torch.Tensor + offset: int + reserved_bytes: int diff --git a/test/utils/test_tensor_buffer_manager.py b/test/utils/test_tensor_buffer_manager.py new file mode 100644 index 0000000000..67ec102445 --- /dev/null +++ b/test/utils/test_tensor_buffer_manager.py @@ -0,0 +1,133 @@ +import pytest +import torch + +from lightllm.utils.tensor_buffer_manager import TensorBufferManager + + +def _aligned_buffer(size: int) -> torch.Tensor: + storage = torch.empty(size + 255, dtype=torch.uint8) + aligned_offset = -storage.data_ptr() % 256 + return storage[aligned_offset : aligned_offset + size] + + +def test_allocate_different_shapes_and_dtypes_from_one_buffer(): + buffer = torch.empty(1024, dtype=torch.uint8) + manager = TensorBufferManager(buffer) + + fp32_tensor = manager.alloc((3, 5), torch.float32) + int16_tensor = manager.alloc((7,), torch.int16) + + assert fp32_tensor.shape == (3, 5) + assert fp32_tensor.dtype == torch.float32 + assert fp32_tensor.is_contiguous() + assert int16_tensor.shape == (7,) + assert int16_tensor.dtype == torch.int16 + assert fp32_tensor.untyped_storage().data_ptr() == buffer.untyped_storage().data_ptr() + assert int16_tensor.untyped_storage().data_ptr() == buffer.untyped_storage().data_ptr() + assert fp32_tensor.data_ptr() % 256 == 0 + assert int16_tensor.data_ptr() % 256 == 0 + + +def test_non_byte_backing_tensor_is_used_as_byte_storage(): + buffer = torch.empty(256, dtype=torch.float32) + manager = TensorBufferManager(buffer) + tensor = manager.alloc((96,), torch.int64) + + assert tensor.untyped_storage().data_ptr() == buffer.untyped_storage().data_ptr() + assert tensor.nbytes == 96 * torch.int64.itemsize + + +def test_released_block_is_reused(): + manager = TensorBufferManager(_aligned_buffer(512)) + first = manager.alloc((16,), torch.float32) + first_ptr = first.data_ptr() + + manager.free(first) + replacement = manager.alloc((8, 2), torch.float32) + + assert replacement.data_ptr() == first_ptr + + +def test_adjacent_free_blocks_are_merged(): + manager = TensorBufferManager(_aligned_buffer(1024)) + first = manager.alloc((32,), torch.uint8) + second = manager.alloc((32,), torch.uint8) + third = manager.alloc((32,), torch.uint8) + first_ptr = first.data_ptr() + + manager.free(second) + manager.free(first) + merged = manager.alloc((512,), torch.uint8) + + assert merged.data_ptr() == first_ptr + manager.free(merged) + manager.free(third) + assert manager.alloc((768,), torch.uint8).data_ptr() == first_ptr + + +def test_release_rejects_unknown_and_already_released_tensors(): + manager = TensorBufferManager(_aligned_buffer(256)) + tensor = manager.alloc((8,), torch.float32) + manager.free(tensor) + + with pytest.raises(ValueError, match="already been freed"): + manager.free(tensor) + with pytest.raises(ValueError, match="not allocated"): + manager.free(torch.empty(1)) + + +def test_stale_tensor_cannot_release_reused_address(): + manager = TensorBufferManager(_aligned_buffer(256)) + stale_tensor = manager.alloc((8,), torch.float32) + manager.free(stale_tensor) + current_tensor = manager.alloc((8,), torch.float32) + + assert stale_tensor.data_ptr() == current_tensor.data_ptr() + with pytest.raises(ValueError, match="already been freed"): + manager.free(stale_tensor) + + manager.free(current_tensor) + + +def test_allocation_reports_fragmentation_on_failure(): + manager = TensorBufferManager(_aligned_buffer(1024)) + first = manager.alloc((256,), torch.uint8) + middle = manager.alloc((512,), torch.uint8) + last = manager.alloc((256,), torch.uint8) + manager.free(first) + manager.free(last) + + with pytest.raises(MemoryError, match="largest_free_block=256 bytes"): + manager.alloc((384,), torch.uint8) + + manager.free(middle) + + +def test_invalid_buffer_is_rejected(): + with pytest.raises(ValueError, match="contiguous"): + TensorBufferManager(torch.empty((4, 4)).t()) + with pytest.raises(ValueError, match="positive power of two"): + TensorBufferManager(torch.empty(256, dtype=torch.uint8), alignment_bytes=0) + with pytest.raises(ValueError, match="positive power of two"): + TensorBufferManager(torch.empty(256, dtype=torch.uint8), alignment_bytes=3) + with pytest.raises(AssertionError, match="no usable bytes"): + TensorBufferManager(torch.empty(0, dtype=torch.uint8)) + + +def test_unaligned_buffer_prefix_is_skipped(): + buffer = torch.empty(512, dtype=torch.uint8)[1:] + manager = TensorBufferManager(buffer) + tensor = manager.alloc((128,), torch.uint8) + + assert tensor.data_ptr() % 256 == 0 + assert tensor.untyped_storage().data_ptr() == buffer.untyped_storage().data_ptr() + + +def test_empty_tensor_does_not_consume_buffer_space(): + manager = TensorBufferManager(_aligned_buffer(512)) + tensor = manager.alloc((0, 4), torch.float16) + full_buffer = manager.alloc((256,), torch.uint8) + + assert tensor.shape == (0, 4) + manager.free(tensor) + manager.free(full_buffer) From 06f9b17768b4a25ec60b64a1ea89c90b25e506f6 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 5 Aug 2026 10:52:07 +0000 Subject: [PATCH 12/20] refactor: manage DeepEP MoE workspace allocations --- .../fused_moe/impl/deepgemm_impl.py | 48 ++-- .../fused_moe/grouped_fused_moe_ep.py | 219 ++++++++++++------ 2 files changed, 185 insertions(+), 82 deletions(-) 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 7976f0bcef..380f0e98f3 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 @@ -14,6 +14,8 @@ chunked_expanded_moe_forward, quantize_fused_experts_input, ) +from lightllm.common.basemodel.triton_kernel.fused_moe.moe_silu_and_mul import silu_and_mul_fwd +from lightllm.common.triton_utils.autotuner import Autotuner from lightllm.common.basemodel.triton_kernel.redundancy_topk_ids_repair import redundancy_topk_ids_repair @@ -222,20 +224,38 @@ 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 = chunked_expanded_moe_forward( - 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, - dist_group_manager.get_deep_ep_prefill_moe_workspace(microbatch_index), - hidden_dtype, - ) + all_tokens = sum(num_recv_tokens_per_expert_list) + if all_tokens > 0: + gather_out = chunked_expanded_moe_forward( + 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, + dist_group_manager.get_deep_ep_prefill_moe_workspace(microbatch_index), + hidden_dtype, + ) + else: + gather_out = torch.empty( + (recv_src_metadata.shape[0], w2_weight.shape[1]), + device=recv_x[0].device, + dtype=hidden_dtype, + ) + ######################################## warning ################################################## + # A rank may receive no tokens during autotune warmup. Run one dummy token through + # silu_and_mul_fwd so the empty rank matches the first kernel call made by non-empty ranks. + # This branch does not synchronize additional calls caused by different positive chunk counts. + if Autotuner.is_autotune_warmup(): + N = w13_weight.shape[1] + _gemm_out_a = torch.zeros((1, N), device=recv_x[0].device, dtype=hidden_dtype) + _silu_out = torch.zeros((1, N // 2), device=recv_x[0].device, dtype=hidden_dtype) + silu_and_mul_fwd(_gemm_out_a.view(-1, N), _silu_out) + _gemm_out_a, _silu_out = None, None del recv_x return gather_out 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 55c06c537b..7bd3643445 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 @@ -25,6 +25,8 @@ ) from lightllm.common.triton_utils.autotuner import Autotuner from lightllm.utils.device_utils import is_sm100_gpu +from lightllm.utils.sgl_utils import HAS_SGL_KERNEL +from lightllm.utils.tensor_buffer_manager import TensorBufferManager logger = init_logger(__name__) _MEGA_MOE_STATES: Dict[Tuple[int, int, int, int], Dict[str, Any]] = {} @@ -292,20 +294,38 @@ def fused_experts_impl( # needed once the received tensors have been produced. del qinput_tensor, input_scale - gather_out = chunked_expanded_moe_forward( - 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, - dist_group_manager.get_deep_ep_prefill_moe_workspace(), - hidden_states.dtype, - ) + all_tokens = sum(handle.num_recv_tokens_per_expert_list) + if all_tokens > 0: + gather_out = chunked_expanded_moe_forward( + 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, + dist_group_manager.get_deep_ep_prefill_moe_workspace(), + hidden_states.dtype, + ) + else: + gather_out = torch.empty( + (handle.recv_src_metadata.shape[0], w2.shape[1]), + device=recv_x[0].device, + dtype=hidden_states.dtype, + ) + ######################################## warning ################################################## + # A rank may receive no tokens during autotune warmup. Run one dummy token through + # silu_and_mul_fwd so the empty rank matches the first kernel call made by non-empty ranks. + # This branch does not synchronize additional calls caused by different positive chunk counts. + if Autotuner.is_autotune_warmup(): + N = w1.shape[1] + _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 del recv_x # normal combine @@ -352,6 +372,87 @@ 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_max_chunk_rows( + workspace: torch.Tensor, + gather_rows: int, + hidden_size: int, + intermediate_size: int, + intermediate_twice: int, + scale_cols: int, + hidden_dtype: torch.dtype, + quant_dtype: torch.dtype, + expert_alignment: int, +) -> int: + """计算并缓存当前 workspace 配置能够容纳的最大 chunk 行数。""" + if not hasattr(_get_max_chunk_rows, "cache"): + _get_max_chunk_rows.cache = {} + max_chunk_rows_cache = _get_max_chunk_rows.cache + + # 同一 1024 行区间共用一个缓存项,并按区间上界探测,保证复用结果不会高估可用空间。 + cached_gather_rows = (gather_rows + 1023) // 1024 * 1024 + cache_key = ( + workspace.numel(), + workspace.device, + cached_gather_rows, + hidden_size, + intermediate_size, + intermediate_twice, + scale_cols, + hidden_dtype, + quant_dtype, + expert_alignment, + HAS_SGL_KERNEL, + ) + if cache_key in max_chunk_rows_cache: + return max_chunk_rows_cache[cache_key] + + def can_allocate(chunk_rows: int) -> bool: + """按实际计算阶段的生命周期申请 buffer,探测该 chunk 是否能够执行。""" + try: + probe_manager = TensorBufferManager(workspace) + probe_manager.alloc((cached_gather_rows, hidden_size), hidden_dtype) + + # W1 阶段同时保存 GEMM 输出和 SwiGLU 输出。 + silu_out = probe_manager.alloc((chunk_rows, intermediate_size), hidden_dtype) + gemm_out_a = probe_manager.alloc((chunk_rows, intermediate_twice), hidden_dtype) + probe_manager.free(gemm_out_a) + + # 量化阶段复用 W1 输出空间,并继续保留 SwiGLU 输出。 + probe_manager.alloc((chunk_rows, intermediate_size), quant_dtype) + aligned_chunk_rows = (chunk_rows + 3) // 4 * 4 + if HAS_SGL_KERNEL: + probe_manager.alloc((scale_cols, aligned_chunk_rows), torch.float32) + else: + # LightLLM fallback 通过 alloc_func 申请 row-major scale;后续 TMA 转置不占用 workspace。 + probe_manager.alloc((chunk_rows, scale_cols), torch.float32) + probe_manager.free(silu_out) + + # W2 阶段释放 SwiGLU 输出后申请最终 GEMM 输出。 + probe_manager.alloc((chunk_rows, hidden_size), hidden_dtype) + except MemoryError: + return False + return True + + # W1 阶段必须同时保存 gemm_out_a 和 silu_out,可据此得到 chunk 数量的绝对上界。 + w1_row_bytes = (intermediate_twice + intermediate_size) * hidden_dtype.itemsize + left = 1 + right = workspace.numel() // w1_row_bytes // expert_alignment + max_chunk_rows = 0 + + while left <= right: + chunk_count = (left + right) // 2 + chunk_rows = chunk_count * expert_alignment + if can_allocate(chunk_rows): + max_chunk_rows = chunk_rows + left = chunk_count + 1 + else: + right = chunk_count - 1 + + max_chunk_rows_cache[cache_key] = max_chunk_rows + logger.info("cache DeepEP max_chunk_rows: key=%s, max_chunk_rows=%s", cache_key, max_chunk_rows) + return max_chunk_rows + + 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 @@ -373,13 +474,8 @@ def chunked_expanded_moe_forward( 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 all_tokens > 0, "chunked_expanded_moe_forward requires non-empty input" 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) - 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) # 与 m_indices 一一对应:0 表示真实 token,1 表示 expert 对齐产生的 padding 行。 @@ -395,50 +491,35 @@ def chunked_expanded_moe_forward( del padding_mask 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 - 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] - 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_row_bytes, - gemm_b_row_bytes + quant_row_bytes, + max_chunk_rows = _get_max_chunk_rows( + workspace=workspace, + gather_rows=gather_rows, + hidden_size=hidden_size, + intermediate_size=intermediate_size, + intermediate_twice=intermediate_twice, + scale_cols=scale_cols, + hidden_dtype=hidden_dtype, + quant_dtype=w2.dtype, + expert_alignment=alignment, ) - 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 + + if max_chunk_rows == 0: raise RuntimeError( - f"DeepEP workspace needs at least {minimum_bytes} bytes " - f"({gather_bytes} dense + {alignment * temp_row_bytes} temporary), have {workspace.numel()} bytes" + f"DeepEP workspace with {workspace.numel()} bytes cannot hold the dense output and " + f"one {alignment}-row temporary chunk" ) + max_chunk_rows = min(all_tokens, max_chunk_rows) - gather_out = workspace[:gather_bytes].view(hidden_dtype).view(gather_rows, hidden_size) + workspace_manager = TensorBufferManager(workspace) + gather_out = workspace_manager.alloc((gather_rows, hidden_size), hidden_dtype) 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 - 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[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) + silu_out = workspace_manager.alloc((chunk_rows, intermediate_size), hidden_dtype) + gemm_out_a = workspace_manager.alloc((chunk_rows, intermediate_twice), hidden_dtype) deepgemm_grouped_fp8_nt_contiguous( (recv_x[0][chunk_start:chunk_end], recv_x[1][chunk_start:chunk_end]), (w1, w1_scale), @@ -446,20 +527,17 @@ def chunked_expanded_moe_forward( m_indices[chunk_start:chunk_end], ) silu_and_mul_fwd(gemm_out_a, silu_out) + workspace_manager.free(gemm_out_a) del gemm_out_a - 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) + quant_buffers = [] 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}") + if device != workspace.device: + raise RuntimeError(f"quant buffer must be allocated on {workspace.device}, got {device}") + quant_buffer = workspace_manager.alloc(shape, dtype) + quant_buffers.append(quant_buffer) + return quant_buffer qsilu_out, qsilu_out_scale = per_token_group_quant_fp8( silu_out, @@ -469,17 +547,22 @@ def workspace_quant_alloc(shape, dtype, device): scale_tma_aligned=True, alloc_func=workspace_quant_alloc, ) - gemm_out_b = ( - workspace[temp_offset : temp_offset + gemm_b_bytes].view(hidden_dtype).view(chunk_rows, hidden_size) - ) + workspace_manager.free(silu_out) + del silu_out + + gemm_out_b = workspace_manager.alloc((chunk_rows, hidden_size), hidden_dtype) 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 + del qsilu_out, qsilu_out_scale + for quant_buffer in quant_buffers: + workspace_manager.free(quant_buffer) + ep_gather_chunk(gemm_out_b, chunk_start, recv_topk_weights, recv_src_metadata, gather_out) + workspace_manager.free(gemm_out_b) ep_compact_metadata(recv_src_metadata) return gather_out From 114c8de87c758e6b091677c62c5b0d171a424da4 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 5 Aug 2026 10:58:20 +0000 Subject: [PATCH 13/20] fix: align DeepEP autotune chunk calls --- .../fused_moe/grouped_fused_moe_ep.py | 106 ++++++++++-------- 1 file changed, 59 insertions(+), 47 deletions(-) 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 7bd3643445..e276bfbab2 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 @@ -515,54 +515,66 @@ def chunked_expanded_moe_forward( gather_out = workspace_manager.alloc((gather_rows, hidden_size), hidden_dtype) gather_out.zero_() - 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_out = workspace_manager.alloc((chunk_rows, intermediate_size), hidden_dtype) - gemm_out_a = workspace_manager.alloc((chunk_rows, intermediate_twice), hidden_dtype) - 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) - workspace_manager.free(gemm_out_a) - del gemm_out_a - - quant_buffers = [] - - def workspace_quant_alloc(shape, dtype, device): - if device != workspace.device: - raise RuntimeError(f"quant buffer must be allocated on {workspace.device}, got {device}") - quant_buffer = workspace_manager.alloc(shape, dtype) - quant_buffers.append(quant_buffer) - return quant_buffer - - 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, - ) - workspace_manager.free(silu_out) - del silu_out - - gemm_out_b = workspace_manager.alloc((chunk_rows, hidden_size), hidden_dtype) - 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 - for quant_buffer in quant_buffers: - workspace_manager.free(quant_buffer) + # 不同 rank 接收到的 token 数不同,因此实际 chunk 数也可能不同。Autotuner warmup + # 中的分布式通信要求各 rank 进入 autotuning 的次数一致,否则容易发生通信错位。 + # 所以只允许第一个 chunk 保持 autotuning;从第二个 chunk 开始临时关闭,循环结束 + # 后再恢复进入函数时的 warmup 状态。零 token rank 的首次调用由外层特殊分支补齐。 + is_autotune_warmup = Autotuner.is_autotune_warmup() + try: + for chunk_index, chunk_start in enumerate(range(0, all_tokens, max_chunk_rows)): + if is_autotune_warmup and chunk_index == 1: + Autotuner.end_autotune_warmup() + + chunk_end = min(chunk_start + max_chunk_rows, all_tokens) + chunk_rows = chunk_end - chunk_start + silu_out = workspace_manager.alloc((chunk_rows, intermediate_size), hidden_dtype) + gemm_out_a = workspace_manager.alloc((chunk_rows, intermediate_twice), hidden_dtype) + 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) + workspace_manager.free(gemm_out_a) + del gemm_out_a + + quant_buffers = [] - ep_gather_chunk(gemm_out_b, chunk_start, recv_topk_weights, recv_src_metadata, gather_out) - workspace_manager.free(gemm_out_b) + def workspace_quant_alloc(shape, dtype, device): + if device != workspace.device: + raise RuntimeError(f"quant buffer must be allocated on {workspace.device}, got {device}") + quant_buffer = workspace_manager.alloc(shape, dtype) + quant_buffers.append(quant_buffer) + return quant_buffer + + 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, + ) + workspace_manager.free(silu_out) + del silu_out + + gemm_out_b = workspace_manager.alloc((chunk_rows, hidden_size), hidden_dtype) + 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 + for quant_buffer in quant_buffers: + workspace_manager.free(quant_buffer) + + ep_gather_chunk(gemm_out_b, chunk_start, recv_topk_weights, recv_src_metadata, gather_out) + workspace_manager.free(gemm_out_b) + finally: + if is_autotune_warmup: + Autotuner.start_autotune_warmup() ep_compact_metadata(recv_src_metadata) return gather_out From 2076102aff32dd9f2985bad19fd2b5d30ab02de1 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 6 Aug 2026 02:20:24 +0000 Subject: [PATCH 14/20] fix: support non-aligned DeepEP gather widths --- .../deepep_expanded_layout_kernels.py | 52 +++++++++++++------ 1 file changed, 35 insertions(+), 17 deletions(-) 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 index c63d091864..573b68087c 100644 --- 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 @@ -166,25 +166,39 @@ def _ep_gather_chunk_kernel( output, output_stride_m, output_stride_k, + hidden_size, TOPK: tl.constexpr, BLOCK_D: tl.constexpr, + NEED_HIDDEN_MASK: 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) + if NEED_HIDDEN_MASK: + hidden_mask = hidden_offsets < hidden_size 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) + if NEED_HIDDEN_MASK: + accumulator = tl.load(output_ptrs, mask=hidden_mask, other=0.0).to(tl.float32) + else: + 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) + chunk_ptrs = chunk + local_row * chunk_stride_m + hidden_offsets * chunk_stride_k + if NEED_HIDDEN_MASK: + value = tl.load(chunk_ptrs, mask=hidden_mask, other=0.0) + else: + value = tl.load(chunk_ptrs) weight = tl.load(weights + slot) accumulator += value.to(tl.float32) * weight - tl.store(output_ptrs, accumulator) + if NEED_HIDDEN_MASK: + tl.store(output_ptrs, accumulator, mask=hidden_mask) + else: + tl.store(output_ptrs, accumulator) @torch.no_grad() @@ -204,24 +218,28 @@ def ep_gather_chunk( """ topk = recv_src_metadata.shape[1] - 2 block_d = 1024 - assert chunk.shape[1] == output.shape[1] and output.shape[1] % block_d == 0 + hidden_size = output.shape[1] + assert chunk.shape[1] == hidden_size 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), + total_recv_tokens=output.shape[0], + chunk=chunk, + chunk_stride_m=chunk.stride(0), + chunk_stride_k=chunk.stride(1), + chunk_start=chunk_start, + chunk_end=chunk_start + chunk.shape[0], + weights=weights, + recv_src_metadata=recv_src_metadata, + metadata_stride_m=recv_src_metadata.stride(0), + metadata_stride_k=recv_src_metadata.stride(1), + output=output, + output_stride_m=output.stride(0), + output_stride_k=output.stride(1), + hidden_size=hidden_size, TOPK=topk, BLOCK_D=block_d, + # 常见整块 hidden size 保持原来的无 mask kernel;仅尾块不完整时启用边界保护。 + NEED_HIDDEN_MASK=hidden_size % block_d != 0, num_warps=2, ) From a017c1e0f78a72d3720d009fb68457be33d04c59 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 6 Aug 2026 02:25:40 +0000 Subject: [PATCH 15/20] refactor: use named arguments for expanded MoE --- .../fused_moe/impl/deepgemm_impl.py | 24 +++++++++---------- .../fused_moe/grouped_fused_moe_ep.py | 24 +++++++++---------- 2 files changed, 24 insertions(+), 24 deletions(-) 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 380f0e98f3..a5ba656c9c 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 @@ -227,18 +227,18 @@ def prefilled_group_gemm( all_tokens = sum(num_recv_tokens_per_expert_list) if all_tokens > 0: gather_out = chunked_expanded_moe_forward( - 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, - dist_group_manager.get_deep_ep_prefill_moe_workspace(microbatch_index), - hidden_dtype, + 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_x=recv_x, + recv_topk_weights=recv_topk_weights, + recv_src_metadata=recv_src_metadata, + w1=w13_weight, + w1_scale=w13_scale, + w2=w2_weight, + w2_scale=w2_scale, + block_size_k=self.quant_method.block_size, + workspace=dist_group_manager.get_deep_ep_prefill_moe_workspace(microbatch_index), + hidden_dtype=hidden_dtype, ) else: gather_out = torch.empty( 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 e276bfbab2..58d4d45514 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 @@ -297,18 +297,18 @@ def fused_experts_impl( all_tokens = sum(handle.num_recv_tokens_per_expert_list) if all_tokens > 0: gather_out = chunked_expanded_moe_forward( - 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, - dist_group_manager.get_deep_ep_prefill_moe_workspace(), - hidden_states.dtype, + num_recv_tokens_per_expert_list=handle.num_recv_tokens_per_expert_list, + num_unaligned_recv_tokens_per_expert=handle.num_unaligned_recv_tokens_per_expert, + recv_x=recv_x, + recv_topk_weights=recv_topk_weights, + recv_src_metadata=handle.recv_src_metadata, + w1=w1, + w1_scale=w1_scale, + w2=w2, + w2_scale=w2_scale, + block_size_k=block_size_k, + workspace=dist_group_manager.get_deep_ep_prefill_moe_workspace(), + hidden_dtype=hidden_states.dtype, ) else: gather_out = torch.empty( From 8e4a0a8e94e4b15876dc57f616a005a481d5b82e Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 6 Aug 2026 02:52:43 +0000 Subject: [PATCH 16/20] fix: initialize required DeepEP MoE buffers --- lightllm/distributed/communication_op.py | 89 +++++++++++++++++++++--- lightllm/models/deepseek2/model.py | 9 +-- lightllm/models/gemma4/model.py | 9 +-- lightllm/models/glm4_moe_lite/model.py | 9 +-- lightllm/models/qwen3_moe/model.py | 9 +-- 5 files changed, 98 insertions(+), 27 deletions(-) diff --git a/lightllm/distributed/communication_op.py b/lightllm/distributed/communication_op.py index 2eb2ea5a79..d003f3f3a1 100644 --- a/lightllm/distributed/communication_op.py +++ b/lightllm/distributed/communication_op.py @@ -22,7 +22,7 @@ import torch import torch.distributed as dist from torch.distributed import ReduceOp, ProcessGroup -from typing import List, Dict, Optional, Union +from typing import List, Dict, Optional, Set, Union from lightllm.utils.log_utils import init_logger from lightllm.utils.device_utils import has_nvlink from lightllm.utils.envs_utils import ( @@ -132,13 +132,43 @@ def get_default_group(self) -> CustomProcessGroup: def get_group(self, group_index: int) -> CustomProcessGroup: return self.groups[group_index] + @staticmethod + def get_moe_quant_methods(layer_weights: List) -> Set[str]: + """收集实际绑定到各 MoE 层 expert weight 上的量化方法名称。 + + expert 量化类型可能分别来自启动参数、quant_cfg 和模型 config。调用本函数 + 时 layer weights 已经构造完成,每层 ``experts.quant_method`` 保存的是按照 + 既定优先级解析后的最终结果,因此这里不再重复解析配置。 + + 返回方法名称集合是为了去重并支持混合量化。例如部分 MoE 层使用 FP4、 + 其余层使用 FP8 时,可以据此同时初始化两条执行路径所需的 buffer;普通 + dense 层没有 ``experts``,会被自然跳过。 + """ + quant_method_names = set() + for layer_weight in layer_weights: + # dense 层没有 experts;这里只关心真正参与 MoE 计算的层。 + experts = getattr(layer_weight, "experts", None) + quant_method = getattr(experts, "quant_method", None) + method_name = getattr(quant_method, "method_name", None) + if method_name is not None: + quant_method_names.add(method_name) + return quant_method_names + def new_deepep_group( self, n_routed_experts, hidden_size, + expert_quant_method_names: Set[str], num_experts_per_tok: int = 1, moe_intermediate_size: Optional[int] = None, ): + """初始化 DeepEP 通信组以及当前模型实际需要的 MoE buffer。 + + ``expert_quant_method_names`` 是各 MoE 层最终绑定的 quant method 名称集合。 + 同一个模型可能逐层混用 FP4 和 FP8:SM100 FP4 层走 Mega MoE,其他层走 + DeepEP legacy low-latency 路径。这里只为实际存在的执行路径分配 buffer, + 避免为未使用的路径长期占用显存。 + """ enable_ep_moe = get_env_start_args().enable_ep_moe prefill_num_max_dispatch_tokens_per_rank = get_deepep_num_max_dispatch_tokens_per_rank_prefill() decode_num_max_dispatch_tokens_per_rank = get_deepep_num_max_dispatch_tokens_per_rank_decode() @@ -173,16 +203,47 @@ def new_deepep_group( ) self.ep_mega_moe_buffer = None self.ep_low_latency_buffer = None - 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), - ) - if is_sm100_gpu(): + + if not expert_quant_method_names: + raise ValueError("No valid MoE quant method was found while initializing DeepEP buffers") + + mega_moe_quant_method = "fp4fp8-b32-deepgemm" + is_sm100 = is_sm100_gpu() + + # Buffer 选择规则: + # 1. 非 SM100 不支持 Mega MoE,只初始化 legacy low-latency buffer; + # 2. SM100 全部 MoE 层为 FP4,只初始化 Mega MoE buffer; + # 3. SM100 全部 MoE 层为 FP8,只初始化 legacy low-latency buffer; + # 4. SM100 逐层混合 FP4/FP8,两套 buffer 都要初始化。 + if is_sm100: + # 只要存在一个 FP4 MoE 层,就需要 Mega MoE buffer;只要存在一个非 FP4 + # MoE 层,就需要 legacy low-latency buffer。FP4/FP8 逐层混用时两者都会初始化。 + has_mega_moe_layer = mega_moe_quant_method in expert_quant_method_names + has_legacy_moe_layer = any( + method_name != mega_moe_quant_method for method_name in expert_quant_method_names + ) + enable_mega_moe_buffer = has_mega_moe_layer + enable_low_latency_buffer = has_legacy_moe_layer + else: + enable_mega_moe_buffer = False + enable_low_latency_buffer = True + + if enable_low_latency_buffer: + # FP8 MoE 的 decode 使用 legacy low-latency buffer;prefill 阶段还会将其 + # 空闲的本地 RDMA storage 复用为分块 grouped GEMM 的临时 workspace。 + 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), + ) + + if enable_mega_moe_buffer: + # SM100 FP4 层通过 DeepGEMM Mega MoE 完成通信和计算,不使用 legacy + # low-latency buffer,因此纯 FP4 模型无需承担后者的大块 RDMA 显存。 if moe_intermediate_size is None: raise ValueError("SM100 Mega MoE requires moe_intermediate_size or intermediate_size in model config") @@ -196,6 +257,12 @@ def new_deepep_group( self.ll_hidden, moe_intermediate_size, ) + logger.info( + "Initialize DeepEP MoE buffers: low_latency=%s, mega_moe=%s, expert_quant_method_names=%s", + enable_low_latency_buffer, + enable_mega_moe_buffer, + sorted(expert_quant_method_names), + ) theoretical_sms = self.ep_buffer.get_theoretical_num_sms(self.ll_num_experts, num_experts_per_tok) self._set_num_sms_for_deep_gemm(theoretical_sms) diff --git a/lightllm/models/deepseek2/model.py b/lightllm/models/deepseek2/model.py index ea6620b4e4..29f209f0e8 100644 --- a/lightllm/models/deepseek2/model.py +++ b/lightllm/models/deepseek2/model.py @@ -49,10 +49,11 @@ def _init_some_value(self): def _init_custom(self): self._init_to_get_yarn_rotary() dist_group_manager.new_deepep_group( - self.config["n_routed_experts"], - self.config["hidden_size"], - self.config.get("num_experts_per_tok", 1), - self.config.get("moe_intermediate_size", self.config.get("intermediate_size")), + n_routed_experts=self.config["n_routed_experts"], + hidden_size=self.config["hidden_size"], + expert_quant_method_names=dist_group_manager.get_moe_quant_methods(self.trans_layers_weight), + num_experts_per_tok=self.config.get("num_experts_per_tok", 1), + moe_intermediate_size=self.config.get("moe_intermediate_size", self.config.get("intermediate_size")), ) def _verify_params(self): diff --git a/lightllm/models/gemma4/model.py b/lightllm/models/gemma4/model.py index 10b1958b0e..061c135b4c 100644 --- a/lightllm/models/gemma4/model.py +++ b/lightllm/models/gemma4/model.py @@ -131,10 +131,11 @@ def _init_custom(self): self._init_to_get_rotary_gemma4() if self.config.get("enable_moe_block", False): dist_group_manager.new_deepep_group( - self.config["num_experts"], - self.config["hidden_size"], - self.config.get("num_experts_per_tok", self.config.get("top_k_experts", 1)), - self.config.get("moe_intermediate_size", self.config.get("intermediate_size")), + n_routed_experts=self.config["num_experts"], + hidden_size=self.config["hidden_size"], + expert_quant_method_names=dist_group_manager.get_moe_quant_methods(self.trans_layers_weight), + num_experts_per_tok=self.config.get("num_experts_per_tok", self.config.get("top_k_experts", 1)), + moe_intermediate_size=self.config.get("moe_intermediate_size", self.config.get("intermediate_size")), ) self._init_ple_static_buffer() diff --git a/lightllm/models/glm4_moe_lite/model.py b/lightllm/models/glm4_moe_lite/model.py index 1e31306aea..1866e0de62 100644 --- a/lightllm/models/glm4_moe_lite/model.py +++ b/lightllm/models/glm4_moe_lite/model.py @@ -26,10 +26,11 @@ def _init_config(self): def _init_custom(self): self._init_to_get_yarn_rotary() dist_group_manager.new_deepep_group( - self.config["n_routed_experts"], - self.config["hidden_size"], - self.config.get("num_experts_per_tok", 1), - self.config.get("moe_intermediate_size", self.config.get("intermediate_size")), + n_routed_experts=self.config["n_routed_experts"], + hidden_size=self.config["hidden_size"], + expert_quant_method_names=dist_group_manager.get_moe_quant_methods(self.trans_layers_weight), + num_experts_per_tok=self.config.get("num_experts_per_tok", 1), + moe_intermediate_size=self.config.get("moe_intermediate_size", self.config.get("intermediate_size")), ) def _init_to_get_yarn_rotary(self): diff --git a/lightllm/models/qwen3_moe/model.py b/lightllm/models/qwen3_moe/model.py index 0d4b45bfe6..5d91755137 100644 --- a/lightllm/models/qwen3_moe/model.py +++ b/lightllm/models/qwen3_moe/model.py @@ -28,8 +28,9 @@ def _init_custom(self): # Only initialize DeepEP group for MoE models with num_experts if "num_experts" in self.config and self.config["num_experts"] > 0: dist_group_manager.new_deepep_group( - self.config["num_experts"], - self.config["hidden_size"], - self.config.get("num_experts_per_tok", 1), - self.config.get("moe_intermediate_size", self.config.get("intermediate_size")), + n_routed_experts=self.config["num_experts"], + hidden_size=self.config["hidden_size"], + expert_quant_method_names=dist_group_manager.get_moe_quant_methods(self.trans_layers_weight), + num_experts_per_tok=self.config.get("num_experts_per_tok", 1), + moe_intermediate_size=self.config.get("moe_intermediate_size", self.config.get("intermediate_size")), ) From b5d68ac7c36d2de49e5cb45e5c4b245bf404f827 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 6 Aug 2026 03:22:27 +0000 Subject: [PATCH 17/20] fix --- .../triton_kernel/fused_moe/deepep_expanded_layout_kernels.py | 3 +++ 1 file changed, 3 insertions(+) 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 index 573b68087c..62e36eb31b 100644 --- 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 @@ -151,6 +151,9 @@ def ep_zero_padding( ) +# TODO: 当前实现会为每个 chunk 重新遍历所有接收 token 的 top-k metadata,并反复 +# 读取、累加和写回长期保留的 gather_out,仍有较大的性能提升空间。后续可以考虑 +# 预先构建 chunk-local 的路由映射,减少无效 metadata 扫描和全局内存读写。 @triton.jit def _ep_gather_chunk_kernel( total_recv_tokens, From 470066d5375ca3535c9761fe0226592b813f9612 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 6 Aug 2026 06:22:07 +0000 Subject: [PATCH 18/20] perf: batch DeepEP metadata compaction --- .../deepep_expanded_layout_kernels.py | 34 ++++++++++++------- 1 file changed, 21 insertions(+), 13 deletions(-) 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 index 62e36eb31b..ab54ea5b79 100644 --- 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 @@ -252,17 +252,21 @@ def _ep_compact_metadata_kernel( recv_src_metadata, metadata_stride_m, metadata_stride_k, + num_recv_tokens, TOPK: tl.constexpr, + BLOCK_TOKEN: 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, - ) + token_offsets = tl.program_id(0) * BLOCK_TOKEN + tl.arange(0, BLOCK_TOKEN) + topk_offsets = tl.arange(0, BLOCK_TOPK) + metadata_offsets = token_offsets[:, None] * metadata_stride_m + (topk_offsets[None, :] + 2) * metadata_stride_k + valid_token_mask = token_offsets < num_recv_tokens + valid_topk_mask = topk_offsets < TOPK + metadata_mask = valid_token_mask[:, None] & valid_topk_mask[None, :] + + # 每个 token 只保留其在稠密输出中的同序行号,其余 top-k 位置全部置为无效。 + slots = tl.where(topk_offsets[None, :] == 0, token_offsets[:, None], -1) + tl.store(recv_src_metadata + metadata_offsets, slots, mask=metadata_mask) @torch.no_grad() @@ -278,11 +282,15 @@ def ep_compact_metadata( 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), + block_token = 128 + grid = (triton.cdiv(recv_src_metadata.shape[0], block_token),) + _ep_compact_metadata_kernel[grid]( + recv_src_metadata=recv_src_metadata, + metadata_stride_m=recv_src_metadata.stride(0), + metadata_stride_k=recv_src_metadata.stride(1), + num_recv_tokens=recv_src_metadata.shape[0], TOPK=topk, + BLOCK_TOKEN=block_token, BLOCK_TOPK=triton.next_power_of_2(topk), - num_warps=1, + num_warps=4, ) From 5c0c0e5fea21d0db2757af9d43509edfdf29ea38 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 6 Aug 2026 06:57:03 +0000 Subject: [PATCH 19/20] docs: note potential NVLink buffer tuning --- lightllm/distributed/communication_op.py | 1 + 1 file changed, 1 insertion(+) diff --git a/lightllm/distributed/communication_op.py b/lightllm/distributed/communication_op.py index d003f3f3a1..3864601f05 100644 --- a/lightllm/distributed/communication_op.py +++ b/lightllm/distributed/communication_op.py @@ -234,6 +234,7 @@ def new_deepep_group( 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 ) + # TODO: 评估是否应在部分场景下设置 num_nvl_bytes,以提升单机性能。 self.ep_low_latency_buffer = deep_ep.Buffer( deepep_group, num_rdma_bytes=num_rdma_bytes, From be5e8c019264de0565e422129ac7797031494c95 Mon Sep 17 00:00:00 2001 From: niushengxiao Date: Thu, 6 Aug 2026 18:31:58 +0800 Subject: [PATCH 20/20] remove todo --- lightllm/distributed/communication_op.py | 1 - 1 file changed, 1 deletion(-) diff --git a/lightllm/distributed/communication_op.py b/lightllm/distributed/communication_op.py index 3864601f05..d003f3f3a1 100644 --- a/lightllm/distributed/communication_op.py +++ b/lightllm/distributed/communication_op.py @@ -234,7 +234,6 @@ def new_deepep_group( 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 ) - # TODO: 评估是否应在部分场景下设置 num_nvl_bytes,以提升单机性能。 self.ep_low_latency_buffer = deep_ep.Buffer( deepep_group, num_rdma_bytes=num_rdma_bytes,