diff --git a/lightllm/distributed/communication_op.py b/lightllm/distributed/communication_op.py index d003f3f3a1..93c603212d 100644 --- a/lightllm/distributed/communication_op.py +++ b/lightllm/distributed/communication_op.py @@ -20,6 +20,7 @@ import os import torch +import triton import torch.distributed as dist from torch.distributed import ReduceOp, ProcessGroup from typing import List, Dict, Optional, Set, Union @@ -231,9 +232,24 @@ def new_deepep_group( 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( + decode_size_hint = deep_ep.Buffer.get_low_latency_rdma_size_hint( self.ll_decode_num_tokens, self.ll_hidden, global_world_size, self.ll_num_experts ) + microbatch_count = len(self.groups) + min_prefill_reuse_buffer_bytes = _calculate_min_chunked_expanded_moe_reuse_buffer_bytes( + global_world_size=global_world_size, + prefill_tokens_per_rank=prefill_num_max_dispatch_tokens_per_rank, + hidden_size=hidden_size, + moe_intermediate_size=moe_intermediate_size, + hidden_dtype=get_torch_dtype(get_env_start_args().data_type), + microbatch_count=microbatch_count, + ) + # Decode 和 prefill 不会同时使用 local RDMA storage,容量取两条路径的较大值。 + num_rdma_bytes = max(decode_size_hint, min_prefill_reuse_buffer_bytes) + # DeepEP 返回的 decode hint 不保证能被 microbatch 均分;最终再对齐一次, + # 确保每个 workspace slice 的容量及起始位置仍保持 256-byte 对齐。 + rdma_alignment = microbatch_count * 256 + num_rdma_bytes = triton.cdiv(num_rdma_bytes, rdma_alignment) * rdma_alignment self.ep_low_latency_buffer = deep_ep.Buffer( deepep_group, num_rdma_bytes=num_rdma_bytes, @@ -408,4 +424,63 @@ def _is_single_group(group: Optional[Union[ProcessGroup, CustomProcessGroup]]) - return dist.get_world_size(group=group) == 1 +def _calculate_min_chunked_expanded_moe_reuse_buffer_bytes( + global_world_size: int, + prefill_tokens_per_rank: int, + hidden_size: int, + moe_intermediate_size: int, + hidden_dtype: torch.dtype, + microbatch_count: int, +) -> int: + """计算 ``chunked_expanded_moe_forward`` 极端情况下所需的最小复用 buffer。 + + Prefill 阶段会将空闲的 legacy DeepEP local RDMA storage 复用为 expanded + grouped GEMM 的 workspace。该函数只计算这条 prefill 路径的容量下限,不包含 + low-latency decode 自身所需的 RDMA 容量。 + + 每个 prefill workspace 由全程常驻的 dense gather 输出和分块计算的临时峰值组成。 + gather 按 ``global_world_size * prefill_tokens_per_rank`` 估算最大可能接收行数, + 并向 1024 行对齐;这与运行期 workspace 探测使用的保守分桶保持一致。临时计算 + 按最小 128 行 chunk 估算:W1 阶段同时需要 ``gemm_out_a`` 和 ``silu_out``;后续 + 量化/W2 阶段涉及 ``silu_out``、FP8 量化结果、scale 及 ``gemm_out_b``。这里有意 + 沿用保守公式,避免因生命周期估计不足而低估空间。 + + 多个 microbatch/communication group 并行时,每组独占一个 workspace slice,故 + prefill 容量乘以 ``microbatch_count``。最后按 ``microbatch_count * 256`` 对齐, + 保证均分后每个 slice 仍满足 256-byte 对齐。 + """ + + def align(value: int, alignment: int) -> int: + return triton.cdiv(value, alignment) * alignment + + def tensor_bytes(rows: int, columns: int, itemsize: int) -> int: + # 沿用TensorBufferManager的256对齐规则 + return align(rows * columns * itemsize, 256) + + chunk_rows = 128 + hidden_itemsize = hidden_dtype.itemsize + scale_cols = moe_intermediate_size // 128 + + silu_bytes = tensor_bytes(chunk_rows, moe_intermediate_size, hidden_itemsize) + gemm_out_a_bytes = tensor_bytes(chunk_rows, 2 * moe_intermediate_size, hidden_itemsize) + quant_bytes = tensor_bytes(chunk_rows, moe_intermediate_size, 1) + scale_bytes = tensor_bytes(chunk_rows, scale_cols, torch.float32.itemsize) + gemm_out_b_bytes = tensor_bytes(chunk_rows, hidden_size, hidden_itemsize) + + w1_peak_bytes = silu_bytes + gemm_out_a_bytes + + # silu 释放后,B 较小时可复用其 first-fit 空洞;B 更大则需另行预留完整空间。 + quant_w2_peak_bytes = ( + silu_bytes + quant_bytes + scale_bytes + (gemm_out_b_bytes if gemm_out_b_bytes > silu_bytes else 0) + ) + temporary_peak_bytes = max(w1_peak_bytes, quant_w2_peak_bytes) + + # 按 _get_max_chunk_rows 的 1024 行对齐规则 + gather_bytes = tensor_bytes(align(global_world_size * prefill_tokens_per_rank, 1024), hidden_size, hidden_itemsize) + per_workspace_bytes = gather_bytes + temporary_peak_bytes + + reuse_buffer_bytes = per_workspace_bytes * microbatch_count + return align(reuse_buffer_bytes, microbatch_count * 256) + + dist_group_manager = DistributeGroupManager()