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/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..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 @@ -226,20 +226,26 @@ 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, + 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( 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, + 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 cfc82facee..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 @@ -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,11 @@ fused_experts, get_ep_num_sms, masked_group_gemm, - deepgemm_grouped_fp8_nt_contiguous, + chunked_expanded_moe_forward, 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.triton_utils.autotuner import Autotuner from lightllm.common.basemodel.triton_kernel.redundancy_topk_ids_repair import redundancy_topk_ids_repair @@ -182,6 +177,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 +211,52 @@ 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, + microbatch_index: int = 0, ): - 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) + assert recv_topk_idx is None + all_tokens = sum(num_recv_tokens_per_expert_list) 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, + 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, ) - 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: + gather_out = torch.empty( + (recv_src_metadata.shape[0], w2_weight.shape[1]), + device=recv_x[0].device, + dtype=hidden_dtype, + ) ######################################## 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. + # 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(): - _gemm_out_a = torch.zeros((1, N), device=device, dtype=hidden_dtype) - _silu_out = torch.zeros((1, N // 2), device=device, dtype=hidden_dtype) + 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 def low_latency_combine( 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..c63d091864 --- /dev/null +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/deepep_expanded_layout_kernels.py @@ -0,0 +1,267 @@ +import torch +import triton +import triton.language as tl + + +@triton.jit +def _ep_build_m_indices_kernel( + num_unaligned_recv_tokens_per_expert, + m_indices, + padding_mask, + 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 = 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_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_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 aligned expert layout used by contiguous grouped GEMM. + + 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. + + ``padding_mask`` is filled in-place with ``1`` for alignment-padding rows + and ``0`` for real token rows. + """ + 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] % expert_alignment == 0 + assert padding_mask.dtype == torch.int32 and padding_mask.shape == m_indices.shape + + _ep_build_m_indices_kernel[(num_experts,)]( + num_unaligned_recv_tokens_per_expert, + m_indices, + padding_mask, + num_experts=num_experts, + num_warps=8, + BLOCK_E=expert_alignment, + BLOCK_EXPERT_NUM=triton.next_power_of_2(num_experts), + ) + + +@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, + padding_mask, + hidden_size: tl.constexpr, + scale_hidden_size: tl.constexpr, + BLOCK_M: tl.constexpr, + BLOCK_K: tl.constexpr, + BLOCK_SCALE_K: tl.constexpr, +): + 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_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_mask[None, :]) + 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] + padding_mask: torch.Tensor, # [num_expanded_tokens], 1 for padding rows +): + """Zero the alignment-padding rows in DeepEP's expanded receive layout. + + 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 = ( + triton.cdiv(padding_mask.shape[0], 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, + padding_mask, + 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 101d316937..0000000000 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/deepep_scatter_gather.py +++ /dev/null @@ -1,232 +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, -): - 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) - cur_expert_token_num = tl.load(num_recv_tokens_per_expert + cur_expert) - - 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), - ) - - 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 - - -@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 4671329840..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 @@ -1,7 +1,9 @@ """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,20 +12,27 @@ ) 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_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, get_deepep_num_max_dispatch_tokens_per_rank_decode, ) 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]] = {} SUPPORTED_EP_EXPERT_DTYPES = ("fp8w8a8-b128-deepgemm", "fp4fp8-b32-deepgemm") + try: from deep_ep import Buffer, EventOverlap import deep_gemm @@ -75,12 +84,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 +249,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,12 +263,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 - # 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 - 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, @@ -274,74 +287,46 @@ 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 - # 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) + all_tokens = sum(handle.num_recv_tokens_per_expert_list) 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, + 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, ) - - # 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: + gather_out = torch.empty( + (handle.recv_src_metadata.shape[0], w2.shape[1]), + device=recv_x[0].device, + dtype=hidden_states.dtype, + ) ######################################## 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. + # 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 combined_x, _, event = buffer.combine( @@ -387,6 +372,214 @@ 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 + 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, # [workspace_bytes], uint8 + hidden_dtype: torch.dtype, # scalar dtype descriptor +): + """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 all_tokens > 0, "chunked_expanded_moe_forward requires non-empty input" + assert workspace.dtype == torch.uint8 and workspace.ndim == 1 and workspace.is_contiguous() + + m_indices = torch.empty(all_tokens, device=recv_x[0].device, dtype=torch.int32) + # 与 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, + padding_mask, + ) + del padding_mask + + gather_rows = recv_src_metadata.shape[0] + scale_cols = intermediate_size // block_size_k + 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, + ) + + if max_chunk_rows == 0: + raise RuntimeError( + 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) + + workspace_manager = TensorBufferManager(workspace) + gather_out = workspace_manager.alloc((gather_rows, hidden_size), hidden_dtype) + gather_out.zero_() + + # 不同 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 = [] + + 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 + + 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..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) @@ -43,6 +45,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") @@ -51,6 +54,30 @@ def _parse_network_config(self, network_config): self.hf_quantization_method = hf_quantization_config["quant_method"] self._mapping_quant_method() + def _mapping_expert_quant_method(self): + 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): + 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: + # 优先级 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": block_size = self.hf_quantization_config.get("weight_block_size", None) @@ -62,19 +89,8 @@ 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() - # 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..2eb2ea5a79 100644 --- a/lightllm/distributed/communication_op.py +++ b/lightllm/distributed/communication_op.py @@ -169,22 +169,20 @@ 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), + ) + 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") @@ -217,9 +215,40 @@ 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 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( diff --git a/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py index dae79cc8a6..3254031056 100644 --- a/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py @@ -7,7 +7,9 @@ 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 ( + 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 @@ -501,7 +503,13 @@ 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, + microbatch_index=0, ) # 1 dispatch execute @@ -527,7 +535,13 @@ 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, + 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 7edfd5a6f9..7311c4d141 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,9 @@ 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 ( + use_sm100_mega_moe, +) from lightllm.utils.dist_utils import get_global_world_size from lightllm.utils.envs_utils import get_env_start_args @@ -315,7 +317,13 @@ 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, + microbatch_index=0, ) # 1 dispatch execute @@ -341,7 +349,13 @@ 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, + microbatch_index=1, ) # wait 0 combine 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) 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()